diff --git a/api/endpoints.py b/api/endpoints.py index 6e2ef0d..9173279 100644 --- a/api/endpoints.py +++ b/api/endpoints.py @@ -13,6 +13,7 @@ from schemas.tool_input import ToolInput from schemas.tool_output import ToolOutput from schemas.chat_message_response import ChatMessageResponseDTO from schemas.chat_message_request import ChatMessageRequestDTO +from schemas.super_agent import SuperAgentRequest, SuperAgentResponse, SuperAgentStreamEvent from workflows.workflow_manager import WorkflowType from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager from services.app_errors import AppError, ErrorCode @@ -190,6 +191,54 @@ def run_workflow_stream(payload: ChatMessageRequestDTO, workflow_manager=Depends return StreamingResponse(event_stream(), media_type="text/event-stream") +@router.get("/api/workflows/list") +def list_workflows(workflow_manager=Depends(get_workflow_manager)): + """列出所有可用工作流""" + workflows = workflow_manager.get_available_workflows() + result = [] + for name in workflows: + info = workflow_manager.get_workflow_info(name) + if info: + result.append(info) + return {"workflows": result} + + +@router.get("/api/workflows/{workflow_name}") +def get_workflow_detail(workflow_name: str, workflow_manager=Depends(get_workflow_manager)): + """获取工作流详情""" + info = workflow_manager.get_workflow_info(workflow_name) + if not info: + raise HTTPException(status_code=404, detail=f"工作流不存在: {workflow_name}") + return info + + +@router.get("/api/tools/list") +def list_tools(tool_router=Depends(get_tool_router)): + """列出所有可用工具""" + tools = tool_router.list_tools() + result = [] + for name in tools: + info = tool_router.get_tool_info(name) + if info: + result.append(info) + return {"tools": result} + + +@router.get("/api/tools/{tool_name}") +def get_tool_detail(tool_name: str, tool_router=Depends(get_tool_router)): + """获取工具详情""" + info = tool_router.get_tool_info(tool_name) + if not info: + raise HTTPException(status_code=404, detail=f"工具不存在: {tool_name}") + return info + + +@router.get("/api/tools/stats") +def get_tools_stats(tool_router=Depends(get_tool_router)): + """获取工具执行统计""" + return tool_router.get_all_stats() + + @router.post("/api/tools/execute", response_model=ToolOutput) def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)): result = tool_router.call(payload.tool_name, payload.payload) @@ -251,3 +300,135 @@ def update_sql_gen(): return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) + + +@router.post("/api/super-agent/query", response_model=SuperAgentResponse) +def super_agent_query(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)): + """Super Agent 同步查询接口""" + trace_id = uuid.uuid4().hex + slog = get_structured_logger() + slog.log("INFO", "super_agent.query.start", trace_id, { + "query": payload.query[:100], + "workflow_type": payload.workflow_type, + "user_id": payload.user_id, + }) + + conversation_id = payload.conversation_id or uuid.uuid4().hex + + try: + workflow_type = _resolve_workflow_type(payload.workflow_type) + except Exception as e: + slog.log("ERROR", "super_agent.query.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value) + return SuperAgentResponse( + conversation_id=conversation_id, + workflow_type=payload.workflow_type, + status="error", + error=f"不支持的工作流类型: {payload.workflow_type}", + ) + + try: + result = workflow_manager.execute_workflow( + workflow_type=workflow_type, + user_input=payload.query, + session_id=conversation_id, + ) + + context = (result.get("result") or {}).get("context") or {} + sql_text = context.get("final_sql") + sr_api_result = context.get("sr_api_result") + + slog.log("INFO", "super_agent.query.success", trace_id, { + "conversation_id": conversation_id, + "has_sql": bool(sql_text), + "has_result": bool(sr_api_result), + }) + + return SuperAgentResponse( + conversation_id=conversation_id, + workflow_type=workflow_type.value, + status="success", + sql=sql_text, + result=str(sr_api_result) if sr_api_result else None, + metadata={"trace_id": trace_id}, + ) + + except Exception as e: + slog.log("ERROR", "super_agent.query.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) + return SuperAgentResponse( + conversation_id=conversation_id, + workflow_type=payload.workflow_type, + status="error", + error=str(e), + metadata={"trace_id": trace_id}, + ) + + +@router.post("/api/super-agent/stream") +def super_agent_stream(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)): + """Super Agent 流式查询接口""" + trace_id = uuid.uuid4().hex + slog = get_structured_logger() + stream_cfg = Config.get_section("stream") + progress_interval = float(stream_cfg.get("progress_interval", 0.3)) + + conversation_id = payload.conversation_id or uuid.uuid4().hex + + def _build_sse_event(event: str, data: str) -> str: + dto = SuperAgentStreamEvent( + conversation_id=conversation_id, + event=event, + data=data, + timestamp=int(time.time() * 1000), + ) + return f"event: {event}\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n" + + async def event_stream(): + try: + slog.log("INFO", "super_agent.stream.start", trace_id, { + "query": payload.query[:100], + "user_id": payload.user_id, + }) + + workflow_type = _resolve_workflow_type(payload.workflow_type) + + result = await asyncio.to_thread( + workflow_manager.execute_workflow, + workflow_type, + payload.query, + conversation_id, + skip_sr_api=True, + ) + + context = (result.get("result") or {}).get("context") or {} + sql_text = context.get("final_sql") + + if not sql_text: + slog.log("ERROR", "super_agent.stream.sql_failed", trace_id, error_code=ErrorCode.SQL_GENERATION_FAILED.value) + yield _build_sse_event("error", "SQL 生成失败") + yield _build_sse_event("done", "") + return + + yield _build_sse_event("sql_generated", sql_text) + + yield _build_sse_event("sql_executing", "") + + tool = SrApiQueryTool() + task = asyncio.create_task( + asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, ensure_ascii=False)) + ) + + while not task.done(): + yield _build_sse_event("sql_executing", "") + await asyncio.sleep(progress_interval) + + sql_result = await task + slog.log("INFO", "super_agent.stream.success", trace_id, {"result_len": len(str(sql_result))}) + yield _build_sse_event("result", str(sql_result)) + + except Exception as e: + slog.log("ERROR", "super_agent.stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) + yield _build_sse_event("error", str(e)) + + yield _build_sse_event("done", "") + + return StreamingResponse(event_stream(), media_type="text/event-stream") diff --git a/config/config.ini b/config/config.ini index 3541dc6..75a752c 100644 --- a/config/config.ini +++ b/config/config.ini @@ -35,28 +35,9 @@ cache_ttl = 600 table_retrieval_dataset_id = 9945baf512ea11f18ccb6a681b3130b2 sql_gen_dataset_id = ee68f53a12ec11f18e436a681b3130b2 -[redis] -enabled = true -host = led-redis.lenovo.com -port = 30398 -password = bgs123456 -database = 0 -sql_prompt_ttl = 600 - [stream] progress_interval = 0.3 -[logging_mysql] -enabled = false -host = 127.0.0.1 -port = 3306 -user = root -password = -database = more_dots -table = structured_logs -connect_timeout = 5 - - [app] service_name = local-model-streaming-api host = 0.0.0.0 diff --git a/config/config.ini.example b/config/config.ini.example index 5a6c60e..d03f23a 100644 --- a/config/config.ini.example +++ b/config/config.ini.example @@ -38,29 +38,10 @@ retrieval_top_k = 3 table_retrieval_dataset_id = sql_gen_dataset_id = -[redis] -# 是否启用 Redis 缓存(用于 sql_gen_prompts) -enabled = false -url = redis://localhost:6379/0 -db = 0 -# SQL 提示词缓存过期秒数 -sql_prompt_ttl = 600 - [stream] # /api/workflows/stream 进度事件间隔(秒) progress_interval = 0.3 -[logging_mysql] -# 是否启用结构化日志写入 MySQL -enabled = false -host = 127.0.0.1 -port = 3306 -user = root -password = -database = more_dots -table = structured_logs -connect_timeout = 5 - [nacos] # 是否启用 Nacos 注册 enabled = false diff --git a/core/__init__.py b/core/__init__.py new file mode 100644 index 0000000..864400b --- /dev/null +++ b/core/__init__.py @@ -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", +] diff --git a/core/providers.py b/core/providers.py new file mode 100644 index 0000000..2f50419 --- /dev/null +++ b/core/providers.py @@ -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) diff --git a/core/registry.py b/core/registry.py new file mode 100644 index 0000000..29407e3 --- /dev/null +++ b/core/registry.py @@ -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}) diff --git a/core/response.py b/core/response.py new file mode 100644 index 0000000..12bd13b --- /dev/null +++ b/core/response.py @@ -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" diff --git a/core/state.py b/core/state.py new file mode 100644 index 0000000..4b2b509 --- /dev/null +++ b/core/state.py @@ -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, + ) diff --git a/requirements.txt b/requirements.txt index d1158fa..b792430 100644 --- a/requirements.txt +++ b/requirements.txt @@ -8,5 +8,3 @@ uvicorn>=0.30.0 nacos-sdk-python==2.0.9 httpx>=0.27.0 pyyaml>=6.0.1 -redis>=5.0.0 -pymysql>=1.1.1 diff --git a/schemas/super_agent.py b/schemas/super_agent.py new file mode 100644 index 0000000..3c93e6d --- /dev/null +++ b/schemas/super_agent.py @@ -0,0 +1,37 @@ +from typing import Any, Dict, Optional + +from pydantic import BaseModel, ConfigDict, Field + + +class SuperAgentRequest(BaseModel): + """Super Agent 请求模型""" + + model_config = ConfigDict(extra="forbid") + + query: str = Field(..., desc ription="用户查询") + conversation_id: Optional[str] = Field(None, description="会话ID") + user_id: Optional[str] = Field(None, description="用户ID") + workflow_type: str = Field(default="conversation", description="工作流类型") + context: Dict[str, str] = Field(default_factory=dict, description="上下文信息") + timeout_seconds: int = Field(default=30, description="超时时间(秒)") + + +class SuperAgentResponse(BaseModel): + """Super Agent 响应模型""" + + conversation_id: str = Field(..., description="会话ID") + workflow_type: str = Field(..., description="工作流类型") + status: str = Field(default="success", description="状态: success/error") + sql: Optional[str] = Field(None, description="生成的SQL") + result: Optional[str] = Field(None, description="查询结果") + error: Optional[str] = Field(None, description="错误信息") + metadata: Dict[str, str] = Field(default_factory=dict, description="元数据") + + +class SuperAgentStreamEvent(BaseModel): + """Super Agent 流式响应事件""" + + conversation_id: str = Field(..., description="会话ID") + event: str = Field(..., description="事件类型: sql_generated/sql_executing/result/error/done") + data: str = Field(..., description="事件数据") + timestamp: int = Field(..., description="时间戳(毫秒)") diff --git a/services/cache.py b/services/cache.py index 27f9515..2d3cd1a 100644 --- a/services/cache.py +++ b/services/cache.py @@ -2,11 +2,6 @@ from __future__ import annotations from typing import Optional -try: - import redis -except Exception: - redis = None - class CacheBase: """缓存接口""" @@ -26,20 +21,3 @@ class NoopCache(CacheBase): def set(self, key: str, value: str, ttl: int) -> None: return None - - -class RedisCache(CacheBase): - """Redis 缓存实现""" - - def __init__(self, url: str, db: int = 0): - if redis is None: - raise ImportError("未安装 redis 依赖") - self._client = redis.Redis.from_url(url, db=db, decode_responses=True) - - def get(self, key: str) -> Optional[str]: - return self._client.get(key) - - def set(self, key: str, value: str, ttl: int) -> None: - self._client.set(key, value, ex=ttl) - - diff --git a/services/sql_prompt_manager.py b/services/sql_prompt_manager.py index 62dbda1..5484b54 100644 --- a/services/sql_prompt_manager.py +++ b/services/sql_prompt_manager.py @@ -2,9 +2,6 @@ import json import os from typing import Any, Dict, Optional -from config import Config -from services.cache import NoopCache, RedisCache - class SqlPromptManager: """按表名读取 SQL 提示词""" @@ -12,51 +9,11 @@ class SqlPromptManager: def __init__(self, base_dir: Optional[str] = None): root_dir = os.path.dirname(os.path.dirname(__file__)) self._base_dir = base_dir or os.path.join(root_dir, "config", "sql_gen_prompts") - self._cache = self._init_cache() - self._cache_ttl = self._get_cache_ttl() - - @staticmethod - def _get_cache_ttl() -> int: - redis_cfg = Config.get_section("redis") - try: - return int(redis_cfg.get("sql_prompt_ttl", 600)) - except Exception: - return 600 - - @staticmethod - def _init_cache(): - redis_cfg = Config.get_section("redis") - enabled = str(redis_cfg.get("enabled", "false")).lower() in ("1", "true", "yes") - if not enabled: - return NoopCache() - - # 优先使用完整 URL;否则使用 host/port/password/database 拼接 - url = redis_cfg.get("url") - db = int(redis_cfg.get("db", redis_cfg.get("database", 0))) - if not url: - host = redis_cfg.get("host") - port = redis_cfg.get("port", "6379") - password = redis_cfg.get("password", "") - database = redis_cfg.get("database", str(db)) - if host: - auth = f":{password}@" if password else "" - url = f"redis://{auth}{host}:{port}/{database}" - - if not url: - return NoopCache() - try: - return RedisCache(url=url, db=db) - except Exception: - return NoopCache() @staticmethod def _safe_filename(name: str) -> str: return name.replace("..", "").replace("/", "_").replace("\\", "_") - @staticmethod - def _cache_key(table_name: str, mtime: float) -> str: - return f"sql_prompt:{table_name}:{int(mtime)}" - def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]: """读取指定表的提示词 JSON""" if not table_name: @@ -67,19 +24,9 @@ class SqlPromptManager: if not os.path.exists(path): return None - mtime = os.path.getmtime(path) - key = self._cache_key(safe_name, mtime) - cached = self._cache.get(key) - if cached: - try: - return json.loads(cached) - except Exception: - pass - with open(path, "r", encoding="utf-8") as f: prompt = json.load(f) - self._cache.set(key, json.dumps(prompt, ensure_ascii=False), self._cache_ttl) return prompt diff --git a/services/structured_logger.py b/services/structured_logger.py index 96ba37e..134b158 100644 --- a/services/structured_logger.py +++ b/services/structured_logger.py @@ -4,58 +4,10 @@ import json from datetime import datetime from typing import Any, Dict, Optional -import pymysql - -from config import Config - class StructuredLogger: def __init__(self): - cfg = Config.get_section("logging_mysql") - self.enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes") - self.host = cfg.get("host", "127.0.0.1") - self.port = int(cfg.get("port", 3306)) - self.user = cfg.get("user", "root") - self.password = cfg.get("password", "") - self.database = cfg.get("database", "more_dots") - self.table = cfg.get("table", "structured_logs") - self.connect_timeout = int(cfg.get("connect_timeout", 5)) - self._inited = False - - def _get_conn(self): - return pymysql.connect( - host=self.host, - port=self.port, - user=self.user, - password=self.password, - database=self.database, - charset="utf8mb4", - autocommit=True, - connect_timeout=self.connect_timeout, - ) - - def _ensure_table(self) -> None: - if self._inited or not self.enabled: - return - sql = f""" - CREATE TABLE IF NOT EXISTS {self.table} ( - id BIGINT PRIMARY KEY AUTO_INCREMENT, - trace_id VARCHAR(64) NOT NULL, - level VARCHAR(16) NOT NULL, - event VARCHAR(128) NOT NULL, - error_code VARCHAR(64) NULL, - payload JSON NULL, - created_at DATETIME NOT NULL - ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4; - """ - try: - with self._get_conn() as conn: - with conn.cursor() as cur: - cur.execute(sql) - self._inited = True - except Exception: - # 开发阶段容错,避免日志失败影响主流程 - self.enabled = False + pass def log(self, level: str, event: str, trace_id: str, payload: Optional[Dict[str, Any]] = None, error_code: Optional[str] = None) -> None: print(json.dumps({ @@ -67,32 +19,6 @@ class StructuredLogger: "created_at": datetime.now().isoformat(), }, ensure_ascii=False)) - if not self.enabled: - return - - self._ensure_table() - if not self.enabled: - return - - insert_sql = f"INSERT INTO {self.table}(trace_id, level, event, error_code, payload, created_at) VALUES(%s,%s,%s,%s,%s,%s)" - try: - with self._get_conn() as conn: - with conn.cursor() as cur: - cur.execute( - insert_sql, - ( - trace_id, - level, - event, - error_code, - json.dumps(payload or {}, ensure_ascii=False), - datetime.now(), - ), - ) - except Exception: - # 开发阶段容错,避免日志失败影响主流程 - return - _GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None diff --git a/services/tool_router.py b/services/tool_router.py index dbc6284..0e8f8cb 100644 --- a/services/tool_router.py +++ b/services/tool_router.py @@ -1,5 +1,13 @@ +""" +工具路由器模块 + +支持动态注册和管理工具 +""" + import json -from typing import Any, Dict, Optional +import logging +import time +from typing import Any, Callable, Dict, List, Optional, Type from langchain_core.tools import BaseTool @@ -7,26 +15,159 @@ from tools.calculator import CalculatorTool from tools.web_search import WebSearchTool from tools.rest_api_tool import RestApiTool from tools.sr_api_tool import SrApiQueryTool +from core.registry import ToolRegistry, ToolMetadata + +logger = logging.getLogger(__name__) class ToolRouter: - """工具路由器:统一调用入口""" - - def __init__(self, tools: Optional[list[BaseTool]] = None): - if tools is None: - tools = [CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool()] - self._tools: Dict[str, BaseTool] = {tool.name: tool for tool in tools} - - def list_tools(self) -> list[str]: + """ + 工具路由器:统一调用入口 + + 支持特性: + - 动态注册工具 + - 工具元数据管理 + - 执行监控 + """ + + def __init__(self, tools: Optional[List[BaseTool]] = None): + self._tools: Dict[str, BaseTool] = {} + self._tool_metadata: Dict[str, ToolMetadata] = {} + self._execution_stats: Dict[str, Dict[str, Any]] = {} + + if tools is not None: + for tool in tools: + self.register_tool(tool) + else: + self._register_default_tools() + + def _register_default_tools(self) -> None: + """注册默认工具""" + default_tools = [ + CalculatorTool(), + WebSearchTool(), + RestApiTool(), + SrApiQueryTool(), + ] + for tool in default_tools: + self.register_tool(tool) + + def register_tool( + self, + tool: BaseTool, + description: str = "", + version: str = "1.0.0", + timeout: int = 30, + retry: int = 0, + tags: Optional[List[str]] = None, + ) -> None: + """ + 注册工具 + + Args: + tool: 工具实例 + description: 描述(默认使用 tool.description) + version: 版本 + timeout: 超时时间 + retry: 重试次数 + tags: 标签 + """ + name = tool.name + metadata = ToolMetadata( + name=name, + description=description or tool.description, + version=version, + timeout=timeout, + retry=retry, + tags=tags or [], + ) + + self._tools[name] = tool + self._tool_metadata[name] = metadata + self._execution_stats[name] = { + "total_calls": 0, + "success_calls": 0, + "failed_calls": 0, + "total_time_ms": 0, + } + + ToolRegistry._entries[name] = type( + "RegistryEntry", + (), + {"instance": tool, "metadata": {"tool_metadata": metadata}} + )() + + logger.info(f"Registered tool: {name} (v{version})") + + def unregister_tool(self, name: str) -> bool: + """ + 注销工具 + + Args: + name: 工具名称 + + Returns: + 是否成功注销 + """ + if name in self._tools: + del self._tools[name] + del self._tool_metadata[name] + del self._execution_stats[name] + ToolRegistry.unregister(name) + logger.info(f"Unregistered tool: {name}") + return True + return False + + def get_tool(self, name: str) -> Optional[BaseTool]: + """获取工具实例""" + return self._tools.get(name) + + def get_tool_metadata(self, name: str) -> Optional[ToolMetadata]: + """获取工具元数据""" + return self._tool_metadata.get(name) + + def list_tools(self) -> List[str]: """列出可用工具名称""" return list(self._tools.keys()) - + + def get_tool_info(self, name: str) -> Optional[Dict[str, Any]]: + """获取工具详细信息""" + if name not in self._tools: + return None + + tool = self._tools[name] + metadata = self._tool_metadata.get(name) + stats = self._execution_stats.get(name, {}) + + return { + "name": name, + "description": metadata.description if metadata else tool.description, + "version": metadata.version if metadata else "unknown", + "timeout": metadata.timeout if metadata else 30, + "tags": metadata.tags if metadata else [], + "stats": { + "total_calls": stats.get("total_calls", 0), + "success_rate": self._calculate_success_rate(name), + }, + } + def call(self, tool_name: str, payload: Any) -> Dict[str, Any]: - """调用工具并返回标准化结果""" + """ + 调用工具并返回标准化结果 + + Args: + tool_name: 工具名称 + payload: 输入参数 + + Returns: + 标准化结果 {ok, data, error} + """ tool = self._tools.get(tool_name) if not tool: return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"} - + + start_time = time.time() + try: if isinstance(payload, (dict, list)): input_value = json.dumps(payload, ensure_ascii=False) @@ -34,8 +175,98 @@ class ToolRouter: input_value = "" else: input_value = str(payload) - + result = tool.run(input_value) + + self._record_success(tool_name, time.time() - start_time) + return {"ok": True, "data": result, "error": None} + except Exception as e: + self._record_failure(tool_name, time.time() - start_time) return {"ok": False, "data": None, "error": str(e)} + + def call_with_metadata( + self, + tool_name: str, + payload: Any, + ) -> Dict[str, Any]: + """ + 调用工具并返回包含元数据的结果 + + Args: + tool_name: 工具名称 + payload: 输入参数 + + Returns: + 包含元数据的结果 + """ + result = self.call(tool_name, payload) + metadata = self.get_tool_metadata(tool_name) + + return { + **result, + "tool_name": tool_name, + "tool_version": metadata.version if metadata else "unknown", + "execution_time_ms": self._execution_stats.get(tool_name, {}).get("last_time_ms", 0), + } + + def _record_success(self, tool_name: str, elapsed: float) -> None: + """记录成功执行""" + if tool_name in self._execution_stats: + stats = self._execution_stats[tool_name] + stats["total_calls"] += 1 + stats["success_calls"] += 1 + stats["total_time_ms"] += elapsed * 1000 + stats["last_time_ms"] = elapsed * 1000 + + def _record_failure(self, tool_name: str, elapsed: float) -> None: + """记录失败执行""" + if tool_name in self._execution_stats: + stats = self._execution_stats[tool_name] + stats["total_calls"] += 1 + stats["failed_calls"] += 1 + stats["total_time_ms"] += elapsed * 1000 + stats["last_time_ms"] = elapsed * 1000 + + def _calculate_success_rate(self, tool_name: str) -> float: + """计算成功率""" + stats = self._execution_stats.get(tool_name) + if not stats or stats["total_calls"] == 0: + return 0.0 + return stats["success_calls"] / stats["total_calls"] + + def get_all_stats(self) -> Dict[str, Dict[str, Any]]: + """获取所有工具的执行统计""" + result = {} + for name in self._tools: + result[name] = { + **self._execution_stats.get(name, {}), + "success_rate": self._calculate_success_rate(name), + } + return result + + def register_function( + self, + name: str, + func: Callable, + description: str = "", + timeout: int = 30, + ) -> None: + """ + 将普通函数注册为工具 + + Args: + name: 工具名称 + func: 函数 + description: 描述 + timeout: 超时时间 + """ + from langchain_core.tools import Tool + + tool = Tool( + name=name, + description=description, + func=func, + ) + self.register_tool(tool, description=description, timeout=timeout) diff --git a/workflows/workflow_manager.py b/workflows/workflow_manager.py index 626f487..2c85f86 100644 --- a/workflows/workflow_manager.py +++ b/workflows/workflow_manager.py @@ -1,74 +1,211 @@ -from typing import Dict, Any, Optional, List +""" +工作流管理器模块 + +支持动态注册和管理工作流类型 +""" + +from typing import Any, Callable, Dict, List, Optional, Type, Union from enum import Enum +from datetime import datetime, timedelta +import logging + from agent.conversation import ConversationAgent from agent.tool import ToolAgent +from core.registry import WorkflowRegistry, WorkflowMetadata + +logger = logging.getLogger(__name__) class WorkflowType(Enum): """可用的工作流类型""" CONVERSATION = "conversation" TOOL_USING = "tool_using" - + class WorkflowManager: - """管理不同工作流类型及其执行""" + """ + 管理不同工作流类型及其执行 + + 支持特性: + - 动态注册工作流 + - 会话管理 + - 工作流元数据 + """ def __init__(self, default_model_section: Optional[str] = None): - self.workflows = { - WorkflowType.CONVERSATION: ConversationAgent(model_section=default_model_section), - WorkflowType.TOOL_USING: ToolAgent(model_section=default_model_section) - } + self._default_model_section = default_model_section + self._workflows: Dict[str, Any] = {} + self._workflow_metadata: Dict[str, WorkflowMetadata] = {} self.active_sessions: Dict[str, Any] = {} + + self._register_default_workflows() - def get_workflow(self, workflow_type: WorkflowType): - """获取工作流实例""" - return self.workflows.get(workflow_type) + def _register_default_workflows(self) -> None: + """注册默认工作流""" + self.register_workflow( + name=WorkflowType.CONVERSATION.value, + agent=ConversationAgent(model_section=self._default_model_section), + description="多轮对话工作流", + version="1.0.0", + ) + self.register_workflow( + name=WorkflowType.TOOL_USING.value, + agent=ToolAgent(model_section=self._default_model_section), + description="工具调用工作流", + version="1.0.0", + ) - def execute_workflow(self, workflow_type: WorkflowType, user_input: str, - session_id: Optional[str] = None, **kwargs) -> Dict[str, Any]: - """执行指定工作流""" - workflow = self.get_workflow(workflow_type) + def register_workflow( + self, + name: str, + agent: Any, + description: str = "", + version: str = "1.0.0", + default_model: Optional[str] = None, + supported_features: Optional[List[str]] = None, + ) -> None: + """ + 注册工作流 + + Args: + name: 工作流名称 + agent: Agent 实例 + description: 描述 + version: 版本 + default_model: 默认模型 + supported_features: 支持的特性列表 + """ + metadata = WorkflowMetadata( + name=name, + description=description, + version=version, + default_model=default_model, + supported_features=supported_features or [], + ) + + self._workflows[name] = agent + self._workflow_metadata[name] = metadata + + WorkflowRegistry._entries[name] = type( + "RegistryEntry", + (), + {"instance": agent, "metadata": {"workflow_metadata": metadata}} + )() + + logger.info(f"Registered workflow: {name} (v{version})") + + def unregister_workflow(self, name: str) -> bool: + """ + 注销工作流 + + Args: + name: 工作流名称 + + Returns: + 是否成功注销 + """ + if name in self._workflows: + del self._workflows[name] + del self._workflow_metadata[name] + WorkflowRegistry.unregister(name) + logger.info(f"Unregistered workflow: {name}") + return True + return False + + def get_workflow(self, workflow_type: Union[WorkflowType, str]) -> Optional[Any]: + """ + 获取工作流实例 + + Args: + workflow_type: 工作流类型(枚举或字符串) + + Returns: + Agent 实例 + """ + name = workflow_type.value if isinstance(workflow_type, WorkflowType) else workflow_type + return self._workflows.get(name) + + def get_workflow_metadata(self, name: str) -> Optional[WorkflowMetadata]: + """获取工作流元数据""" + return self._workflow_metadata.get(name) + + def execute_workflow( + self, + workflow_type: Union[WorkflowType, str], + user_input: str, + session_id: Optional[str] = None, + **kwargs + ) -> Dict[str, Any]: + """ + 执行指定工作流 + + Args: + workflow_type: 工作流类型 + user_input: 用户输入 + session_id: 会话ID + **kwargs: 额外参数 + + Returns: + 执行结果 + """ + name = workflow_type.value if isinstance(workflow_type, WorkflowType) else workflow_type + workflow = self.get_workflow(name) if not workflow: - return {"error": f"Workflow {workflow_type.value} not found"} + return {"error": f"Workflow {name} not found"} - # 未提供会话 ID 时生成 if not session_id: session_id = f"session_{len(self.active_sessions) + 1}" - # 执行工作流 result = workflow.run(user_input, **kwargs) - # 存储会话数据 self.active_sessions[session_id] = { - "workflow_type": workflow_type, + "workflow_type": name, "last_result": result, "timestamp": self._get_timestamp() } return { "session_id": session_id, - "workflow_type": workflow_type.value, + "workflow_type": name, "result": result } def get_available_workflows(self) -> List[str]: """获取可用工作流列表""" - return [workflow.value for workflow in WorkflowType] + return list(self._workflows.keys()) + + def get_workflow_info(self, name: str) -> Optional[Dict[str, Any]]: + """获取工作流详细信息""" + if name not in self._workflows: + return None + + metadata = self._workflow_metadata.get(name) + return { + "name": name, + "description": metadata.description if metadata else "", + "version": metadata.version if metadata else "unknown", + "supported_features": metadata.supported_features if metadata else [], + } def _get_timestamp(self) -> str: """获取当前时间戳""" - from datetime import datetime return datetime.now().isoformat() def get_session_info(self, session_id: str) -> Optional[Dict[str, Any]]: """获取会话信息""" return self.active_sessions.get(session_id) - def cleanup_sessions(self, older_than_hours: int = 24): - """清理过期会话""" - from datetime import datetime, timedelta + def cleanup_sessions(self, older_than_hours: int = 24) -> int: + """ + 清理过期会话 + Args: + older_than_hours: 超过多少小时的会话将被清理 + + Returns: + 清理的会话数量 + """ cutoff_time = datetime.now() - timedelta(hours=older_than_hours) sessions_to_remove = [] @@ -80,4 +217,31 @@ class WorkflowManager: for session_id in sessions_to_remove: del self.active_sessions[session_id] - return len(sessions_to_remove) \ No newline at end of file + if sessions_to_remove: + logger.info(f"Cleaned up {len(sessions_to_remove)} expired sessions") + + return len(sessions_to_remove) + + def create_agent_instance( + self, + workflow_name: str, + agent_class: Type, + model_section: Optional[str] = None, + **kwargs + ) -> Any: + """ + 创建并注册新的 Agent 实例 + + Args: + workflow_name: 工作流名称 + agent_class: Agent 类 + model_section: 模型配置段 + **kwargs: Agent 构造参数 + + Returns: + Agent 实例 + """ + model = model_section or self._default_model_section + agent = agent_class(model_section=model, **kwargs) + self.register_workflow(name=workflow_name, agent=agent) + return agent