x
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
"""
|
||||
核心模块 - 提供扩展性基础设施
|
||||
|
||||
包含:
|
||||
- Registry: 注册机制
|
||||
- State: 增强的状态管理
|
||||
- Provider: LLM Provider 抽象
|
||||
- Response: 统一响应格式
|
||||
"""
|
||||
|
||||
from .registry import (
|
||||
BaseRegistry,
|
||||
NodeRegistry,
|
||||
ToolRegistry,
|
||||
WorkflowRegistry,
|
||||
ProviderRegistry,
|
||||
RegistryEntry,
|
||||
ToolMetadata,
|
||||
WorkflowMetadata,
|
||||
ProviderMetadata,
|
||||
register_tool,
|
||||
register_workflow,
|
||||
register_provider,
|
||||
)
|
||||
from .state import AgentState, StateContext
|
||||
from .providers import LLMProvider, LLMFactory
|
||||
from .response import ApiResponse, StreamEvent
|
||||
|
||||
__all__ = [
|
||||
# Registry
|
||||
"BaseRegistry",
|
||||
"NodeRegistry",
|
||||
"ToolRegistry",
|
||||
"WorkflowRegistry",
|
||||
"ProviderRegistry",
|
||||
"RegistryEntry",
|
||||
"ToolMetadata",
|
||||
"WorkflowMetadata",
|
||||
"ProviderMetadata",
|
||||
"register_tool",
|
||||
"register_workflow",
|
||||
"register_provider",
|
||||
# State
|
||||
"AgentState",
|
||||
"StateContext",
|
||||
# Provider
|
||||
"LLMProvider",
|
||||
"LLMFactory",
|
||||
# Response
|
||||
"ApiResponse",
|
||||
"StreamEvent",
|
||||
]
|
||||
@@ -0,0 +1,218 @@
|
||||
"""
|
||||
LLM Provider 抽象层
|
||||
|
||||
支持多种 LLM 提供商的统一接口
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Protocol, Type, runtime_checkable
|
||||
from abc import ABC, abstractmethod
|
||||
import logging
|
||||
|
||||
from langchain_core.language_models import BaseChatModel
|
||||
|
||||
from config import Config
|
||||
from .registry import ProviderRegistry, register_provider
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class LLMProvider(Protocol):
|
||||
"""
|
||||
LLM Provider 协议
|
||||
|
||||
定义所有 LLM 提供商必须实现的接口
|
||||
"""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Provider 名称"""
|
||||
...
|
||||
|
||||
@property
|
||||
def supported_models(self) -> List[str]:
|
||||
"""支持的模型列表"""
|
||||
...
|
||||
|
||||
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||
"""
|
||||
创建 LLM 模型实例
|
||||
|
||||
Args:
|
||||
config: 模型配置
|
||||
|
||||
Returns:
|
||||
BaseChatModel 实例
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class BaseLLMProvider(ABC):
|
||||
"""LLM Provider 基类"""
|
||||
|
||||
def __init__(self, name: str, supported_models: List[str]):
|
||||
self._name = name
|
||||
self._supported_models = supported_models
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return self._name
|
||||
|
||||
@property
|
||||
def supported_models(self) -> List[str]:
|
||||
return self._supported_models
|
||||
|
||||
@abstractmethod
|
||||
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||
"""子类实现具体的模型创建逻辑"""
|
||||
pass
|
||||
|
||||
|
||||
class OpenAICompatibleProvider(BaseLLMProvider):
|
||||
"""
|
||||
OpenAI 兼容的 Provider
|
||||
|
||||
支持所有兼容 OpenAI API 的服务:
|
||||
- OpenAI
|
||||
- Azure OpenAI
|
||||
- 本地部署的兼容服务
|
||||
- 国产大模型(通义千问、文心一言等)
|
||||
"""
|
||||
|
||||
def __init__(self, name: str = "openai_compatible", supported_models: Optional[List[str]] = None):
|
||||
super().__init__(name, supported_models or [])
|
||||
|
||||
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||
from langchain_openai import ChatOpenAI
|
||||
|
||||
return ChatOpenAI(
|
||||
model=config.get("model_name") or config.get("MODEL_NAME"),
|
||||
openai_api_key=config.get("api_key") or config.get("OPENAI_API_KEY"),
|
||||
openai_api_base=config.get("base_url") or config.get("URL"),
|
||||
temperature=config.get("temperature", 0.7),
|
||||
max_tokens=config.get("max_tokens"),
|
||||
timeout=config.get("timeout", 30),
|
||||
max_retries=config.get("max_retries", 3),
|
||||
)
|
||||
|
||||
|
||||
class AzureOpenAIProvider(BaseLLMProvider):
|
||||
"""Azure OpenAI Provider"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__("azure", ["gpt-4", "gpt-35-turbo"])
|
||||
|
||||
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||
from langchain_openai import AzureChatOpenAI
|
||||
|
||||
return AzureChatOpenAI(
|
||||
azure_deployment=config.get("deployment_name"),
|
||||
openai_api_version=config.get("api_version", "2024-02-15-preview"),
|
||||
azure_endpoint=config.get("endpoint"),
|
||||
openai_api_key=config.get("api_key"),
|
||||
temperature=config.get("temperature", 0.7),
|
||||
)
|
||||
|
||||
|
||||
class LLMFactory:
|
||||
"""
|
||||
LLM 工厂类
|
||||
|
||||
统一管理 LLM 实例的创建,支持:
|
||||
- 多种 Provider
|
||||
- 配置驱动
|
||||
- 单例缓存
|
||||
"""
|
||||
|
||||
_instances: Dict[str, BaseChatModel] = {}
|
||||
_default_provider: str = "openai_compatible"
|
||||
|
||||
@classmethod
|
||||
def register_provider(cls, name: str, provider: LLMProvider) -> None:
|
||||
"""注册 Provider"""
|
||||
ProviderRegistry._entries[name] = type(
|
||||
"RegistryEntry",
|
||||
(),
|
||||
{"instance": provider, "metadata": {}}
|
||||
)()
|
||||
logger.info(f"Registered LLM provider: {name}")
|
||||
|
||||
@classmethod
|
||||
def get_provider(cls, name: str) -> Optional[LLMProvider]:
|
||||
"""获取 Provider"""
|
||||
entry = ProviderRegistry.get_entry(name)
|
||||
if entry:
|
||||
return entry.instance
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
provider: Optional[str] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
model_section: Optional[str] = None,
|
||||
use_cache: bool = True,
|
||||
) -> BaseChatModel:
|
||||
"""
|
||||
创建 LLM 实例
|
||||
|
||||
Args:
|
||||
provider: Provider 名称,默认使用 openai_compatible
|
||||
config: 模型配置
|
||||
model_section: 配置文件中的模型段名
|
||||
use_cache: 是否使用缓存
|
||||
|
||||
Returns:
|
||||
BaseChatModel 实例
|
||||
"""
|
||||
provider_name = provider or cls._default_provider
|
||||
|
||||
if model_section:
|
||||
config = Config.get_section(model_section)
|
||||
cache_key = f"{provider_name}:{model_section}"
|
||||
else:
|
||||
config = config or {}
|
||||
cache_key = f"{provider_name}:{hash(frozenset(config.items()))}"
|
||||
|
||||
if use_cache and cache_key in cls._instances:
|
||||
return cls._instances[cache_key]
|
||||
|
||||
provider_instance = cls.get_provider(provider_name)
|
||||
|
||||
if provider_instance is None:
|
||||
provider_instance = OpenAICompatibleProvider()
|
||||
|
||||
model = provider_instance.create_model(config)
|
||||
|
||||
if use_cache:
|
||||
cls._instances[cache_key] = model
|
||||
|
||||
return model
|
||||
|
||||
@classmethod
|
||||
def clear_cache(cls) -> None:
|
||||
"""清空缓存"""
|
||||
cls._instances.clear()
|
||||
|
||||
@classmethod
|
||||
def list_providers(cls) -> List[str]:
|
||||
"""列出所有注册的 Provider"""
|
||||
return ProviderRegistry.list_names()
|
||||
|
||||
|
||||
OpenAICompatibleProvider()
|
||||
ProviderRegistry._entries["openai_compatible"] = type(
|
||||
"RegistryEntry",
|
||||
(),
|
||||
{"instance": OpenAICompatibleProvider(), "metadata": {}}
|
||||
)()
|
||||
|
||||
|
||||
def create_chat_model(model_section: Optional[str] = None) -> BaseChatModel:
|
||||
"""
|
||||
创建聊天模型(兼容现有代码)
|
||||
|
||||
这是现有 llm_factory.create_chat_model 的替代实现,
|
||||
使用新的 Provider 架构但保持接口兼容
|
||||
"""
|
||||
return LLMFactory.create(model_section=model_section)
|
||||
@@ -0,0 +1,222 @@
|
||||
"""
|
||||
核心注册机制模块
|
||||
|
||||
提供统一的注册器模式,支持动态扩展:
|
||||
- NodeRegistry: 节点注册器
|
||||
- ToolRegistry: 工具注册器
|
||||
- WorkflowRegistry: 工作流注册器
|
||||
- ProviderRegistry: LLM Provider 注册器
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Type, TypeVar, Protocol, runtime_checkable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
import time
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Registerable(Protocol[T]):
|
||||
"""可注册对象的协议"""
|
||||
name: str
|
||||
description: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegistryEntry:
|
||||
"""注册条目"""
|
||||
name: str
|
||||
instance: Any
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
registered_at: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class BaseRegistry:
|
||||
"""基础注册器"""
|
||||
|
||||
_entries: Dict[str, RegistryEntry] = {}
|
||||
|
||||
@classmethod
|
||||
def register(cls, name: str, metadata: Optional[Dict[str, Any]] = None):
|
||||
"""
|
||||
注册装饰器
|
||||
|
||||
Usage:
|
||||
@NodeRegistry.register("my_node", metadata={"category": "processing"})
|
||||
def my_node(state: AgentState) -> AgentState:
|
||||
...
|
||||
"""
|
||||
def decorator(obj: T) -> T:
|
||||
entry = RegistryEntry(
|
||||
name=name,
|
||||
instance=obj,
|
||||
metadata=metadata or {},
|
||||
)
|
||||
cls._entries[name] = entry
|
||||
return obj
|
||||
return decorator
|
||||
|
||||
@classmethod
|
||||
def get(cls, name: str) -> Optional[Any]:
|
||||
"""获取注册的对象"""
|
||||
entry = cls._entries.get(name)
|
||||
return entry.instance if entry else None
|
||||
|
||||
@classmethod
|
||||
def get_entry(cls, name: str) -> Optional[RegistryEntry]:
|
||||
"""获取注册条目(包含元数据)"""
|
||||
return cls._entries.get(name)
|
||||
|
||||
@classmethod
|
||||
def list_names(cls) -> List[str]:
|
||||
"""列出所有注册名称"""
|
||||
return list(cls._entries.keys())
|
||||
|
||||
@classmethod
|
||||
def list_entries(cls) -> List[RegistryEntry]:
|
||||
"""列出所有注册条目"""
|
||||
return list(cls._entries.values())
|
||||
|
||||
@classmethod
|
||||
def unregister(cls, name: str) -> bool:
|
||||
"""注销注册"""
|
||||
if name in cls._entries:
|
||||
del cls._entries[name]
|
||||
return True
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def clear(cls):
|
||||
"""清空注册表"""
|
||||
cls._entries.clear()
|
||||
|
||||
|
||||
class NodeRegistry(BaseRegistry):
|
||||
"""节点注册器 - 用于注册 Agent 工作流节点"""
|
||||
_entries: Dict[str, RegistryEntry] = {}
|
||||
|
||||
|
||||
class ToolRegistry(BaseRegistry):
|
||||
"""工具注册器 - 用于注册工具"""
|
||||
_entries: Dict[str, RegistryEntry] = {}
|
||||
|
||||
|
||||
class WorkflowRegistry(BaseRegistry):
|
||||
"""工作流注册器 - 用于注册工作流类型"""
|
||||
_entries: Dict[str, RegistryEntry] = {}
|
||||
|
||||
|
||||
class ProviderRegistry(BaseRegistry):
|
||||
"""LLM Provider 注册器"""
|
||||
_entries: Dict[str, RegistryEntry] = {}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolMetadata:
|
||||
"""工具元数据"""
|
||||
name: str
|
||||
description: str
|
||||
version: str = "1.0.0"
|
||||
timeout: int = 30
|
||||
retry: int = 0
|
||||
parameters_schema: Optional[Dict[str, Any]] = None
|
||||
requires_auth: bool = False
|
||||
tags: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class WorkflowMetadata:
|
||||
"""工作流元数据"""
|
||||
name: str
|
||||
description: str
|
||||
version: str = "1.0.0"
|
||||
agent_class: Optional[Type] = None
|
||||
default_model: Optional[str] = None
|
||||
supported_features: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderMetadata:
|
||||
"""LLM Provider 元数据"""
|
||||
name: str
|
||||
description: str
|
||||
provider_type: str
|
||||
supported_models: List[str] = field(default_factory=list)
|
||||
config_schema: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
def register_tool(
|
||||
name: str,
|
||||
description: str = "",
|
||||
version: str = "1.0.0",
|
||||
timeout: int = 30,
|
||||
retry: int = 0,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
"""
|
||||
工具注册装饰器
|
||||
|
||||
Usage:
|
||||
@register_tool("calculator", "数学计算", timeout=10)
|
||||
class CalculatorTool(BaseTool):
|
||||
...
|
||||
"""
|
||||
metadata = ToolMetadata(
|
||||
name=name,
|
||||
description=description,
|
||||
version=version,
|
||||
timeout=timeout,
|
||||
retry=retry,
|
||||
tags=tags or [],
|
||||
)
|
||||
return ToolRegistry.register(name, {"tool_metadata": metadata})
|
||||
|
||||
|
||||
def register_workflow(
|
||||
name: str,
|
||||
description: str = "",
|
||||
version: str = "1.0.0",
|
||||
default_model: Optional[str] = None,
|
||||
supported_features: Optional[List[str]] = None,
|
||||
):
|
||||
"""
|
||||
工作流注册装饰器
|
||||
|
||||
Usage:
|
||||
@register_workflow("data_query", "数据查询工作流")
|
||||
class DataQueryAgent(BaseAgent):
|
||||
...
|
||||
"""
|
||||
metadata = WorkflowMetadata(
|
||||
name=name,
|
||||
description=description,
|
||||
version=version,
|
||||
default_model=default_model,
|
||||
supported_features=supported_features or [],
|
||||
)
|
||||
return WorkflowRegistry.register(name, {"workflow_metadata": metadata})
|
||||
|
||||
|
||||
def register_provider(
|
||||
name: str,
|
||||
provider_type: str,
|
||||
description: str = "",
|
||||
supported_models: Optional[List[str]] = None,
|
||||
):
|
||||
"""
|
||||
LLM Provider 注册装饰器
|
||||
|
||||
Usage:
|
||||
@register_provider("openai", "openai", supported_models=["gpt-4", "gpt-3.5"])
|
||||
class OpenAIProvider:
|
||||
...
|
||||
"""
|
||||
metadata = ProviderMetadata(
|
||||
name=name,
|
||||
description=description,
|
||||
provider_type=provider_type,
|
||||
supported_models=supported_models or [],
|
||||
)
|
||||
return ProviderRegistry.register(name, {"provider_metadata": metadata})
|
||||
@@ -0,0 +1,248 @@
|
||||
"""
|
||||
统一响应格式模块
|
||||
|
||||
提供标准化的 API 响应和流式事件格式
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Generic, List, Optional, TypeVar, Literal
|
||||
from pydantic import BaseModel, Field
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class ApiResponse(BaseModel, Generic[T]):
|
||||
"""
|
||||
统一 API 响应格式
|
||||
|
||||
所有 API 响应都使用这个格式,提供一致的响应结构
|
||||
|
||||
Usage:
|
||||
@router.get("/users/{user_id}")
|
||||
async def get_user(user_id: str) -> ApiResponse[User]:
|
||||
user = await user_service.get(user_id)
|
||||
return ApiResponse.success(data=user)
|
||||
"""
|
||||
|
||||
code: str = Field(default="success", description="响应代码")
|
||||
message: str = Field(default="", description="响应消息")
|
||||
data: Optional[T] = Field(default=None, description="响应数据")
|
||||
trace_id: Optional[str] = Field(default=None, description="追踪ID")
|
||||
timestamp: int = Field(
|
||||
default_factory=lambda: int(time.time() * 1000),
|
||||
description="时间戳(毫秒)"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def success(cls, data: T = None, message: str = "", trace_id: Optional[str] = None) -> "ApiResponse[T]":
|
||||
"""创建成功响应"""
|
||||
return cls(
|
||||
code="success",
|
||||
message=message,
|
||||
data=data,
|
||||
trace_id=trace_id or uuid.uuid4().hex,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def error(
|
||||
cls,
|
||||
code: str = "error",
|
||||
message: str = "",
|
||||
data: T = None,
|
||||
trace_id: Optional[str] = None,
|
||||
) -> "ApiResponse[T]":
|
||||
"""创建错误响应"""
|
||||
return cls(
|
||||
code=code,
|
||||
message=message,
|
||||
data=data,
|
||||
trace_id=trace_id or uuid.uuid4().hex,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_exception(cls, exc: Exception, trace_id: Optional[str] = None) -> "ApiResponse[None]":
|
||||
"""从异常创建错误响应"""
|
||||
return cls.error(
|
||||
code="internal_error",
|
||||
message=str(exc),
|
||||
trace_id=trace_id,
|
||||
)
|
||||
|
||||
def is_success(self) -> bool:
|
||||
"""判断是否成功"""
|
||||
return self.code == "success"
|
||||
|
||||
|
||||
class PagedResponse(BaseModel, Generic[T]):
|
||||
"""
|
||||
分页响应格式
|
||||
|
||||
用于返回分页数据
|
||||
"""
|
||||
|
||||
items: List[T] = Field(default_factory=list, description="数据列表")
|
||||
total: int = Field(default=0, description="总数")
|
||||
page: int = Field(default=1, description="当前页")
|
||||
page_size: int = Field(default=20, description="每页大小")
|
||||
total_pages: int = Field(default=0, description="总页数")
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
items: List[T],
|
||||
total: int,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
) -> "PagedResponse[T]":
|
||||
"""创建分页响应"""
|
||||
total_pages = (total + page_size - 1) // page_size if page_size > 0 else 0
|
||||
return cls(
|
||||
items=items,
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
total_pages=total_pages,
|
||||
)
|
||||
|
||||
|
||||
class StreamEvent(BaseModel):
|
||||
"""
|
||||
流式响应事件
|
||||
|
||||
用于 SSE (Server-Sent Events) 流式响应
|
||||
|
||||
Usage:
|
||||
async def event_stream():
|
||||
yield StreamEvent(event="start", data="Processing started")
|
||||
# ... 处理逻辑
|
||||
yield StreamEvent(event="result", data=json.dumps(result))
|
||||
yield StreamEvent(event="done", data="")
|
||||
"""
|
||||
|
||||
event: str = Field(..., description="事件类型")
|
||||
data: str = Field(default="", description="事件数据")
|
||||
event_id: Optional[str] = Field(default=None, description="事件ID")
|
||||
retry: Optional[int] = Field(default=None, description="重试间隔(毫秒)")
|
||||
|
||||
def to_sse(self) -> str:
|
||||
"""转换为 SSE 格式字符串"""
|
||||
lines = [f"event: {self.event}"]
|
||||
if self.event_id:
|
||||
lines.append(f"id: {self.event_id}")
|
||||
if self.retry:
|
||||
lines.append(f"retry: {self.retry}")
|
||||
lines.append(f"data: {self.data}")
|
||||
lines.append("")
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
@classmethod
|
||||
def message(cls, data: str, event_id: Optional[str] = None) -> "StreamEvent":
|
||||
"""创建消息事件"""
|
||||
return cls(event="message", data=data, event_id=event_id)
|
||||
|
||||
@classmethod
|
||||
def done(cls) -> "StreamEvent":
|
||||
"""创建完成事件"""
|
||||
return cls(event="done", data="[DONE]")
|
||||
|
||||
@classmethod
|
||||
def error(cls, message: str) -> "StreamEvent":
|
||||
"""创建错误事件"""
|
||||
return cls(event="error", data=message)
|
||||
|
||||
|
||||
class WorkflowEvent(BaseModel):
|
||||
"""
|
||||
工作流事件
|
||||
|
||||
用于工作流执行过程中的状态通知
|
||||
"""
|
||||
|
||||
workflow_id: str = Field(..., description="工作流ID")
|
||||
event_type: Literal[
|
||||
"started",
|
||||
"node_started",
|
||||
"node_completed",
|
||||
"node_failed",
|
||||
"completed",
|
||||
"failed",
|
||||
] = Field(..., description="事件类型")
|
||||
node_name: Optional[str] = Field(None, description="节点名称")
|
||||
data: Optional[Dict[str, Any]] = Field(None, description="事件数据")
|
||||
error: Optional[str] = Field(None, description="错误信息")
|
||||
timestamp: int = Field(
|
||||
default_factory=lambda: int(time.time() * 1000),
|
||||
description="时间戳"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def started(cls, workflow_id: str) -> "WorkflowEvent":
|
||||
"""创建开始事件"""
|
||||
return cls(workflow_id=workflow_id, event_type="started")
|
||||
|
||||
@classmethod
|
||||
def node_started(cls, workflow_id: str, node_name: str) -> "WorkflowEvent":
|
||||
"""创建节点开始事件"""
|
||||
return cls(
|
||||
workflow_id=workflow_id,
|
||||
event_type="node_started",
|
||||
node_name=node_name,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def node_completed(
|
||||
cls,
|
||||
workflow_id: str,
|
||||
node_name: str,
|
||||
data: Optional[Dict[str, Any]] = None,
|
||||
) -> "WorkflowEvent":
|
||||
"""创建节点完成事件"""
|
||||
return cls(
|
||||
workflow_id=workflow_id,
|
||||
event_type="node_completed",
|
||||
node_name=node_name,
|
||||
data=data,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def completed(
|
||||
cls,
|
||||
workflow_id: str,
|
||||
data: Optional[Dict[str, Any]] = None,
|
||||
) -> "WorkflowEvent":
|
||||
"""创建完成事件"""
|
||||
return cls(workflow_id=workflow_id, event_type="completed", data=data)
|
||||
|
||||
@classmethod
|
||||
def failed(
|
||||
cls,
|
||||
workflow_id: str,
|
||||
error: str,
|
||||
node_name: Optional[str] = None,
|
||||
) -> "WorkflowEvent":
|
||||
"""创建失败事件"""
|
||||
return cls(
|
||||
workflow_id=workflow_id,
|
||||
event_type="failed",
|
||||
node_name=node_name,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
class ErrorCode:
|
||||
"""错误代码常量"""
|
||||
|
||||
SUCCESS = "success"
|
||||
UNKNOWN_ERROR = "unknown_error"
|
||||
INVALID_REQUEST = "invalid_request"
|
||||
INVALID_WORKFLOW_TYPE = "invalid_workflow_type"
|
||||
SQL_GENERATION_FAILED = "sql_generation_failed"
|
||||
TOOL_NOT_FOUND = "tool_not_found"
|
||||
TOOL_EXECUTION_FAILED = "tool_execution_failed"
|
||||
INTERNAL_ERROR = "internal_error"
|
||||
TIMEOUT = "timeout"
|
||||
RATE_LIMITED = "rate_limited"
|
||||
UNAUTHORIZED = "unauthorized"
|
||||
+199
@@ -0,0 +1,199 @@
|
||||
"""
|
||||
增强的状态管理模块
|
||||
|
||||
使用 Pydantic 提供类型安全和验证
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Literal
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage, SystemMessage
|
||||
|
||||
|
||||
class StateContext(BaseModel):
|
||||
"""状态上下文 - 存储工作流执行过程中的数据"""
|
||||
|
||||
original_input: Optional[str] = Field(None, description="用户原始输入")
|
||||
normalized_input: Optional[str] = Field(None, description="规范化后的输入")
|
||||
intent: Optional[str] = Field(None, description="识别的意图")
|
||||
table_match: Optional[Dict[str, Any]] = Field(None, description="表名匹配结果")
|
||||
final_sql: Optional[str] = Field(None, description="生成的 SQL")
|
||||
sr_api_result: Optional[Any] = Field(None, description="API 执行结果")
|
||||
|
||||
class Config:
|
||||
extra = "allow"
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
"""获取上下文值"""
|
||||
return getattr(self, key, default)
|
||||
|
||||
def set(self, key: str, value: Any) -> None:
|
||||
"""设置上下文值"""
|
||||
setattr(self, key, value)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""转换为字典"""
|
||||
return self.model_dump(exclude_none=True)
|
||||
|
||||
|
||||
class AgentState(BaseModel):
|
||||
"""
|
||||
Agent 工作流状态定义
|
||||
|
||||
使用 Pydantic 提供类型安全和验证
|
||||
"""
|
||||
|
||||
messages: List[BaseMessage] = Field(default_factory=list, description="消息历史")
|
||||
current_step: str = Field(default="start", description="当前步骤")
|
||||
context: StateContext = Field(default_factory=StateContext, description="上下文数据")
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
"extra": "forbid",
|
||||
}
|
||||
|
||||
@field_validator("messages", mode="before")
|
||||
@classmethod
|
||||
def validate_messages(cls, v):
|
||||
"""验证并转换消息列表"""
|
||||
if not isinstance(v, list):
|
||||
return []
|
||||
|
||||
result = []
|
||||
for msg in v:
|
||||
if isinstance(msg, BaseMessage):
|
||||
result.append(msg)
|
||||
elif isinstance(msg, dict):
|
||||
msg_type = msg.get("type", "human")
|
||||
content = msg.get("content", "")
|
||||
if msg_type == "human":
|
||||
result.append(HumanMessage(content=content))
|
||||
elif msg_type == "ai":
|
||||
result.append(AIMessage(content=content))
|
||||
elif msg_type == "system":
|
||||
result.append(SystemMessage(content=content))
|
||||
return result
|
||||
|
||||
def add_message(self, message: BaseMessage) -> "AgentState":
|
||||
"""添加消息并返回新状态"""
|
||||
return AgentState(
|
||||
messages=[*self.messages, message],
|
||||
current_step=self.current_step,
|
||||
context=self.context,
|
||||
)
|
||||
|
||||
def add_human_message(self, content: str) -> "AgentState":
|
||||
"""添加用户消息"""
|
||||
return self.add_message(HumanMessage(content=content))
|
||||
|
||||
def add_ai_message(self, content: str) -> "AgentState":
|
||||
"""添加 AI 消息"""
|
||||
return self.add_message(AIMessage(content=content))
|
||||
|
||||
def update_step(self, step: str) -> "AgentState":
|
||||
"""更新当前步骤"""
|
||||
return AgentState(
|
||||
messages=self.messages,
|
||||
current_step=step,
|
||||
context=self.context,
|
||||
)
|
||||
|
||||
def update_context(self, **kwargs) -> "AgentState":
|
||||
"""更新上下文"""
|
||||
new_context = self.context.model_copy()
|
||||
for key, value in kwargs.items():
|
||||
new_context.set(key, value)
|
||||
return AgentState(
|
||||
messages=self.messages,
|
||||
current_step=self.current_step,
|
||||
context=new_context,
|
||||
)
|
||||
|
||||
def get_last_message(self) -> Optional[BaseMessage]:
|
||||
"""获取最后一条消息"""
|
||||
return self.messages[-1] if self.messages else None
|
||||
|
||||
def get_context(self, key: str, default: Any = None) -> Any:
|
||||
"""获取上下文值"""
|
||||
return self.context.get(key, default)
|
||||
|
||||
def to_legacy_format(self) -> Dict[str, Any]:
|
||||
"""
|
||||
转换为旧格式(兼容现有代码)
|
||||
|
||||
现有代码期望 state 是一个可修改的对象,
|
||||
这个方法返回一个兼容的字典格式
|
||||
"""
|
||||
return {
|
||||
"messages": self.messages,
|
||||
"current_step": self.current_step,
|
||||
"context": self.context.to_dict(),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_legacy_format(cls, data: Dict[str, Any]) -> "AgentState":
|
||||
"""从旧格式创建"""
|
||||
context_data = data.get("context", {})
|
||||
if isinstance(context_data, StateContext):
|
||||
context = context_data
|
||||
else:
|
||||
context = StateContext(**context_data) if context_data else StateContext()
|
||||
|
||||
return cls(
|
||||
messages=data.get("messages", []),
|
||||
current_step=data.get("current_step", "start"),
|
||||
context=context,
|
||||
)
|
||||
|
||||
|
||||
class MutableAgentState:
|
||||
"""
|
||||
可变的 Agent 状态包装器
|
||||
|
||||
用于兼容现有代码中直接修改 state 的模式
|
||||
"""
|
||||
|
||||
def __init__(self, state: Optional[AgentState] = None):
|
||||
self._state = state or AgentState()
|
||||
self._context_overrides: Dict[str, Any] = {}
|
||||
|
||||
@property
|
||||
def messages(self) -> List[BaseMessage]:
|
||||
return self._state.messages
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value: List[BaseMessage]):
|
||||
self._state = AgentState(
|
||||
messages=value,
|
||||
current_step=self._state.current_step,
|
||||
context=self._state.context,
|
||||
)
|
||||
|
||||
@property
|
||||
def current_step(self) -> str:
|
||||
return self._state.current_step
|
||||
|
||||
@current_step.setter
|
||||
def current_step(self, value: str):
|
||||
self._state = AgentState(
|
||||
messages=self._state.messages,
|
||||
current_step=value,
|
||||
context=self._state.context,
|
||||
)
|
||||
|
||||
@property
|
||||
def context(self) -> Dict[str, Any]:
|
||||
"""返回可修改的上下文字典"""
|
||||
result = self._state.context.to_dict()
|
||||
result.update(self._context_overrides)
|
||||
return result
|
||||
|
||||
def to_immutable(self) -> AgentState:
|
||||
"""转换为不可变状态"""
|
||||
context = self._state.context.model_copy()
|
||||
for key, value in self._context_overrides.items():
|
||||
context.set(key, value)
|
||||
return AgentState(
|
||||
messages=self._state.messages,
|
||||
current_step=self._state.current_step,
|
||||
context=context,
|
||||
)
|
||||
Reference in New Issue
Block a user