x
This commit is contained in:
@@ -13,6 +13,7 @@ from schemas.tool_input import ToolInput
|
|||||||
from schemas.tool_output import ToolOutput
|
from schemas.tool_output import ToolOutput
|
||||||
from schemas.chat_message_response import ChatMessageResponseDTO
|
from schemas.chat_message_response import ChatMessageResponseDTO
|
||||||
from schemas.chat_message_request import ChatMessageRequestDTO
|
from schemas.chat_message_request import ChatMessageRequestDTO
|
||||||
|
from schemas.super_agent import SuperAgentRequest, SuperAgentResponse, SuperAgentStreamEvent
|
||||||
from workflows.workflow_manager import WorkflowType
|
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 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
|
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")
|
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)
|
@router.post("/api/tools/execute", response_model=ToolOutput)
|
||||||
def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)):
|
def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)):
|
||||||
result = tool_router.call(payload.tool_name, payload.payload)
|
result = tool_router.call(payload.tool_name, payload.payload)
|
||||||
@@ -251,3 +300,135 @@ def update_sql_gen():
|
|||||||
return {"ok": True, "result": result}
|
return {"ok": True, "result": result}
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise HTTPException(status_code=500, detail=str(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")
|
||||||
|
|||||||
@@ -35,28 +35,9 @@ cache_ttl = 600
|
|||||||
table_retrieval_dataset_id = 9945baf512ea11f18ccb6a681b3130b2
|
table_retrieval_dataset_id = 9945baf512ea11f18ccb6a681b3130b2
|
||||||
sql_gen_dataset_id = ee68f53a12ec11f18e436a681b3130b2
|
sql_gen_dataset_id = ee68f53a12ec11f18e436a681b3130b2
|
||||||
|
|
||||||
[redis]
|
|
||||||
enabled = true
|
|
||||||
host = led-redis.lenovo.com
|
|
||||||
port = 30398
|
|
||||||
password = bgs123456
|
|
||||||
database = 0
|
|
||||||
sql_prompt_ttl = 600
|
|
||||||
|
|
||||||
[stream]
|
[stream]
|
||||||
progress_interval = 0.3
|
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]
|
[app]
|
||||||
service_name = local-model-streaming-api
|
service_name = local-model-streaming-api
|
||||||
host = 0.0.0.0
|
host = 0.0.0.0
|
||||||
|
|||||||
@@ -38,29 +38,10 @@ retrieval_top_k = 3
|
|||||||
table_retrieval_dataset_id =
|
table_retrieval_dataset_id =
|
||||||
sql_gen_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]
|
[stream]
|
||||||
# /api/workflows/stream 进度事件间隔(秒)
|
# /api/workflows/stream 进度事件间隔(秒)
|
||||||
progress_interval = 0.3
|
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]
|
||||||
# 是否启用 Nacos 注册
|
# 是否启用 Nacos 注册
|
||||||
enabled = false
|
enabled = false
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
"""
|
||||||
|
核心模块 - 提供扩展性基础设施
|
||||||
|
|
||||||
|
包含:
|
||||||
|
- Registry: 注册机制
|
||||||
|
- State: 增强的状态管理
|
||||||
|
- Provider: LLM Provider 抽象
|
||||||
|
- Response: 统一响应格式
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .registry import (
|
||||||
|
BaseRegistry,
|
||||||
|
NodeRegistry,
|
||||||
|
ToolRegistry,
|
||||||
|
WorkflowRegistry,
|
||||||
|
ProviderRegistry,
|
||||||
|
RegistryEntry,
|
||||||
|
ToolMetadata,
|
||||||
|
WorkflowMetadata,
|
||||||
|
ProviderMetadata,
|
||||||
|
register_tool,
|
||||||
|
register_workflow,
|
||||||
|
register_provider,
|
||||||
|
)
|
||||||
|
from .state import AgentState, StateContext
|
||||||
|
from .providers import LLMProvider, LLMFactory
|
||||||
|
from .response import ApiResponse, StreamEvent
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Registry
|
||||||
|
"BaseRegistry",
|
||||||
|
"NodeRegistry",
|
||||||
|
"ToolRegistry",
|
||||||
|
"WorkflowRegistry",
|
||||||
|
"ProviderRegistry",
|
||||||
|
"RegistryEntry",
|
||||||
|
"ToolMetadata",
|
||||||
|
"WorkflowMetadata",
|
||||||
|
"ProviderMetadata",
|
||||||
|
"register_tool",
|
||||||
|
"register_workflow",
|
||||||
|
"register_provider",
|
||||||
|
# State
|
||||||
|
"AgentState",
|
||||||
|
"StateContext",
|
||||||
|
# Provider
|
||||||
|
"LLMProvider",
|
||||||
|
"LLMFactory",
|
||||||
|
# Response
|
||||||
|
"ApiResponse",
|
||||||
|
"StreamEvent",
|
||||||
|
]
|
||||||
@@ -0,0 +1,218 @@
|
|||||||
|
"""
|
||||||
|
LLM Provider 抽象层
|
||||||
|
|
||||||
|
支持多种 LLM 提供商的统一接口
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, List, Optional, Protocol, Type, runtime_checkable
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from langchain_core.language_models import BaseChatModel
|
||||||
|
|
||||||
|
from config import Config
|
||||||
|
from .registry import ProviderRegistry, register_provider
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class LLMProvider(Protocol):
|
||||||
|
"""
|
||||||
|
LLM Provider 协议
|
||||||
|
|
||||||
|
定义所有 LLM 提供商必须实现的接口
|
||||||
|
"""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
"""Provider 名称"""
|
||||||
|
...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supported_models(self) -> List[str]:
|
||||||
|
"""支持的模型列表"""
|
||||||
|
...
|
||||||
|
|
||||||
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||||
|
"""
|
||||||
|
创建 LLM 模型实例
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: 模型配置
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
BaseChatModel 实例
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
class BaseLLMProvider(ABC):
|
||||||
|
"""LLM Provider 基类"""
|
||||||
|
|
||||||
|
def __init__(self, name: str, supported_models: List[str]):
|
||||||
|
self._name = name
|
||||||
|
self._supported_models = supported_models
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
return self._name
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supported_models(self) -> List[str]:
|
||||||
|
return self._supported_models
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||||
|
"""子类实现具体的模型创建逻辑"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class OpenAICompatibleProvider(BaseLLMProvider):
|
||||||
|
"""
|
||||||
|
OpenAI 兼容的 Provider
|
||||||
|
|
||||||
|
支持所有兼容 OpenAI API 的服务:
|
||||||
|
- OpenAI
|
||||||
|
- Azure OpenAI
|
||||||
|
- 本地部署的兼容服务
|
||||||
|
- 国产大模型(通义千问、文心一言等)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str = "openai_compatible", supported_models: Optional[List[str]] = None):
|
||||||
|
super().__init__(name, supported_models or [])
|
||||||
|
|
||||||
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||||
|
from langchain_openai import ChatOpenAI
|
||||||
|
|
||||||
|
return ChatOpenAI(
|
||||||
|
model=config.get("model_name") or config.get("MODEL_NAME"),
|
||||||
|
openai_api_key=config.get("api_key") or config.get("OPENAI_API_KEY"),
|
||||||
|
openai_api_base=config.get("base_url") or config.get("URL"),
|
||||||
|
temperature=config.get("temperature", 0.7),
|
||||||
|
max_tokens=config.get("max_tokens"),
|
||||||
|
timeout=config.get("timeout", 30),
|
||||||
|
max_retries=config.get("max_retries", 3),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AzureOpenAIProvider(BaseLLMProvider):
|
||||||
|
"""Azure OpenAI Provider"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("azure", ["gpt-4", "gpt-35-turbo"])
|
||||||
|
|
||||||
|
def create_model(self, config: Dict[str, Any]) -> BaseChatModel:
|
||||||
|
from langchain_openai import AzureChatOpenAI
|
||||||
|
|
||||||
|
return AzureChatOpenAI(
|
||||||
|
azure_deployment=config.get("deployment_name"),
|
||||||
|
openai_api_version=config.get("api_version", "2024-02-15-preview"),
|
||||||
|
azure_endpoint=config.get("endpoint"),
|
||||||
|
openai_api_key=config.get("api_key"),
|
||||||
|
temperature=config.get("temperature", 0.7),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class LLMFactory:
|
||||||
|
"""
|
||||||
|
LLM 工厂类
|
||||||
|
|
||||||
|
统一管理 LLM 实例的创建,支持:
|
||||||
|
- 多种 Provider
|
||||||
|
- 配置驱动
|
||||||
|
- 单例缓存
|
||||||
|
"""
|
||||||
|
|
||||||
|
_instances: Dict[str, BaseChatModel] = {}
|
||||||
|
_default_provider: str = "openai_compatible"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register_provider(cls, name: str, provider: LLMProvider) -> None:
|
||||||
|
"""注册 Provider"""
|
||||||
|
ProviderRegistry._entries[name] = type(
|
||||||
|
"RegistryEntry",
|
||||||
|
(),
|
||||||
|
{"instance": provider, "metadata": {}}
|
||||||
|
)()
|
||||||
|
logger.info(f"Registered LLM provider: {name}")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_provider(cls, name: str) -> Optional[LLMProvider]:
|
||||||
|
"""获取 Provider"""
|
||||||
|
entry = ProviderRegistry.get_entry(name)
|
||||||
|
if entry:
|
||||||
|
return entry.instance
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
provider: Optional[str] = None,
|
||||||
|
config: Optional[Dict[str, Any]] = None,
|
||||||
|
model_section: Optional[str] = None,
|
||||||
|
use_cache: bool = True,
|
||||||
|
) -> BaseChatModel:
|
||||||
|
"""
|
||||||
|
创建 LLM 实例
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider: Provider 名称,默认使用 openai_compatible
|
||||||
|
config: 模型配置
|
||||||
|
model_section: 配置文件中的模型段名
|
||||||
|
use_cache: 是否使用缓存
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
BaseChatModel 实例
|
||||||
|
"""
|
||||||
|
provider_name = provider or cls._default_provider
|
||||||
|
|
||||||
|
if model_section:
|
||||||
|
config = Config.get_section(model_section)
|
||||||
|
cache_key = f"{provider_name}:{model_section}"
|
||||||
|
else:
|
||||||
|
config = config or {}
|
||||||
|
cache_key = f"{provider_name}:{hash(frozenset(config.items()))}"
|
||||||
|
|
||||||
|
if use_cache and cache_key in cls._instances:
|
||||||
|
return cls._instances[cache_key]
|
||||||
|
|
||||||
|
provider_instance = cls.get_provider(provider_name)
|
||||||
|
|
||||||
|
if provider_instance is None:
|
||||||
|
provider_instance = OpenAICompatibleProvider()
|
||||||
|
|
||||||
|
model = provider_instance.create_model(config)
|
||||||
|
|
||||||
|
if use_cache:
|
||||||
|
cls._instances[cache_key] = model
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def clear_cache(cls) -> None:
|
||||||
|
"""清空缓存"""
|
||||||
|
cls._instances.clear()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_providers(cls) -> List[str]:
|
||||||
|
"""列出所有注册的 Provider"""
|
||||||
|
return ProviderRegistry.list_names()
|
||||||
|
|
||||||
|
|
||||||
|
OpenAICompatibleProvider()
|
||||||
|
ProviderRegistry._entries["openai_compatible"] = type(
|
||||||
|
"RegistryEntry",
|
||||||
|
(),
|
||||||
|
{"instance": OpenAICompatibleProvider(), "metadata": {}}
|
||||||
|
)()
|
||||||
|
|
||||||
|
|
||||||
|
def create_chat_model(model_section: Optional[str] = None) -> BaseChatModel:
|
||||||
|
"""
|
||||||
|
创建聊天模型(兼容现有代码)
|
||||||
|
|
||||||
|
这是现有 llm_factory.create_chat_model 的替代实现,
|
||||||
|
使用新的 Provider 架构但保持接口兼容
|
||||||
|
"""
|
||||||
|
return LLMFactory.create(model_section=model_section)
|
||||||
@@ -0,0 +1,222 @@
|
|||||||
|
"""
|
||||||
|
核心注册机制模块
|
||||||
|
|
||||||
|
提供统一的注册器模式,支持动态扩展:
|
||||||
|
- NodeRegistry: 节点注册器
|
||||||
|
- ToolRegistry: 工具注册器
|
||||||
|
- WorkflowRegistry: 工作流注册器
|
||||||
|
- ProviderRegistry: LLM Provider 注册器
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Callable, Dict, List, Optional, Type, TypeVar, Protocol, runtime_checkable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
|
import time
|
||||||
|
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class Registerable(Protocol[T]):
|
||||||
|
"""可注册对象的协议"""
|
||||||
|
name: str
|
||||||
|
description: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class RegistryEntry:
|
||||||
|
"""注册条目"""
|
||||||
|
name: str
|
||||||
|
instance: Any
|
||||||
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
registered_at: float = field(default_factory=time.time)
|
||||||
|
|
||||||
|
|
||||||
|
class BaseRegistry:
|
||||||
|
"""基础注册器"""
|
||||||
|
|
||||||
|
_entries: Dict[str, RegistryEntry] = {}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register(cls, name: str, metadata: Optional[Dict[str, Any]] = None):
|
||||||
|
"""
|
||||||
|
注册装饰器
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@NodeRegistry.register("my_node", metadata={"category": "processing"})
|
||||||
|
def my_node(state: AgentState) -> AgentState:
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
def decorator(obj: T) -> T:
|
||||||
|
entry = RegistryEntry(
|
||||||
|
name=name,
|
||||||
|
instance=obj,
|
||||||
|
metadata=metadata or {},
|
||||||
|
)
|
||||||
|
cls._entries[name] = entry
|
||||||
|
return obj
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get(cls, name: str) -> Optional[Any]:
|
||||||
|
"""获取注册的对象"""
|
||||||
|
entry = cls._entries.get(name)
|
||||||
|
return entry.instance if entry else None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_entry(cls, name: str) -> Optional[RegistryEntry]:
|
||||||
|
"""获取注册条目(包含元数据)"""
|
||||||
|
return cls._entries.get(name)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_names(cls) -> List[str]:
|
||||||
|
"""列出所有注册名称"""
|
||||||
|
return list(cls._entries.keys())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_entries(cls) -> List[RegistryEntry]:
|
||||||
|
"""列出所有注册条目"""
|
||||||
|
return list(cls._entries.values())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def unregister(cls, name: str) -> bool:
|
||||||
|
"""注销注册"""
|
||||||
|
if name in cls._entries:
|
||||||
|
del cls._entries[name]
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def clear(cls):
|
||||||
|
"""清空注册表"""
|
||||||
|
cls._entries.clear()
|
||||||
|
|
||||||
|
|
||||||
|
class NodeRegistry(BaseRegistry):
|
||||||
|
"""节点注册器 - 用于注册 Agent 工作流节点"""
|
||||||
|
_entries: Dict[str, RegistryEntry] = {}
|
||||||
|
|
||||||
|
|
||||||
|
class ToolRegistry(BaseRegistry):
|
||||||
|
"""工具注册器 - 用于注册工具"""
|
||||||
|
_entries: Dict[str, RegistryEntry] = {}
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowRegistry(BaseRegistry):
|
||||||
|
"""工作流注册器 - 用于注册工作流类型"""
|
||||||
|
_entries: Dict[str, RegistryEntry] = {}
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderRegistry(BaseRegistry):
|
||||||
|
"""LLM Provider 注册器"""
|
||||||
|
_entries: Dict[str, RegistryEntry] = {}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ToolMetadata:
|
||||||
|
"""工具元数据"""
|
||||||
|
name: str
|
||||||
|
description: str
|
||||||
|
version: str = "1.0.0"
|
||||||
|
timeout: int = 30
|
||||||
|
retry: int = 0
|
||||||
|
parameters_schema: Optional[Dict[str, Any]] = None
|
||||||
|
requires_auth: bool = False
|
||||||
|
tags: List[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class WorkflowMetadata:
|
||||||
|
"""工作流元数据"""
|
||||||
|
name: str
|
||||||
|
description: str
|
||||||
|
version: str = "1.0.0"
|
||||||
|
agent_class: Optional[Type] = None
|
||||||
|
default_model: Optional[str] = None
|
||||||
|
supported_features: List[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ProviderMetadata:
|
||||||
|
"""LLM Provider 元数据"""
|
||||||
|
name: str
|
||||||
|
description: str
|
||||||
|
provider_type: str
|
||||||
|
supported_models: List[str] = field(default_factory=list)
|
||||||
|
config_schema: Optional[Dict[str, Any]] = None
|
||||||
|
|
||||||
|
|
||||||
|
def register_tool(
|
||||||
|
name: str,
|
||||||
|
description: str = "",
|
||||||
|
version: str = "1.0.0",
|
||||||
|
timeout: int = 30,
|
||||||
|
retry: int = 0,
|
||||||
|
tags: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
工具注册装饰器
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@register_tool("calculator", "数学计算", timeout=10)
|
||||||
|
class CalculatorTool(BaseTool):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
metadata = ToolMetadata(
|
||||||
|
name=name,
|
||||||
|
description=description,
|
||||||
|
version=version,
|
||||||
|
timeout=timeout,
|
||||||
|
retry=retry,
|
||||||
|
tags=tags or [],
|
||||||
|
)
|
||||||
|
return ToolRegistry.register(name, {"tool_metadata": metadata})
|
||||||
|
|
||||||
|
|
||||||
|
def register_workflow(
|
||||||
|
name: str,
|
||||||
|
description: str = "",
|
||||||
|
version: str = "1.0.0",
|
||||||
|
default_model: Optional[str] = None,
|
||||||
|
supported_features: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
工作流注册装饰器
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@register_workflow("data_query", "数据查询工作流")
|
||||||
|
class DataQueryAgent(BaseAgent):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
metadata = WorkflowMetadata(
|
||||||
|
name=name,
|
||||||
|
description=description,
|
||||||
|
version=version,
|
||||||
|
default_model=default_model,
|
||||||
|
supported_features=supported_features or [],
|
||||||
|
)
|
||||||
|
return WorkflowRegistry.register(name, {"workflow_metadata": metadata})
|
||||||
|
|
||||||
|
|
||||||
|
def register_provider(
|
||||||
|
name: str,
|
||||||
|
provider_type: str,
|
||||||
|
description: str = "",
|
||||||
|
supported_models: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
LLM Provider 注册装饰器
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@register_provider("openai", "openai", supported_models=["gpt-4", "gpt-3.5"])
|
||||||
|
class OpenAIProvider:
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
metadata = ProviderMetadata(
|
||||||
|
name=name,
|
||||||
|
description=description,
|
||||||
|
provider_type=provider_type,
|
||||||
|
supported_models=supported_models or [],
|
||||||
|
)
|
||||||
|
return ProviderRegistry.register(name, {"provider_metadata": metadata})
|
||||||
@@ -0,0 +1,248 @@
|
|||||||
|
"""
|
||||||
|
统一响应格式模块
|
||||||
|
|
||||||
|
提供标准化的 API 响应和流式事件格式
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, Generic, List, Optional, TypeVar, Literal
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class ApiResponse(BaseModel, Generic[T]):
|
||||||
|
"""
|
||||||
|
统一 API 响应格式
|
||||||
|
|
||||||
|
所有 API 响应都使用这个格式,提供一致的响应结构
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@router.get("/users/{user_id}")
|
||||||
|
async def get_user(user_id: str) -> ApiResponse[User]:
|
||||||
|
user = await user_service.get(user_id)
|
||||||
|
return ApiResponse.success(data=user)
|
||||||
|
"""
|
||||||
|
|
||||||
|
code: str = Field(default="success", description="响应代码")
|
||||||
|
message: str = Field(default="", description="响应消息")
|
||||||
|
data: Optional[T] = Field(default=None, description="响应数据")
|
||||||
|
trace_id: Optional[str] = Field(default=None, description="追踪ID")
|
||||||
|
timestamp: int = Field(
|
||||||
|
default_factory=lambda: int(time.time() * 1000),
|
||||||
|
description="时间戳(毫秒)"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def success(cls, data: T = None, message: str = "", trace_id: Optional[str] = None) -> "ApiResponse[T]":
|
||||||
|
"""创建成功响应"""
|
||||||
|
return cls(
|
||||||
|
code="success",
|
||||||
|
message=message,
|
||||||
|
data=data,
|
||||||
|
trace_id=trace_id or uuid.uuid4().hex,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def error(
|
||||||
|
cls,
|
||||||
|
code: str = "error",
|
||||||
|
message: str = "",
|
||||||
|
data: T = None,
|
||||||
|
trace_id: Optional[str] = None,
|
||||||
|
) -> "ApiResponse[T]":
|
||||||
|
"""创建错误响应"""
|
||||||
|
return cls(
|
||||||
|
code=code,
|
||||||
|
message=message,
|
||||||
|
data=data,
|
||||||
|
trace_id=trace_id or uuid.uuid4().hex,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_exception(cls, exc: Exception, trace_id: Optional[str] = None) -> "ApiResponse[None]":
|
||||||
|
"""从异常创建错误响应"""
|
||||||
|
return cls.error(
|
||||||
|
code="internal_error",
|
||||||
|
message=str(exc),
|
||||||
|
trace_id=trace_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
def is_success(self) -> bool:
|
||||||
|
"""判断是否成功"""
|
||||||
|
return self.code == "success"
|
||||||
|
|
||||||
|
|
||||||
|
class PagedResponse(BaseModel, Generic[T]):
|
||||||
|
"""
|
||||||
|
分页响应格式
|
||||||
|
|
||||||
|
用于返回分页数据
|
||||||
|
"""
|
||||||
|
|
||||||
|
items: List[T] = Field(default_factory=list, description="数据列表")
|
||||||
|
total: int = Field(default=0, description="总数")
|
||||||
|
page: int = Field(default=1, description="当前页")
|
||||||
|
page_size: int = Field(default=20, description="每页大小")
|
||||||
|
total_pages: int = Field(default=0, description="总页数")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
items: List[T],
|
||||||
|
total: int,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 20,
|
||||||
|
) -> "PagedResponse[T]":
|
||||||
|
"""创建分页响应"""
|
||||||
|
total_pages = (total + page_size - 1) // page_size if page_size > 0 else 0
|
||||||
|
return cls(
|
||||||
|
items=items,
|
||||||
|
total=total,
|
||||||
|
page=page,
|
||||||
|
page_size=page_size,
|
||||||
|
total_pages=total_pages,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class StreamEvent(BaseModel):
|
||||||
|
"""
|
||||||
|
流式响应事件
|
||||||
|
|
||||||
|
用于 SSE (Server-Sent Events) 流式响应
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
async def event_stream():
|
||||||
|
yield StreamEvent(event="start", data="Processing started")
|
||||||
|
# ... 处理逻辑
|
||||||
|
yield StreamEvent(event="result", data=json.dumps(result))
|
||||||
|
yield StreamEvent(event="done", data="")
|
||||||
|
"""
|
||||||
|
|
||||||
|
event: str = Field(..., description="事件类型")
|
||||||
|
data: str = Field(default="", description="事件数据")
|
||||||
|
event_id: Optional[str] = Field(default=None, description="事件ID")
|
||||||
|
retry: Optional[int] = Field(default=None, description="重试间隔(毫秒)")
|
||||||
|
|
||||||
|
def to_sse(self) -> str:
|
||||||
|
"""转换为 SSE 格式字符串"""
|
||||||
|
lines = [f"event: {self.event}"]
|
||||||
|
if self.event_id:
|
||||||
|
lines.append(f"id: {self.event_id}")
|
||||||
|
if self.retry:
|
||||||
|
lines.append(f"retry: {self.retry}")
|
||||||
|
lines.append(f"data: {self.data}")
|
||||||
|
lines.append("")
|
||||||
|
lines.append("")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def message(cls, data: str, event_id: Optional[str] = None) -> "StreamEvent":
|
||||||
|
"""创建消息事件"""
|
||||||
|
return cls(event="message", data=data, event_id=event_id)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def done(cls) -> "StreamEvent":
|
||||||
|
"""创建完成事件"""
|
||||||
|
return cls(event="done", data="[DONE]")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def error(cls, message: str) -> "StreamEvent":
|
||||||
|
"""创建错误事件"""
|
||||||
|
return cls(event="error", data=message)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowEvent(BaseModel):
|
||||||
|
"""
|
||||||
|
工作流事件
|
||||||
|
|
||||||
|
用于工作流执行过程中的状态通知
|
||||||
|
"""
|
||||||
|
|
||||||
|
workflow_id: str = Field(..., description="工作流ID")
|
||||||
|
event_type: Literal[
|
||||||
|
"started",
|
||||||
|
"node_started",
|
||||||
|
"node_completed",
|
||||||
|
"node_failed",
|
||||||
|
"completed",
|
||||||
|
"failed",
|
||||||
|
] = Field(..., description="事件类型")
|
||||||
|
node_name: Optional[str] = Field(None, description="节点名称")
|
||||||
|
data: Optional[Dict[str, Any]] = Field(None, description="事件数据")
|
||||||
|
error: Optional[str] = Field(None, description="错误信息")
|
||||||
|
timestamp: int = Field(
|
||||||
|
default_factory=lambda: int(time.time() * 1000),
|
||||||
|
description="时间戳"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def started(cls, workflow_id: str) -> "WorkflowEvent":
|
||||||
|
"""创建开始事件"""
|
||||||
|
return cls(workflow_id=workflow_id, event_type="started")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def node_started(cls, workflow_id: str, node_name: str) -> "WorkflowEvent":
|
||||||
|
"""创建节点开始事件"""
|
||||||
|
return cls(
|
||||||
|
workflow_id=workflow_id,
|
||||||
|
event_type="node_started",
|
||||||
|
node_name=node_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def node_completed(
|
||||||
|
cls,
|
||||||
|
workflow_id: str,
|
||||||
|
node_name: str,
|
||||||
|
data: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> "WorkflowEvent":
|
||||||
|
"""创建节点完成事件"""
|
||||||
|
return cls(
|
||||||
|
workflow_id=workflow_id,
|
||||||
|
event_type="node_completed",
|
||||||
|
node_name=node_name,
|
||||||
|
data=data,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def completed(
|
||||||
|
cls,
|
||||||
|
workflow_id: str,
|
||||||
|
data: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> "WorkflowEvent":
|
||||||
|
"""创建完成事件"""
|
||||||
|
return cls(workflow_id=workflow_id, event_type="completed", data=data)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def failed(
|
||||||
|
cls,
|
||||||
|
workflow_id: str,
|
||||||
|
error: str,
|
||||||
|
node_name: Optional[str] = None,
|
||||||
|
) -> "WorkflowEvent":
|
||||||
|
"""创建失败事件"""
|
||||||
|
return cls(
|
||||||
|
workflow_id=workflow_id,
|
||||||
|
event_type="failed",
|
||||||
|
node_name=node_name,
|
||||||
|
error=error,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ErrorCode:
|
||||||
|
"""错误代码常量"""
|
||||||
|
|
||||||
|
SUCCESS = "success"
|
||||||
|
UNKNOWN_ERROR = "unknown_error"
|
||||||
|
INVALID_REQUEST = "invalid_request"
|
||||||
|
INVALID_WORKFLOW_TYPE = "invalid_workflow_type"
|
||||||
|
SQL_GENERATION_FAILED = "sql_generation_failed"
|
||||||
|
TOOL_NOT_FOUND = "tool_not_found"
|
||||||
|
TOOL_EXECUTION_FAILED = "tool_execution_failed"
|
||||||
|
INTERNAL_ERROR = "internal_error"
|
||||||
|
TIMEOUT = "timeout"
|
||||||
|
RATE_LIMITED = "rate_limited"
|
||||||
|
UNAUTHORIZED = "unauthorized"
|
||||||
+199
@@ -0,0 +1,199 @@
|
|||||||
|
"""
|
||||||
|
增强的状态管理模块
|
||||||
|
|
||||||
|
使用 Pydantic 提供类型安全和验证
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, List, Optional, Literal
|
||||||
|
from pydantic import BaseModel, Field, field_validator
|
||||||
|
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage, SystemMessage
|
||||||
|
|
||||||
|
|
||||||
|
class StateContext(BaseModel):
|
||||||
|
"""状态上下文 - 存储工作流执行过程中的数据"""
|
||||||
|
|
||||||
|
original_input: Optional[str] = Field(None, description="用户原始输入")
|
||||||
|
normalized_input: Optional[str] = Field(None, description="规范化后的输入")
|
||||||
|
intent: Optional[str] = Field(None, description="识别的意图")
|
||||||
|
table_match: Optional[Dict[str, Any]] = Field(None, description="表名匹配结果")
|
||||||
|
final_sql: Optional[str] = Field(None, description="生成的 SQL")
|
||||||
|
sr_api_result: Optional[Any] = Field(None, description="API 执行结果")
|
||||||
|
|
||||||
|
class Config:
|
||||||
|
extra = "allow"
|
||||||
|
|
||||||
|
def get(self, key: str, default: Any = None) -> Any:
|
||||||
|
"""获取上下文值"""
|
||||||
|
return getattr(self, key, default)
|
||||||
|
|
||||||
|
def set(self, key: str, value: Any) -> None:
|
||||||
|
"""设置上下文值"""
|
||||||
|
setattr(self, key, value)
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
"""转换为字典"""
|
||||||
|
return self.model_dump(exclude_none=True)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentState(BaseModel):
|
||||||
|
"""
|
||||||
|
Agent 工作流状态定义
|
||||||
|
|
||||||
|
使用 Pydantic 提供类型安全和验证
|
||||||
|
"""
|
||||||
|
|
||||||
|
messages: List[BaseMessage] = Field(default_factory=list, description="消息历史")
|
||||||
|
current_step: str = Field(default="start", description="当前步骤")
|
||||||
|
context: StateContext = Field(default_factory=StateContext, description="上下文数据")
|
||||||
|
|
||||||
|
model_config = {
|
||||||
|
"arbitrary_types_allowed": True,
|
||||||
|
"extra": "forbid",
|
||||||
|
}
|
||||||
|
|
||||||
|
@field_validator("messages", mode="before")
|
||||||
|
@classmethod
|
||||||
|
def validate_messages(cls, v):
|
||||||
|
"""验证并转换消息列表"""
|
||||||
|
if not isinstance(v, list):
|
||||||
|
return []
|
||||||
|
|
||||||
|
result = []
|
||||||
|
for msg in v:
|
||||||
|
if isinstance(msg, BaseMessage):
|
||||||
|
result.append(msg)
|
||||||
|
elif isinstance(msg, dict):
|
||||||
|
msg_type = msg.get("type", "human")
|
||||||
|
content = msg.get("content", "")
|
||||||
|
if msg_type == "human":
|
||||||
|
result.append(HumanMessage(content=content))
|
||||||
|
elif msg_type == "ai":
|
||||||
|
result.append(AIMessage(content=content))
|
||||||
|
elif msg_type == "system":
|
||||||
|
result.append(SystemMessage(content=content))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def add_message(self, message: BaseMessage) -> "AgentState":
|
||||||
|
"""添加消息并返回新状态"""
|
||||||
|
return AgentState(
|
||||||
|
messages=[*self.messages, message],
|
||||||
|
current_step=self.current_step,
|
||||||
|
context=self.context,
|
||||||
|
)
|
||||||
|
|
||||||
|
def add_human_message(self, content: str) -> "AgentState":
|
||||||
|
"""添加用户消息"""
|
||||||
|
return self.add_message(HumanMessage(content=content))
|
||||||
|
|
||||||
|
def add_ai_message(self, content: str) -> "AgentState":
|
||||||
|
"""添加 AI 消息"""
|
||||||
|
return self.add_message(AIMessage(content=content))
|
||||||
|
|
||||||
|
def update_step(self, step: str) -> "AgentState":
|
||||||
|
"""更新当前步骤"""
|
||||||
|
return AgentState(
|
||||||
|
messages=self.messages,
|
||||||
|
current_step=step,
|
||||||
|
context=self.context,
|
||||||
|
)
|
||||||
|
|
||||||
|
def update_context(self, **kwargs) -> "AgentState":
|
||||||
|
"""更新上下文"""
|
||||||
|
new_context = self.context.model_copy()
|
||||||
|
for key, value in kwargs.items():
|
||||||
|
new_context.set(key, value)
|
||||||
|
return AgentState(
|
||||||
|
messages=self.messages,
|
||||||
|
current_step=self.current_step,
|
||||||
|
context=new_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_last_message(self) -> Optional[BaseMessage]:
|
||||||
|
"""获取最后一条消息"""
|
||||||
|
return self.messages[-1] if self.messages else None
|
||||||
|
|
||||||
|
def get_context(self, key: str, default: Any = None) -> Any:
|
||||||
|
"""获取上下文值"""
|
||||||
|
return self.context.get(key, default)
|
||||||
|
|
||||||
|
def to_legacy_format(self) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
转换为旧格式(兼容现有代码)
|
||||||
|
|
||||||
|
现有代码期望 state 是一个可修改的对象,
|
||||||
|
这个方法返回一个兼容的字典格式
|
||||||
|
"""
|
||||||
|
return {
|
||||||
|
"messages": self.messages,
|
||||||
|
"current_step": self.current_step,
|
||||||
|
"context": self.context.to_dict(),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_legacy_format(cls, data: Dict[str, Any]) -> "AgentState":
|
||||||
|
"""从旧格式创建"""
|
||||||
|
context_data = data.get("context", {})
|
||||||
|
if isinstance(context_data, StateContext):
|
||||||
|
context = context_data
|
||||||
|
else:
|
||||||
|
context = StateContext(**context_data) if context_data else StateContext()
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
messages=data.get("messages", []),
|
||||||
|
current_step=data.get("current_step", "start"),
|
||||||
|
context=context,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class MutableAgentState:
|
||||||
|
"""
|
||||||
|
可变的 Agent 状态包装器
|
||||||
|
|
||||||
|
用于兼容现有代码中直接修改 state 的模式
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, state: Optional[AgentState] = None):
|
||||||
|
self._state = state or AgentState()
|
||||||
|
self._context_overrides: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def messages(self) -> List[BaseMessage]:
|
||||||
|
return self._state.messages
|
||||||
|
|
||||||
|
@messages.setter
|
||||||
|
def messages(self, value: List[BaseMessage]):
|
||||||
|
self._state = AgentState(
|
||||||
|
messages=value,
|
||||||
|
current_step=self._state.current_step,
|
||||||
|
context=self._state.context,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def current_step(self) -> str:
|
||||||
|
return self._state.current_step
|
||||||
|
|
||||||
|
@current_step.setter
|
||||||
|
def current_step(self, value: str):
|
||||||
|
self._state = AgentState(
|
||||||
|
messages=self._state.messages,
|
||||||
|
current_step=value,
|
||||||
|
context=self._state.context,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def context(self) -> Dict[str, Any]:
|
||||||
|
"""返回可修改的上下文字典"""
|
||||||
|
result = self._state.context.to_dict()
|
||||||
|
result.update(self._context_overrides)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def to_immutable(self) -> AgentState:
|
||||||
|
"""转换为不可变状态"""
|
||||||
|
context = self._state.context.model_copy()
|
||||||
|
for key, value in self._context_overrides.items():
|
||||||
|
context.set(key, value)
|
||||||
|
return AgentState(
|
||||||
|
messages=self._state.messages,
|
||||||
|
current_step=self._state.current_step,
|
||||||
|
context=context,
|
||||||
|
)
|
||||||
@@ -8,5 +8,3 @@ uvicorn>=0.30.0
|
|||||||
nacos-sdk-python==2.0.9
|
nacos-sdk-python==2.0.9
|
||||||
httpx>=0.27.0
|
httpx>=0.27.0
|
||||||
pyyaml>=6.0.1
|
pyyaml>=6.0.1
|
||||||
redis>=5.0.0
|
|
||||||
pymysql>=1.1.1
|
|
||||||
|
|||||||
@@ -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="时间戳(毫秒)")
|
||||||
@@ -2,11 +2,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
try:
|
|
||||||
import redis
|
|
||||||
except Exception:
|
|
||||||
redis = None
|
|
||||||
|
|
||||||
|
|
||||||
class CacheBase:
|
class CacheBase:
|
||||||
"""缓存接口"""
|
"""缓存接口"""
|
||||||
@@ -26,20 +21,3 @@ class NoopCache(CacheBase):
|
|||||||
|
|
||||||
def set(self, key: str, value: str, ttl: int) -> None:
|
def set(self, key: str, value: str, ttl: int) -> None:
|
||||||
return 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)
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,9 +2,6 @@ import json
|
|||||||
import os
|
import os
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from config import Config
|
|
||||||
from services.cache import NoopCache, RedisCache
|
|
||||||
|
|
||||||
|
|
||||||
class SqlPromptManager:
|
class SqlPromptManager:
|
||||||
"""按表名读取 SQL 提示词"""
|
"""按表名读取 SQL 提示词"""
|
||||||
@@ -12,51 +9,11 @@ class SqlPromptManager:
|
|||||||
def __init__(self, base_dir: Optional[str] = None):
|
def __init__(self, base_dir: Optional[str] = None):
|
||||||
root_dir = os.path.dirname(os.path.dirname(__file__))
|
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._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
|
@staticmethod
|
||||||
def _safe_filename(name: str) -> str:
|
def _safe_filename(name: str) -> str:
|
||||||
return name.replace("..", "").replace("/", "_").replace("\\", "_")
|
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]]:
|
def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]:
|
||||||
"""读取指定表的提示词 JSON"""
|
"""读取指定表的提示词 JSON"""
|
||||||
if not table_name:
|
if not table_name:
|
||||||
@@ -67,19 +24,9 @@ class SqlPromptManager:
|
|||||||
if not os.path.exists(path):
|
if not os.path.exists(path):
|
||||||
return None
|
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:
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
prompt = json.load(f)
|
prompt = json.load(f)
|
||||||
|
|
||||||
self._cache.set(key, json.dumps(prompt, ensure_ascii=False), self._cache_ttl)
|
|
||||||
return prompt
|
return prompt
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -4,58 +4,10 @@ import json
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
import pymysql
|
|
||||||
|
|
||||||
from config import Config
|
|
||||||
|
|
||||||
|
|
||||||
class StructuredLogger:
|
class StructuredLogger:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
cfg = Config.get_section("logging_mysql")
|
pass
|
||||||
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
|
|
||||||
|
|
||||||
def log(self, level: str, event: str, trace_id: str, payload: Optional[Dict[str, Any]] = None, error_code: Optional[str] = None) -> None:
|
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({
|
print(json.dumps({
|
||||||
@@ -67,32 +19,6 @@ class StructuredLogger:
|
|||||||
"created_at": datetime.now().isoformat(),
|
"created_at": datetime.now().isoformat(),
|
||||||
}, ensure_ascii=False))
|
}, 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
|
_GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None
|
||||||
|
|
||||||
|
|||||||
+239
-8
@@ -1,5 +1,13 @@
|
|||||||
|
"""
|
||||||
|
工具路由器模块
|
||||||
|
|
||||||
|
支持动态注册和管理工具
|
||||||
|
"""
|
||||||
|
|
||||||
import json
|
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
|
from langchain_core.tools import BaseTool
|
||||||
|
|
||||||
@@ -7,26 +15,159 @@ from tools.calculator import CalculatorTool
|
|||||||
from tools.web_search import WebSearchTool
|
from tools.web_search import WebSearchTool
|
||||||
from tools.rest_api_tool import RestApiTool
|
from tools.rest_api_tool import RestApiTool
|
||||||
from tools.sr_api_tool import SrApiQueryTool
|
from tools.sr_api_tool import SrApiQueryTool
|
||||||
|
from core.registry import ToolRegistry, ToolMetadata
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class ToolRouter:
|
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())
|
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]:
|
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)
|
tool = self._tools.get(tool_name)
|
||||||
if not tool:
|
if not tool:
|
||||||
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
|
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
|
||||||
|
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if isinstance(payload, (dict, list)):
|
if isinstance(payload, (dict, list)):
|
||||||
input_value = json.dumps(payload, ensure_ascii=False)
|
input_value = json.dumps(payload, ensure_ascii=False)
|
||||||
@@ -36,6 +177,96 @@ class ToolRouter:
|
|||||||
input_value = str(payload)
|
input_value = str(payload)
|
||||||
|
|
||||||
result = tool.run(input_value)
|
result = tool.run(input_value)
|
||||||
|
|
||||||
|
self._record_success(tool_name, time.time() - start_time)
|
||||||
|
|
||||||
return {"ok": True, "data": result, "error": None}
|
return {"ok": True, "data": result, "error": None}
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
self._record_failure(tool_name, time.time() - start_time)
|
||||||
return {"ok": False, "data": None, "error": str(e)}
|
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
@@ -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 enum import Enum
|
||||||
|
from datetime import datetime, timedelta
|
||||||
|
import logging
|
||||||
|
|
||||||
from agent.conversation import ConversationAgent
|
from agent.conversation import ConversationAgent
|
||||||
from agent.tool import ToolAgent
|
from agent.tool import ToolAgent
|
||||||
|
from core.registry import WorkflowRegistry, WorkflowMetadata
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class WorkflowType(Enum):
|
class WorkflowType(Enum):
|
||||||
@@ -11,64 +23,189 @@ class WorkflowType(Enum):
|
|||||||
|
|
||||||
|
|
||||||
class WorkflowManager:
|
class WorkflowManager:
|
||||||
"""管理不同工作流类型及其执行"""
|
"""
|
||||||
|
管理不同工作流类型及其执行
|
||||||
|
|
||||||
|
支持特性:
|
||||||
|
- 动态注册工作流
|
||||||
|
- 会话管理
|
||||||
|
- 工作流元数据
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self, default_model_section: Optional[str] = None):
|
def __init__(self, default_model_section: Optional[str] = None):
|
||||||
self.workflows = {
|
self._default_model_section = default_model_section
|
||||||
WorkflowType.CONVERSATION: ConversationAgent(model_section=default_model_section),
|
self._workflows: Dict[str, Any] = {}
|
||||||
WorkflowType.TOOL_USING: ToolAgent(model_section=default_model_section)
|
self._workflow_metadata: Dict[str, WorkflowMetadata] = {}
|
||||||
}
|
|
||||||
self.active_sessions: Dict[str, Any] = {}
|
self.active_sessions: Dict[str, Any] = {}
|
||||||
|
|
||||||
def get_workflow(self, workflow_type: WorkflowType):
|
self._register_default_workflows()
|
||||||
"""获取工作流实例"""
|
|
||||||
return self.workflows.get(workflow_type)
|
|
||||||
|
|
||||||
def execute_workflow(self, workflow_type: WorkflowType, user_input: str,
|
def _register_default_workflows(self) -> None:
|
||||||
session_id: Optional[str] = None, **kwargs) -> Dict[str, Any]:
|
"""注册默认工作流"""
|
||||||
"""执行指定工作流"""
|
self.register_workflow(
|
||||||
workflow = self.get_workflow(workflow_type)
|
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:
|
if not workflow:
|
||||||
return {"error": f"Workflow {workflow_type.value} not found"}
|
return {"error": f"Workflow {name} not found"}
|
||||||
|
|
||||||
# 未提供会话 ID 时生成
|
|
||||||
if not session_id:
|
if not session_id:
|
||||||
session_id = f"session_{len(self.active_sessions) + 1}"
|
session_id = f"session_{len(self.active_sessions) + 1}"
|
||||||
|
|
||||||
# 执行工作流
|
|
||||||
result = workflow.run(user_input, **kwargs)
|
result = workflow.run(user_input, **kwargs)
|
||||||
|
|
||||||
# 存储会话数据
|
|
||||||
self.active_sessions[session_id] = {
|
self.active_sessions[session_id] = {
|
||||||
"workflow_type": workflow_type,
|
"workflow_type": name,
|
||||||
"last_result": result,
|
"last_result": result,
|
||||||
"timestamp": self._get_timestamp()
|
"timestamp": self._get_timestamp()
|
||||||
}
|
}
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"session_id": session_id,
|
"session_id": session_id,
|
||||||
"workflow_type": workflow_type.value,
|
"workflow_type": name,
|
||||||
"result": result
|
"result": result
|
||||||
}
|
}
|
||||||
|
|
||||||
def get_available_workflows(self) -> List[str]:
|
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:
|
def _get_timestamp(self) -> str:
|
||||||
"""获取当前时间戳"""
|
"""获取当前时间戳"""
|
||||||
from datetime import datetime
|
|
||||||
return datetime.now().isoformat()
|
return datetime.now().isoformat()
|
||||||
|
|
||||||
def get_session_info(self, session_id: str) -> Optional[Dict[str, Any]]:
|
def get_session_info(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||||
"""获取会话信息"""
|
"""获取会话信息"""
|
||||||
return self.active_sessions.get(session_id)
|
return self.active_sessions.get(session_id)
|
||||||
|
|
||||||
def cleanup_sessions(self, older_than_hours: int = 24):
|
def cleanup_sessions(self, older_than_hours: int = 24) -> int:
|
||||||
"""清理过期会话"""
|
"""
|
||||||
from datetime import datetime, timedelta
|
清理过期会话
|
||||||
|
|
||||||
|
Args:
|
||||||
|
older_than_hours: 超过多少小时的会话将被清理
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
清理的会话数量
|
||||||
|
"""
|
||||||
cutoff_time = datetime.now() - timedelta(hours=older_than_hours)
|
cutoff_time = datetime.now() - timedelta(hours=older_than_hours)
|
||||||
|
|
||||||
sessions_to_remove = []
|
sessions_to_remove = []
|
||||||
@@ -80,4 +217,31 @@ class WorkflowManager:
|
|||||||
for session_id in sessions_to_remove:
|
for session_id in sessions_to_remove:
|
||||||
del self.active_sessions[session_id]
|
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)
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user