refactor(workflow): enforce typed query boundary

This commit is contained in:
jxxghp
2026-08-28 01:16:08 +08:00
parent 2f41780893
commit b4f8736541
35 changed files with 733 additions and 223 deletions
+7 -4
View File
@@ -7,7 +7,7 @@ from pydantic import BaseModel, Field
from app.agent.tools.base import MoviePilotTool
from app.agent.tools.tags import ToolTag
from app.application.agentdata import get_agent_workflow_port
from app.application.workflow import get_configured_workflow_query
from app.runtime.log import logger
@@ -54,8 +54,7 @@ class QueryWorkflowsTool(MoviePilotTool):
logger.info(f"执行工具: {self.name}, 参数: state={state}, name={name}, trigger_type={trigger_type}")
try:
workflow_oper = get_agent_workflow_port()
workflows = await workflow_oper.async_list()
workflows = await get_configured_workflow_query().list()
# 过滤工作流
filtered_workflows = []
@@ -101,7 +100,11 @@ class QueryWorkflowsTool(MoviePilotTool):
"event": "事件触发",
"manual": "手动触发"
}
trigger_type_desc = trigger_type_map.get(wf.trigger_type, wf.trigger_type or "定时触发")
trigger_type_key = wf.trigger_type or "timer"
trigger_type_desc = trigger_type_map.get(
trigger_type_key,
trigger_type_key,
)
simplified = {
"id": wf.id,
+2 -3
View File
@@ -61,8 +61,7 @@ def get_workflow_definition_command(
def get_workflow_query_service(
db: AsyncSession = Depends(get_async_session),
runtime: HostRuntime = Depends(get_host_runtime),
) -> WorkflowQueryService:
"""组装工作流只读查询用例,避免端点直接持有数据库操作器"""
return WorkflowQueryService(repository=runtime.workflow.repository(db))
"""返回组合根装配的工作流只读查询用例。"""
return runtime.workflow.query
-13
View File
@@ -5,7 +5,6 @@ from __future__ import annotations
from collections.abc import Callable
from typing import Any
AgentDataFactory = Callable[[], Any]
@@ -76,12 +75,6 @@ class DownloadHistoryPort(_PortProxy):
port_name = "download_history"
class WorkflowPort(_PortProxy):
"""工作流数据端口代理。"""
port_name = "workflow"
class PluginDataPort(_PortProxy):
"""插件数据端口代理。"""
@@ -110,7 +103,6 @@ def configure_agent_data_ports(**factories: AgentDataFactory) -> None:
"subscribe_history",
"transfer_history",
"download_history",
"workflow",
"plugin_data",
}
missing = sorted(required - factories.keys())
@@ -167,11 +159,6 @@ def get_agent_download_history_port() -> Any:
return get_agent_data_ports().download_history()
def get_agent_workflow_port() -> Any:
"""创建 Agent 工作流数据端口实例。"""
return get_agent_data_ports().workflow()
def get_agent_plugin_data_port() -> Any:
"""创建 Agent 插件数据端口实例。"""
return get_agent_data_ports().plugin_data()
+15 -5
View File
@@ -4,8 +4,10 @@ from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import asdict
from typing import Any, Optional
from app.application.workflow import WorkflowSnapshot
from app.schemas.media import resolve_media_identity
@@ -26,8 +28,10 @@ class ServerSharingService:
*,
subscribe_provider: Callable[[int], Any],
async_subscribe_provider: Callable[[int], Awaitable[Any]],
workflow_provider: Callable[[int], Any],
async_workflow_provider: Callable[[int], Awaitable[Any]],
workflow_provider: Callable[[int], Optional[WorkflowSnapshot]],
async_workflow_provider: Callable[
[int], Awaitable[Optional[WorkflowSnapshot]]
],
user_uuid_provider: Callable[[], str],
subscribe_sender: Callable[[dict], Any],
async_subscribe_sender: Callable[[dict], Awaitable[Any]],
@@ -68,9 +72,9 @@ class ServerSharingService:
return payload
@staticmethod
def prepare_workflow(workflow: Any) -> dict:
def prepare_workflow(workflow: WorkflowSnapshot) -> dict:
"""移除本地字段并把动作和流程编码为中心服务兼容格式。"""
workflow_dict = workflow.to_dict()
workflow_dict = asdict(workflow)
workflow_dict.pop("id", None)
workflow_dict.pop("context", None)
workflow_dict["actions"] = json.dumps(workflow_dict["actions"] or [])
@@ -78,7 +82,9 @@ class ServerSharingService:
return workflow_dict
@staticmethod
def validate_workflow(workflow: Any) -> tuple[bool, str]:
def validate_workflow(
workflow: Optional[WorkflowSnapshot],
) -> tuple[bool, str]:
"""验证工作流存在且同时包含动作与流程。"""
if not workflow:
return False, "工作流不存在"
@@ -160,6 +166,8 @@ class ServerSharingService:
valid, message = self.validate_workflow(workflow)
if not valid:
return False, message
if workflow is None:
return False, "工作流不存在"
payload = {
"share_title": share_title,
"share_comment": share_comment,
@@ -188,6 +196,8 @@ class ServerSharingService:
valid, message = self.validate_workflow(workflow)
if not valid:
return False, message
if workflow is None:
return False, "工作流不存在"
payload = {
"share_title": share_title,
"share_comment": share_comment,
+72 -15
View File
@@ -1,11 +1,12 @@
"""工作流状态与定义写操作应用用例。"""
from dataclasses import dataclass
import json
from collections.abc import Awaitable
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Callable, Mapping, Optional, Protocol, TypeVar
from typing import Any, Callable, List, Mapping, Optional, Protocol, TypeVar
from app.schemas.common import JsonData
WORKFLOW_TRIGGER_TIMER = "timer"
WORKFLOW_TRIGGER_EVENT = "event"
@@ -17,6 +18,30 @@ SUPPORTED_WORKFLOW_TRIGGERS = {
}
@dataclass(frozen=True, slots=True)
class WorkflowSnapshot:
"""工作流查询返回的脱离数据库会话的冻结快照。"""
id: int
name: str
description: Optional[str]
timer: Optional[str]
trigger_type: Optional[str]
event_type: Optional[str]
event_conditions: Mapping[str, JsonData]
state: str
current_action: Optional[str]
result: Optional[str]
run_count: Optional[int]
actions: tuple[Mapping[str, JsonData], ...]
flows: tuple[Mapping[str, JsonData], ...]
context: Mapping[str, JsonData]
execution_config: Mapping[str, JsonData]
execution_state: Mapping[str, JsonData]
add_time: Optional[str]
last_time: Optional[str]
class WorkflowRuntime(Protocol):
"""声明宿主入口与 Chain 消费的工作流运行时能力。"""
@@ -40,7 +65,7 @@ class WorkflowRuntime(Protocol):
"""移除全部或指定工作流的事件触发器。"""
...
def update_workflow_event(self, workflow: Any) -> None:
def update_workflow_event(self, workflow: WorkflowSnapshot) -> None:
"""按最新定义刷新工作流事件触发器。"""
...
@@ -87,15 +112,31 @@ def get_workflow_manager() -> WorkflowRuntime:
return _workflow_runtime_provider()
class AsyncWorkflowQueryRepository(Protocol):
"""工作流查询用例需要的异步读取端口。"""
class WorkflowQueryRepository(Protocol):
"""工作流查询用例需要的同步与异步快照端口。"""
async def async_list(self) -> list[Any]:
"""读取全部工作流。"""
def get(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""按 ID 读取工作流快照"""
...
async def async_get(self, workflow_id: int) -> Optional[Any]:
"""按 ID 读取工作流。"""
def list_enabled(self) -> List[WorkflowSnapshot]:
"""读取全部启用的工作流快照"""
...
def list_timer_enabled(self) -> List[WorkflowSnapshot]:
"""读取启用的定时工作流快照。"""
...
def list_event_enabled(self) -> List[WorkflowSnapshot]:
"""读取启用的事件工作流快照。"""
...
async def async_list(self) -> List[WorkflowSnapshot]:
"""异步读取全部工作流快照。"""
...
async def async_get(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""异步按 ID 读取工作流快照。"""
...
@@ -114,18 +155,34 @@ class WorkflowCachePort(Protocol):
class WorkflowQueryService:
"""提供工作流列表和详情查询,隔离 API 与数据库会话。"""
def __init__(self, repository: AsyncWorkflowQueryRepository) -> None:
"""保存请求级异步查询端口。"""
def __init__(self, repository: WorkflowQueryRepository) -> None:
"""保存可返回脱离会话快照的查询端口。"""
self._repository = repository
async def list(self) -> list[Any]:
"""返回全部工作流。"""
async def list(self) -> List[WorkflowSnapshot]:
"""返回全部工作流快照"""
return await self._repository.async_list()
async def get(self, workflow_id: int) -> Optional[Any]:
"""返回指定工作流。"""
async def get(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""返回指定工作流快照"""
return await self._repository.async_get(workflow_id)
def get_sync(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""同步返回指定工作流快照。"""
return self._repository.get(workflow_id)
def list_enabled(self) -> List[WorkflowSnapshot]:
"""同步返回全部启用的工作流快照。"""
return self._repository.list_enabled()
def list_timer_enabled(self) -> List[WorkflowSnapshot]:
"""同步返回启用的定时工作流快照。"""
return self._repository.list_timer_enabled()
def list_event_enabled(self) -> List[WorkflowSnapshot]:
"""同步返回启用的事件工作流快照。"""
return self._repository.list_event_enabled()
_configured_workflow_query: WorkflowQueryService | None = None
+48 -34
View File
@@ -14,7 +14,11 @@ from typing import Any, Callable, List, Optional, Tuple
from pydantic import BaseModel
from app.application.chain.data import get_chain_workflow_port
from app.application.workflow import get_workflow_manager
from app.application.workflow import (
WorkflowSnapshot,
get_configured_workflow_query,
get_workflow_manager,
)
from app.chain import ChainBase
from app.runtime.events import Event, eventmanager
from app.runtime.execution import OwnedThreadPoolExecutor
@@ -27,9 +31,6 @@ ARTIFACT_FIELDS = {"torrents", "medias", "fileitems", "downloads", "sites", "sub
DEFAULT_WORKFLOW_MAX_WORKERS = 4
WORKFLOW_EXECUTOR_STOP_TIMEOUT_SECONDS = 10.0
CIRCULAR_REFERENCE_PLACEHOLDER = "[Circular]"
Workflow = Any
def _serialize_workflow_key(key: Any) -> Any:
"""将映射键转换为 JSON 安全值。"""
if key is None or isinstance(key, (str, int, float, bool)):
@@ -118,7 +119,11 @@ class WorkflowExecutor:
工作流执行器
"""
def __init__(self, workflow: Workflow, step_callback: Callable = None):
def __init__(
self,
workflow: WorkflowSnapshot,
step_callback: Callable = None,
):
"""
初始化工作流执行器
:param workflow: 工作流对象
@@ -132,8 +137,13 @@ class WorkflowExecutor:
if step_callback
else False
)
self.actions = {action['id']: Action(**action) for action in workflow.actions}
self.flows = [ActionFlow(**flow) for flow in workflow.flows]
self.actions: dict[str, Action] = {}
for action_data in workflow.actions:
action = Action(**dict(action_data))
if not action.id:
raise ValueError("工作流动作缺少 ID")
self.actions[action.id] = action
self.flows = [ActionFlow(**dict(flow)) for flow in workflow.flows]
execution_config = getattr(workflow, "execution_config", None) or {}
execution_state = getattr(workflow, "execution_state", None) or {}
self.execution_config = (
@@ -652,7 +662,8 @@ class WorkflowExecutor:
self.flow_satisfied.add(flow_key)
if not source_success and self.node_states.get(source_id) == "failed":
self.flow_failed.add(flow_key)
self.evaluate_target_state(flow.target)
if flow.target:
self.evaluate_target_state(flow.target)
def evaluate_target_state(self, target_id: str) -> None:
"""
@@ -1277,10 +1288,29 @@ class WorkflowChain(ChainBase):
"""
workflowoper = get_chain_workflow_port()
def save_step(action: Action, context: ActionContext, execution_state: dict, completed: bool):
"""
保存上下文到数据库
"""
# 重置工作流
if from_begin:
workflowoper.reset(workflow_id)
# 查询工作流数据
workflow = get_configured_workflow_query().get_sync(workflow_id)
if not workflow:
logger.warn(f"工作流 {workflow_id} 不存在")
return False, "工作流不存在"
if not workflow.actions:
logger.warn(f"工作流 {workflow.name} 无动作")
return False, "工作流无动作"
if not workflow.flows:
logger.warn(f"工作流 {workflow.name} 无流程")
return False, "工作流无流程"
def save_step(
action: Action,
context: ActionContext,
execution_state: dict,
completed: bool,
) -> None:
"""保存动作上下文和结构化执行状态。"""
get_chain_workflow_port().step(
workflow_id,
action_id=action.id if completed else "",
@@ -1305,22 +1335,6 @@ class WorkflowChain(ChainBase):
},
)
# 重置工作流
if from_begin:
workflowoper.reset(workflow_id)
# 查询工作流数据
workflow = workflowoper.get(workflow_id)
if not workflow:
logger.warn(f"工作流 {workflow_id} 不存在")
return False, "工作流不存在"
if not workflow.actions:
logger.warn(f"工作流 {workflow.name} 无动作")
return False, "工作流无动作"
if not workflow.flows:
logger.warn(f"工作流 {workflow.name} 无流程")
return False, "工作流无流程"
logger.info(f"开始执行工作流 {workflow.name},共 {len(workflow.actions)} 个动作 ...")
if progress_callback:
progress_callback(
@@ -1355,22 +1369,22 @@ class WorkflowChain(ChainBase):
return True, ""
@staticmethod
def get_workflows() -> List[Workflow]:
def get_workflows() -> List[WorkflowSnapshot]:
"""
获取工作流列表
"""
return get_chain_workflow_port().list_enabled()
return get_configured_workflow_query().list_enabled()
@staticmethod
def get_timer_workflows() -> List[Workflow]:
def get_timer_workflows() -> List[WorkflowSnapshot]:
"""
获取定时触发的工作流列表
"""
return get_chain_workflow_port().get_timer_triggered_workflows()
return get_configured_workflow_query().list_timer_enabled()
@staticmethod
def get_event_workflows() -> List[Workflow]:
def get_event_workflows() -> List[WorkflowSnapshot]:
"""
获取事件触发的工作流列表
"""
return get_chain_workflow_port().get_event_triggered_workflows()
return get_configured_workflow_query().list_event_enabled()
+131 -5
View File
@@ -1,18 +1,144 @@
"""工作流执行状态事务适配器。"""
"""工作流查询与执行状态事务适配器。"""
from collections.abc import Callable
from typing import Any, TypeVar
from collections.abc import Callable, Iterable
from contextlib import AbstractAsyncContextManager
from copy import deepcopy
from typing import Any, Optional, TypeVar, cast
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import Session
from app.application.workflow import WorkflowExecutionCommand
from app.application.workflow import WorkflowExecutionCommand, WorkflowSnapshot
from app.db.oper.workflow import WorkflowOper
from app.db.uow import SqlAlchemyUnitOfWork
from app.schemas.common import JsonData
_Result = TypeVar("_Result")
def _copy_json_mapping(value: object, field_name: str) -> dict[str, JsonData]:
"""复制 ORM JSON 对象,拒绝把损坏结构带出 Session。"""
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"工作流 {field_name} 必须是 JSON 对象")
return cast(dict[str, JsonData], deepcopy(value))
def _copy_json_sequence(
value: object,
field_name: str,
) -> tuple[dict[str, JsonData], ...]:
"""复制 ORM JSON 对象序列,确保快照不共享可变容器。"""
if value is None:
return ()
if not isinstance(value, list) or any(not isinstance(item, dict) for item in value):
raise ValueError(f"工作流 {field_name} 必须是 JSON 对象数组")
return tuple(cast(dict[str, JsonData], deepcopy(item)) for item in value)
def _project_workflow(record: object) -> WorkflowSnapshot:
"""在持有数据库会话时把 ORM 记录投影为稳定快照。"""
workflow_id = getattr(record, "id", None)
name = getattr(record, "name", None)
state = getattr(record, "state", None)
if not isinstance(workflow_id, int) or not isinstance(name, str) or not isinstance(state, str):
raise ValueError("工作流记录缺少稳定身份或状态")
return WorkflowSnapshot(
id=workflow_id,
name=name,
description=getattr(record, "description", None),
timer=getattr(record, "timer", None),
trigger_type=getattr(record, "trigger_type", None),
event_type=getattr(record, "event_type", None),
event_conditions=_copy_json_mapping(
getattr(record, "event_conditions", None),
"event_conditions",
),
state=state,
current_action=getattr(record, "current_action", None),
result=getattr(record, "result", None),
run_count=getattr(record, "run_count", None),
actions=_copy_json_sequence(getattr(record, "actions", None), "actions"),
flows=_copy_json_sequence(getattr(record, "flows", None), "flows"),
context=_copy_json_mapping(getattr(record, "context", None), "context"),
execution_config=_copy_json_mapping(
getattr(record, "execution_config", None),
"execution_config",
),
execution_state=_copy_json_mapping(
getattr(record, "execution_state", None),
"execution_state",
),
add_time=getattr(record, "add_time", None),
last_time=getattr(record, "last_time", None),
)
class TransactionalWorkflowQueryRepository:
"""在自有短 Session 内查询并投影工作流快照。"""
def __init__(
self,
sync_session: Callable[[], Session],
async_session: Callable[[], AbstractAsyncContextManager[AsyncSession]],
) -> None:
"""保存同步 Session 工厂与异步 Session 作用域。"""
self._sync_session = sync_session
self._async_session = async_session
def get(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""在同步短 Session 内读取并投影单条工作流。"""
session = self._sync_session()
try:
record = WorkflowOper(session).get(workflow_id)
return _project_workflow(record) if record else None
finally:
session.close()
def list_enabled(self) -> list[WorkflowSnapshot]:
"""在同步短 Session 内投影全部启用工作流。"""
return self._list_sync(lambda repository: repository.list_enabled())
def list_timer_enabled(self) -> list[WorkflowSnapshot]:
"""在同步短 Session 内投影启用的定时工作流。"""
return self._list_sync(
lambda repository: repository.get_timer_triggered_workflows()
)
def list_event_enabled(self) -> list[WorkflowSnapshot]:
"""在同步短 Session 内投影启用的事件工作流。"""
return self._list_sync(
lambda repository: repository.get_event_triggered_workflows()
)
async def async_list(self) -> list[WorkflowSnapshot]:
"""在异步短 Session 内投影全部工作流。"""
async with self._async_session() as session:
records = await WorkflowOper(session).async_list()
return [_project_workflow(record) for record in records]
async def async_get(self, workflow_id: int) -> Optional[WorkflowSnapshot]:
"""在异步短 Session 内读取并投影单条工作流。"""
async with self._async_session() as session:
record = await WorkflowOper(session).async_get(workflow_id)
return _project_workflow(record) if record else None
def _list_sync(
self,
operation: Callable[[WorkflowOper], Iterable[object]],
) -> list[WorkflowSnapshot]:
"""在同步短 Session 内执行列表查询并完成投影。"""
session = self._sync_session()
try:
return [
_project_workflow(record)
for record in operation(WorkflowOper(session))
]
finally:
session.close()
class TransactionalWorkflowExecutionService:
"""为每次工作流执行状态写入创建独立短会话和 UnitOfWork。"""
+7
View File
@@ -774,6 +774,13 @@ _MESSAGE_NOTIFICATION_SYMBOL_ALIASES: Dict[str, SymbolAlias] = {
}
SYMBOL_ALIASES: Dict[str, Dict[str, SymbolAlias]] = {
"app.workflow": {
"WorkFlowManager": SymbolAlias(
target_module="app.workflow",
target_name="WorkflowManager",
replacement="app.workflow.WorkflowManager",
),
},
"app.application.transfer": {
name: SymbolAlias(
target_module="app.sdk._legacy.transfer",
+3 -3
View File
@@ -33,6 +33,7 @@ from app.application.outbox import dispatch_pending_outbox
from app.application.plugin.routes import register_plugin_api
from app.application.plugin.runtime import get_plugin_manager
from app.application.site.sites import SitesHelper # pylint: disable=import-error,no-name-in-module
from app.application.workflow import WorkflowSnapshot
from app.chain import ChainBase
from app.chain.mediaserver import MediaServerChain
from app.chain.recommend import RecommendChain
@@ -56,7 +57,6 @@ from app.schemas.dashboard import ScheduleProgress as _SchemaScheduleProgress
from app.schemas.message import Message, MessageType
from app.schemas.system import MediaServerConf as _SchemaMediaServerConf
from app.schemas.types import EventType, SystemConfigKey
from app.schemas.workflow import Workflow
lock = threading.Lock()
SCHEDULER_PROGRESS_PREFIX = "scheduler"
@@ -1703,7 +1703,7 @@ class Scheduler(ConfigReloadMixin, metaclass=SingletonClass):
for workflow in WorkflowChain().get_timer_workflows() or []:
self.update_workflow_job(workflow)
def remove_workflow_job(self, workflow: Workflow):
def remove_workflow_job(self, workflow: WorkflowSnapshot):
"""
移除工作流服务
"""
@@ -1787,7 +1787,7 @@ class Scheduler(ConfigReloadMixin, metaclass=SingletonClass):
role="system",
)
def update_workflow_job(self, workflow: Workflow):
def update_workflow_job(self, workflow: WorkflowSnapshot):
"""
更新工作流定时服务
"""
+2 -1
View File
@@ -17,7 +17,7 @@ from app.application.subscription.mutation import (
SubscriptionHistoryMutationRepository,
SubscriptionMutationRepository,
)
from app.application.workflow import WorkflowCachePort
from app.application.workflow import WorkflowCachePort, WorkflowQueryService
from app.runtime.tasks import TaskRegistry
@@ -165,6 +165,7 @@ class SiteRuntime:
class WorkflowRuntime:
"""工作流定义、状态与缓存操作所需的数据工厂。"""
query: WorkflowQueryService
repository: RepositoryFactory
system_config: Callable[[], WorkflowCachePort]
+16 -12
View File
@@ -114,7 +114,10 @@ from app.db.adapters.transfer.admission import TransactionalTransferAdmissionRep
from app.db.adapters.transfer.execution import (
TransactionalTransferExecutionRepository,
)
from app.db.adapters.workflow import TransactionalWorkflowExecutionService
from app.db.adapters.workflow import (
TransactionalWorkflowExecutionService,
TransactionalWorkflowQueryRepository,
)
from app.db.oper.agentchat import AgentChatOper
from app.db.oper.agenttask import AgentTaskOper
from app.db.oper.downloadhistory import DownloadHistoryOper
@@ -239,11 +242,6 @@ async def _async_get_subscribe(subscribe_id: int):
return await SubscribeOper().async_get(subscribe_id)
async def _async_get_workflow(workflow_id: int):
"""通过数据库操作器异步读取工作流,供服务端共享用例使用。"""
return await WorkflowOper().async_get(workflow_id)
def _execute_legacy_transfer_command(**kwargs: Any) -> Any:
"""把旧 Chain ABI 延迟转入唯一 TransferChain durable command。"""
from app.chain.transfer import TransferChain
@@ -273,7 +271,7 @@ def _build_chain_runtime_context() -> ChainRuntimeContext:
)
def configure_runtime_data_providers() -> None:
def configure_runtime_data_providers(workflow_query: WorkflowQueryService) -> None:
"""在启动组合层装配运行时和外部服务所需的数据库读取能力。"""
configure_service_config_reader(lambda key: get_configured_system_config().get(key))
configure_module_runtime(lambda: ModuleManager())
@@ -309,8 +307,8 @@ def configure_runtime_data_providers() -> None:
subscribe_id
),
async_subscribe_provider=_async_get_subscribe,
workflow_provider=lambda workflow_id: WorkflowOper().get(workflow_id),
async_workflow_provider=_async_get_workflow,
workflow_provider=workflow_query.get_sync,
async_workflow_provider=workflow_query.get,
user_uuid_provider=MoviePilotServerHelper.get_user_uuid,
subscribe_sender=MoviePilotServerHelper.subscribe_share,
async_subscribe_sender=MoviePilotServerHelper.async_subscribe_share,
@@ -799,6 +797,13 @@ async def init_modules() -> HostRuntime:
chain=lambda: build_chain_runtime_config(legacy_settings),
)
runtime_settings = _build_runtime_settings_service()
workflow_query = WorkflowQueryService(
repository=TransactionalWorkflowQueryRepository(
sync_session=SessionFactory,
async_session=async_session_scope,
)
)
configure_workflow_query(workflow_query)
agent_chat_persistence = AgentChatPersistenceService(
repository=lambda session: AgentChatOper(session),
async_executor=database_worker,
@@ -839,6 +844,7 @@ async def init_modules() -> HostRuntime:
outbox=SqlAlchemyAsyncOutboxStager,
),
workflow=WorkflowRuntime(
query=workflow_query,
repository=WorkflowOper,
system_config=get_configured_system_config,
),
@@ -853,7 +859,7 @@ async def init_modules() -> HostRuntime:
configure_token_runtime_config(lambda: build_token_runtime_config(legacy_settings))
# 旧 app.api.data 导入只保留 ABI 转发,正式 API 依赖全部读取 HostRuntime。
configure_api_data_runtime(api_data)
configure_runtime_data_providers()
configure_runtime_data_providers(workflow_query)
workflow_execution = TransactionalWorkflowExecutionService(SessionFactory)
configure_workflow_legacy_writer(workflow_execution)
configure_chain_data_ports(
@@ -908,7 +914,6 @@ async def init_modules() -> HostRuntime:
sync_session=SessionFactory,
async_session=async_session_scope,
)))
configure_workflow_query(WorkflowQueryService(repository=WorkflowOper()))
configure_agent_data_ports(
agent_chat=lambda: AgentChatOper(),
agent_task=lambda: AgentTaskOper(),
@@ -921,7 +926,6 @@ async def init_modules() -> HostRuntime:
subscribe_history=lambda: SubscribeHistoryOper(),
transfer_history=lambda: TransferHistoryOper(),
download_history=lambda: DownloadHistoryOper(),
workflow=lambda: WorkflowOper(),
plugin_data=lambda: PluginDataOper(),
)
configure_agent_task_execution(AgentTaskExecutionService(
+5 -5
View File
@@ -1,19 +1,19 @@
from app.application.workflow import configure_workflow_runtime
from app.workflow import WorkFlowManager
from app.workflow import WorkflowManager
# 启动模块是 concrete WorkFlowManager 的唯一宿主装配边界。
configure_workflow_runtime(lambda: WorkFlowManager())
# 启动模块是 concrete WorkflowManager 的唯一宿主装配边界。
configure_workflow_runtime(lambda: WorkflowManager())
def init_workflow():
"""
初始化工作流
"""
WorkFlowManager()
WorkflowManager()
def stop_workflow() -> bool:
"""
停止工作流并返回全部活动执行是否收敛
"""
return WorkFlowManager().stop()
return WorkflowManager().stop()
+11 -9
View File
@@ -4,21 +4,23 @@ from typing import Any, Dict, List, Optional, Tuple
from pydantic import BaseModel
from app.application.chain.data import get_chain_workflow_port
from app.application.workflow import WorkflowExecutionOwner
from app.application.workflow import (
WorkflowExecutionOwner,
WorkflowSnapshot,
get_configured_workflow_query,
)
from app.foundation.reflection import ModuleHelper
from app.foundation.singleton import Singleton
from app.runtime.config import global_vars
from app.runtime.events import Event, eventmanager
from app.runtime.log import logger
from app.runtime.stop import runtime_stop_state
from app.schemas.types import EventType
from app.schemas.workflow import Action, ActionContext, ActionResult, Workflow
from app.schemas.workflow import Action, ActionContext, ActionResult
_WORKFLOW_STOP_TIMEOUT_SECONDS = 10.0
class WorkFlowManager(metaclass=Singleton):
class WorkflowManager(metaclass=Singleton):
"""
工作流管理器
"""
@@ -316,7 +318,7 @@ class WorkFlowManager(metaclass=Singleton):
return {}
return action.get_contract()
def update_workflow_event(self, workflow: Workflow):
def update_workflow_event(self, workflow: WorkflowSnapshot):
"""
更新工作流事件触发器
"""
@@ -333,11 +335,11 @@ class WorkFlowManager(metaclass=Singleton):
"""
workflows = []
if workflow_id:
workflow = get_chain_workflow_port().get(workflow_id)
workflow = get_configured_workflow_query().get_sync(workflow_id)
if workflow:
workflows = [workflow]
else:
workflows = get_chain_workflow_port().get_event_triggered_workflows()
workflows = get_configured_workflow_query().list_event_enabled()
try:
for workflow in workflows:
self.update_workflow_event(workflow)
@@ -410,7 +412,7 @@ class WorkFlowManager(metaclass=Singleton):
"""
try:
# 检查工作流是否存在且启用
workflow = get_chain_workflow_port().get(workflow_id)
workflow = get_configured_workflow_query().get_sync(workflow_id)
if not workflow or workflow.state == 'P':
return