mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-29 03:56:43 +08:00
297 lines
9.2 KiB
Python
297 lines
9.2 KiB
Python
"""Agent 自主定时任务执行的异步应用边界。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import threading
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass
|
|
from typing import Protocol, TypeVar
|
|
from uuid import uuid4
|
|
|
|
from app.application.database import AsyncDatabaseExecutor
|
|
from app.runtime.execution import await_task_to_terminal
|
|
from app.schemas.exception import DatabaseWorkerOverloadedError
|
|
|
|
|
|
T = TypeVar("T")
|
|
SyncTransaction = Callable[[Callable[[object], T]], T]
|
|
|
|
|
|
class AgentTaskRecord(Protocol):
|
|
"""执行认领与终态判定所需的任务投影。"""
|
|
|
|
id: int
|
|
enabled: bool
|
|
last_run_id: str | None
|
|
last_status: str
|
|
|
|
|
|
class AgentTaskRunRecord(Protocol):
|
|
"""执行期间需要脱离数据库会话持有的运行记录字段。"""
|
|
|
|
run_id: str
|
|
task_id: int
|
|
trigger_source: str
|
|
name: str
|
|
content: str
|
|
trigger_type: str
|
|
cron_expression: str | None
|
|
run_at: str | None
|
|
user_id: str
|
|
username: str | None
|
|
session_id: str
|
|
|
|
|
|
class AgentTaskRepository(Protocol):
|
|
"""AgentTask 执行用例使用的同步短事务仓储合同。"""
|
|
|
|
def get(self, task_id: int) -> AgentTaskRecord | None:
|
|
"""读取任务当前投影。"""
|
|
|
|
def begin_run(
|
|
self,
|
|
task_id: int,
|
|
trigger_source: str = "scheduled",
|
|
*,
|
|
run_id: str | None = None,
|
|
) -> AgentTaskRunRecord | None:
|
|
"""原子认领任务并创建运行快照。"""
|
|
|
|
def finish_run_outcome(
|
|
self,
|
|
run_id: str,
|
|
success: bool,
|
|
result: str,
|
|
) -> AgentTaskFinishRecord:
|
|
"""收口运行并返回事务确认的终态事实。"""
|
|
|
|
|
|
class AgentTaskFinishRecord(Protocol):
|
|
"""同步仓储返回的结构化运行终态。"""
|
|
|
|
run_finalized: bool
|
|
task_projection_updated: bool
|
|
date_task_disabled: bool
|
|
|
|
|
|
AgentTaskRepositoryFactory = Callable[[object], AgentTaskRepository]
|
|
AgentTaskScheduleRemover = Callable[[int, int, str], bool]
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class AgentTaskRunSnapshot:
|
|
"""任务认领成功后可安全跨越数据库会话的执行快照。"""
|
|
|
|
run_id: str
|
|
task_id: int
|
|
trigger_source: str
|
|
name: str
|
|
content: str
|
|
trigger_type: str
|
|
cron_expression: str | None
|
|
run_at: str | None
|
|
user_id: str
|
|
username: str | None
|
|
session_id: str
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class AgentTaskClaim:
|
|
"""任务认领结果;拒绝原因保持现有 Agent 用户提示合同。"""
|
|
|
|
run: AgentTaskRunSnapshot | None
|
|
rejection: str | None = None
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class AgentTaskFinishOutcome:
|
|
"""区分运行收口、任务投影更新与一次任务停用三个事实。"""
|
|
|
|
run_finalized: bool
|
|
task_projection_updated: bool
|
|
date_task_disabled: bool
|
|
|
|
|
|
class AgentTaskExecutionService:
|
|
"""通过有界数据库 worker 认领并收口一次 AgentTask 执行。"""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
repository: AgentTaskRepositoryFactory,
|
|
async_executor: AsyncDatabaseExecutor,
|
|
sync_transaction: SyncTransaction,
|
|
) -> None:
|
|
"""保存同步仓储、事务和异步执行器。"""
|
|
self._repository = repository
|
|
self._async_executor = async_executor
|
|
self._sync_transaction = sync_transaction
|
|
|
|
@staticmethod
|
|
def _snapshot(run: AgentTaskRunRecord) -> AgentTaskRunSnapshot:
|
|
"""在事务内复制运行字段,避免 ORM 对象越过会话边界。"""
|
|
return AgentTaskRunSnapshot(
|
|
run_id=run.run_id,
|
|
task_id=run.task_id,
|
|
trigger_source=run.trigger_source,
|
|
name=run.name,
|
|
content=run.content,
|
|
trigger_type=run.trigger_type,
|
|
cron_expression=run.cron_expression,
|
|
run_at=run.run_at,
|
|
user_id=run.user_id,
|
|
username=run.username,
|
|
session_id=run.session_id,
|
|
)
|
|
|
|
async def claim(
|
|
self,
|
|
task_id: int,
|
|
trigger_source: str = "scheduled",
|
|
*,
|
|
scheduler_generation: int | None = None,
|
|
remove_schedule: AgentTaskScheduleRemover | None = None,
|
|
) -> AgentTaskClaim:
|
|
"""认领一次执行;取消发生在提交后时先补偿收口再传播取消。"""
|
|
|
|
run_id = uuid4().hex
|
|
run_created = threading.Event()
|
|
|
|
def transaction(session: object) -> AgentTaskClaim:
|
|
repository = self._repository(session)
|
|
run = repository.begin_run(
|
|
task_id=task_id,
|
|
trigger_source=trigger_source,
|
|
run_id=run_id,
|
|
)
|
|
if not run:
|
|
task = repository.get(task_id)
|
|
return AgentTaskClaim(
|
|
run=None,
|
|
rejection=(
|
|
"Agent 定时任务不存在或已停用"
|
|
if not task or not task.enabled
|
|
else "Agent 定时任务当前不可执行"
|
|
),
|
|
)
|
|
run_created.set()
|
|
return AgentTaskClaim(run=self._snapshot(run))
|
|
|
|
async def claim_to_terminal() -> AgentTaskClaim:
|
|
"""容量瞬时耗尽时保留本轮调度,直到认领取得 admission。"""
|
|
while True:
|
|
try:
|
|
return await self._async_executor.run(
|
|
lambda: self._sync_transaction(transaction)
|
|
)
|
|
except DatabaseWorkerOverloadedError:
|
|
await asyncio.sleep(0.01)
|
|
|
|
claim_task = asyncio.create_task(claim_to_terminal())
|
|
try:
|
|
return await claim_task
|
|
except asyncio.CancelledError as cancellation:
|
|
# 纯容量拒绝尚未进入事务,不存在需要等待数据库容量的补偿对象。
|
|
if not run_created.is_set():
|
|
raise cancellation
|
|
finalize_task = asyncio.create_task(self._finalize(
|
|
run_id=run_id,
|
|
task_id=task_id,
|
|
success=False,
|
|
result="Agent 定时任务已取消",
|
|
scheduler_generation=scheduler_generation,
|
|
remove_schedule=remove_schedule,
|
|
))
|
|
await await_task_to_terminal(finalize_task)
|
|
raise cancellation
|
|
|
|
async def finalize(
|
|
self,
|
|
run: AgentTaskRunSnapshot,
|
|
*,
|
|
success: bool,
|
|
result: str,
|
|
scheduler_generation: int | None = None,
|
|
remove_schedule: AgentTaskScheduleRemover | None = None,
|
|
) -> AgentTaskFinishOutcome:
|
|
"""等待终态事务完成,并仅清理仍属于该 generation 的一次任务。"""
|
|
|
|
return await self._finalize(
|
|
run_id=run.run_id,
|
|
task_id=run.task_id,
|
|
success=success,
|
|
result=result,
|
|
scheduler_generation=scheduler_generation,
|
|
remove_schedule=remove_schedule,
|
|
)
|
|
|
|
async def _finalize(
|
|
self,
|
|
*,
|
|
run_id: str,
|
|
task_id: int,
|
|
success: bool,
|
|
result: str,
|
|
scheduler_generation: int | None,
|
|
remove_schedule: AgentTaskScheduleRemover | None,
|
|
) -> AgentTaskFinishOutcome:
|
|
"""按稳定运行 ID 收口,供正常路径和取消补偿共享。"""
|
|
|
|
def transaction(session: object) -> AgentTaskFinishOutcome:
|
|
repository = self._repository(session)
|
|
outcome = repository.finish_run_outcome(
|
|
run_id=run_id,
|
|
success=success,
|
|
result=result,
|
|
)
|
|
return AgentTaskFinishOutcome(
|
|
run_finalized=outcome.run_finalized,
|
|
task_projection_updated=outcome.task_projection_updated,
|
|
date_task_disabled=outcome.date_task_disabled,
|
|
)
|
|
|
|
async def finish_to_terminal() -> AgentTaskFinishOutcome:
|
|
"""容量瞬时耗尽时保留 owner,直到收口取得 admission。"""
|
|
while True:
|
|
try:
|
|
return await self._async_executor.run(
|
|
lambda: self._sync_transaction(transaction)
|
|
)
|
|
except DatabaseWorkerOverloadedError:
|
|
await asyncio.sleep(0.01)
|
|
|
|
finish_task = asyncio.create_task(finish_to_terminal())
|
|
cancellation: asyncio.CancelledError | None = None
|
|
try:
|
|
outcome = await asyncio.shield(finish_task)
|
|
except asyncio.CancelledError as error:
|
|
cancellation = error
|
|
outcome = await await_task_to_terminal(finish_task)
|
|
|
|
if (
|
|
outcome.date_task_disabled
|
|
and scheduler_generation is not None
|
|
and remove_schedule is not None
|
|
):
|
|
remove_schedule(task_id, scheduler_generation, run_id)
|
|
if cancellation is not None:
|
|
raise cancellation
|
|
return outcome
|
|
|
|
|
|
_service: AgentTaskExecutionService | None = None
|
|
|
|
|
|
def configure_agent_task_execution(service: AgentTaskExecutionService) -> None:
|
|
"""由启动组合根登记 AgentTask 执行服务。"""
|
|
global _service
|
|
_service = service
|
|
|
|
|
|
def get_agent_task_execution_service() -> AgentTaskExecutionService:
|
|
"""返回已登记的 AgentTask 执行服务。"""
|
|
if _service is None:
|
|
raise RuntimeError("AgentTask 执行服务尚未配置")
|
|
return _service
|