mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 23:47:41 +08:00
refactor(workflow): enforce typed query boundary
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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。"""
|
||||
|
||||
|
||||
@@ -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
@@ -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):
|
||||
"""
|
||||
更新工作流定时服务
|
||||
"""
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user