from typing import Dict, Any, Optional, List, cast from enum import Enum 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, 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.enable_multi_turn = CONVERSATION_ENABLE_MULTI_TURN if enable_multi_turn is None else bool(enable_multi_turn) 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 {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, **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": workflow_type.value, "result": result } def get_available_workflows(self) -> List[str]: """获取可用工作流列表""" return [workflow.value for workflow in WorkflowType] def _get_timestamp(self) -> str: """获取当前时间戳""" 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): """清理过期会话""" from datetime import timedelta 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 = 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] return len(sessions_to_remove)