This commit is contained in:
2026-03-24 18:07:22 +08:00
parent e062368ef2
commit 9a16f738d8
121 changed files with 8904 additions and 3940 deletions
+77 -200
View File
@@ -1,247 +1,124 @@
"""
工作流管理器模块
支持动态注册和管理工作流类型
"""
from typing import Any, Callable, Dict, List, Optional, Type, Union
from typing import Dict, Any, Optional, List, cast
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__)
from config import CONVERSATION_ENABLE_MULTI_TURN
from agent.agents.conversation import ConversationAgent
from agent.agents.tool import ToolAgent
from services.common.app_errors import AppError, ErrorCode
from services.common.datetime_utils import DateTimeGenerator
class WorkflowType(Enum):
"""可用的工作流类型"""
CONVERSATION = "conversation"
TOOL_USING = "tool_using"
class WorkflowManager:
"""
管理不同工作流类型及其执行
"""管理不同工作流类型及其执行"""
支持特性:
- 动态注册工作流
- 会话管理
- 工作流元数据
"""
def __init__(self, default_model_section: Optional[str] = None):
self._default_model_section = default_model_section
self._workflows: Dict[str, Any] = {}
self._workflow_metadata: Dict[str, WorkflowMetadata] = {}
def __init__(self, default_model_section: Optional[str] = None, enable_multi_turn: Optional[bool] = None):
self.workflows = {
WorkflowType.CONVERSATION: ConversationAgent(model_section=default_model_section),
WorkflowType.TOOL_USING: ToolAgent(model_section=default_model_section)
}
self.active_sessions: Dict[str, Any] = {}
self._register_default_workflows()
self.enable_multi_turn = CONVERSATION_ENABLE_MULTI_TURN if enable_multi_turn is None else bool(enable_multi_turn)
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)
def get_workflow(self, workflow_type: WorkflowType):
"""获取工作流实例"""
return self.workflows.get(workflow_type)
@staticmethod
def _validate_user_input(user_input: Any) -> str:
if not isinstance(user_input, str) or not user_input.strip():
raise AppError(
code=ErrorCode.INVALID_REQUEST,
message="user_input 不能为空",
status_code=400,
detail={"field": "user_input", "reason": "missing_or_blank"},
)
return user_input
def execute_workflow(self, workflow_type: WorkflowType, user_input: str,
session_id: Optional[str] = None, **kwargs) -> Dict[str, Any]:
"""执行指定工作流"""
user_input = self._validate_user_input(user_input)
workflow = self.get_workflow(workflow_type)
if not workflow:
return {"error": f"Workflow {name} not found"}
return {"error": f"Workflow {workflow_type.value} not found"}
# 未提供会话 ID 时生成
if not session_id:
session_id = f"session_{len(self.active_sessions) + 1}"
existing_session: Optional[Dict[str, Any]] = self.active_sessions.get(session_id)
if existing_session and existing_session.get("workflow_type") != workflow_type:
raise ValueError(
f"Session '{session_id}' is already bound to workflow '{existing_session['workflow_type'].value}'"
)
run_kwargs: Dict[str, Any] = dict(kwargs)
if "conversation_id" not in run_kwargs:
run_kwargs["conversation_id"] = session_id
if workflow_type == WorkflowType.CONVERSATION and self.enable_multi_turn:
if "conversation_history" not in run_kwargs:
history_value = cast(Any, list((existing_session or {}).get("conversation_history") or []))
run_kwargs["conversation_history"] = history_value
if "last_context" not in run_kwargs:
context_value = cast(Any, dict((existing_session or {}).get("last_context") or {}))
run_kwargs["last_context"] = context_value
result = workflow.run(user_input, **kwargs)
self.active_sessions[session_id] = {
"workflow_type": name,
# 执行工作流
result = workflow.run(user_input, **run_kwargs)
session_record: Dict[str, Any] = {
"workflow_type": workflow_type,
"last_result": result,
"timestamp": self._get_timestamp()
}
if workflow_type == WorkflowType.CONVERSATION and self.enable_multi_turn:
context = (result or {}).get("context") or {}
session_record["conversation_history"] = list(
(result or {}).get("conversation_history") or context.get("conversation_history") or []
)
session_record["last_context"] = dict(context.get("last_context") or {})
# 存储会话数据
self.active_sessions[session_id] = session_record
return {
"session_id": session_id,
"workflow_type": name,
"workflow_type": workflow_type.value,
"result": result
}
def get_available_workflows(self) -> List[str]:
"""获取可用工作流列表"""
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 [],
}
return [workflow.value for workflow in WorkflowType]
def _get_timestamp(self) -> str:
"""获取当前时间戳"""
return datetime.now().isoformat()
return DateTimeGenerator.now().iso_str
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) -> int:
"""
清理过期会话
def cleanup_sessions(self, older_than_hours: int = 24):
"""清理过期会话"""
from datetime import timedelta
Args:
older_than_hours: 超过多少小时的会话将被清理
Returns:
清理的会话数量
"""
cutoff_time = datetime.now() - timedelta(hours=older_than_hours)
cutoff_time = DateTimeGenerator.now().dt - timedelta(hours=older_than_hours)
sessions_to_remove = []
for session_id, session_data in self.active_sessions.items():
session_time = datetime.fromisoformat(session_data["timestamp"])
session_time = DateTimeGenerator.parse(session_data["timestamp"], default_to_now=False)
if session_time < cutoff_time:
sessions_to_remove.append(session_id)
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
return len(sessions_to_remove)