This commit is contained in:
2026-03-11 23:40:39 +08:00
parent db25d61026
commit e062368ef2
15 changed files with 1592 additions and 229 deletions
+52
View File
@@ -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",
]
+218
View File
@@ -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)
+222
View File
@@ -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})
+248
View File
@@ -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
View File
@@ -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,
)