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
+181
View File
@@ -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")
-19
View File
@@ -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
-19
View File
@@ -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
+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,
)
-2
View File
@@ -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
+37
View File
@@ -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="时间戳(毫秒)")
-22
View File
@@ -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)
-53
View File
@@ -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
+1 -75
View File
@@ -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
+239 -8
View File
@@ -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)
@@ -36,6 +177,96 @@ class ToolRouter:
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)
+188 -24
View File
@@ -1,7 +1,19 @@
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):
@@ -11,64 +23,189 @@ class WorkflowType(Enum):
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] = {}
def get_workflow(self, workflow_type: WorkflowType):
"""获取工作流实例"""
return self.workflows.get(workflow_type)
self._register_default_workflows()
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_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 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]
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