Files
more_dots/workflows/workflow_manager.py
T

124 lines
5.1 KiB
Python
Raw Normal View History

2026-03-24 18:07:22 +08:00
from typing import Dict, Any, Optional, List, cast
2026-02-17 02:31:39 +08:00
from enum import Enum
2026-03-24 18:07:22 +08:00
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
2026-02-17 02:31:39 +08:00
class WorkflowType(Enum):
2026-02-26 13:43:44 +08:00
"""可用的工作流类型"""
2026-02-17 02:31:39 +08:00
CONVERSATION = "conversation"
TOOL_USING = "tool_using"
2026-03-24 18:07:22 +08:00
2026-02-17 02:31:39 +08:00
class WorkflowManager:
2026-03-24 18:07:22 +08:00
"""管理不同工作流类型及其执行"""
2026-02-17 02:31:39 +08:00
2026-03-24 18:07:22 +08:00
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)
}
2026-02-17 02:31:39 +08:00
self.active_sessions: Dict[str, Any] = {}
2026-03-24 18:07:22 +08:00
self.enable_multi_turn = CONVERSATION_ENABLE_MULTI_TURN if enable_multi_turn is None else bool(enable_multi_turn)
2026-03-11 23:40:39 +08:00
2026-03-24 18:07:22 +08:00
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)
2026-02-17 02:31:39 +08:00
if not workflow:
2026-03-24 18:07:22 +08:00
return {"error": f"Workflow {workflow_type.value} not found"}
2026-02-17 02:31:39 +08:00
2026-03-24 18:07:22 +08:00
# 未提供会话 ID 时生成
2026-02-17 02:31:39 +08:00
if not session_id:
session_id = f"session_{len(self.active_sessions) + 1}"
2026-03-24 18:07:22 +08:00
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,
2026-02-17 02:31:39 +08:00
"last_result": result,
"timestamp": self._get_timestamp()
}
2026-03-24 18:07:22 +08:00
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
2026-02-17 02:31:39 +08:00
return {
"session_id": session_id,
2026-03-24 18:07:22 +08:00
"workflow_type": workflow_type.value,
2026-02-17 02:31:39 +08:00
"result": result
}
def get_available_workflows(self) -> List[str]:
2026-02-26 13:43:44 +08:00
"""获取可用工作流列表"""
2026-03-24 18:07:22 +08:00
return [workflow.value for workflow in WorkflowType]
2026-02-17 02:31:39 +08:00
def _get_timestamp(self) -> str:
2026-02-26 13:43:44 +08:00
"""获取当前时间戳"""
2026-03-24 18:07:22 +08:00
return DateTimeGenerator.now().iso_str
2026-02-17 02:31:39 +08:00
def get_session_info(self, session_id: str) -> Optional[Dict[str, Any]]:
2026-02-26 13:43:44 +08:00
"""获取会话信息"""
2026-02-17 02:31:39 +08:00
return self.active_sessions.get(session_id)
2026-03-24 18:07:22 +08:00
def cleanup_sessions(self, older_than_hours: int = 24):
"""清理过期会话"""
from datetime import timedelta
2026-02-17 02:31:39 +08:00
2026-03-24 18:07:22 +08:00
cutoff_time = DateTimeGenerator.now().dt - timedelta(hours=older_than_hours)
2026-02-17 02:31:39 +08:00
sessions_to_remove = []
for session_id, session_data in self.active_sessions.items():
2026-03-24 18:07:22 +08:00
session_time = DateTimeGenerator.parse(session_data["timestamp"], default_to_now=False)
2026-02-17 02:31:39 +08:00
if session_time < cutoff_time:
sessions_to_remove.append(session_id)
for session_id in sessions_to_remove:
del self.active_sessions[session_id]
2026-03-24 18:07:22 +08:00
return len(sessions_to_remove)