init
This commit is contained in:
@@ -0,0 +1,9 @@
|
||||
# Workflows 模块
|
||||
|
||||
## 作用
|
||||
|
||||
管理不同类型工作流的注册、路由和执行入口。
|
||||
|
||||
## 文件
|
||||
|
||||
- `workflow_manager.py`:工作流类型映射与执行调度
|
||||
@@ -0,0 +1,5 @@
|
||||
"""工作流管理模块"""
|
||||
|
||||
from .workflow_manager import WorkflowManager, WorkflowType
|
||||
|
||||
__all__ = ["WorkflowManager", "WorkflowType"]
|
||||
+77
-200
@@ -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)
|
||||
Reference in New Issue
Block a user