"""Web 运行生命周期的统一入口:事务抢占、消息落库与后台 worker 调度。 只收口所有 Web 形态入口共有的正确性边界;模型降级、自动标题、渠道回复和定时 结果统计仍由各自调用方负责。不引入持久化 run 实体或队列。 """ from __future__ import annotations import asyncio from dataclasses import dataclass, field from typing import Any, Callable, Optional from uuid import UUID from sqlalchemy import select, update from core.storage import session_scope from core.storage.message_index import allocate_message_idx from core.storage.models import Message, Task from .broker import broker from .common import INSTANCE from .runs import run_agent_bg class RunTaskNotFound(Exception): """目标 task 不存在或不属于指定用户。""" class RunTaskBusy(Exception): """目标 task 已有活跃 run。""" def __init__(self, status: str) -> None: self.status = status super().__init__(f"task already has an active run (status={status})") class RunScheduleError(Exception): """消息已持久化,但后台 coroutine 未能登记到 event loop。""" PrepareClaim = Callable[[Any, Task], tuple[dict[str, Any], dict[str, Any]]] @dataclass(frozen=True) class RunClaim: """抢占成功后返回给调用方的领域元数据。""" metadata: dict[str, Any] = field(default_factory=dict) def claim_run_with_message( task_id: UUID, user_id: UUID, user_message: str, *, prepare: Optional[PrepareClaim] = None, ) -> RunClaim: """在 task 行锁下原子提交 user 消息和 `run_status=running`。 `prepare` 在持锁事务内执行,可读取 task/用户状态并返回: `(额外 task 更新字段, 调用方元数据)`。它用于 Web 模型降级和自动标题快照, 不把这些领域规则塞进通用生命周期层。回调抛错会使整个事务回滚。 """ with session_scope() as s: task = s.execute( select(Task) .where(Task.task_id == task_id, Task.user_id == user_id) .with_for_update() ).scalar_one_or_none() if task is None: raise RunTaskNotFound(str(task_id)) if task.run_status in ("running", "cancelling"): raise RunTaskBusy(task.run_status) extra_values: dict[str, Any] = {} metadata: dict[str, Any] = {} if prepare is not None: extra_values, metadata = prepare(s, task) # task 行锁串行化同一 task 的 idx 分配,无需额外序列表或 advisory lock。 next_idx = allocate_message_idx(s, task_id, locked_task=task) s.add(Message( task_id=task_id, idx=int(next_idx), payload={"role": "user", "content": user_message}, )) values = { "run_status": "running", "run_error": None, "run_owner": INSTANCE or None, **extra_values, } s.execute(update(Task).where(Task.task_id == task_id).values(**values)) return RunClaim(metadata=dict(metadata)) def schedule_claimed_run( app, task_id: UUID, user_id: UUID, user_message: str, *, image_variant: str = "", video_variant: str = "", scheduled: bool = False, ) -> asyncio.Task: """调度已完成事务抢占的 run,并统一登记 broker / inflight。 调度失败时消息不能回滚(事务已经提交),因此把 task 收敛为 error 并抛出 `RunScheduleError`,调用方可转成 HTTP 500、渠道错误回复或定时失败记录。 """ broker.start(task_id) run_coro = asyncio.to_thread( run_agent_bg, task_id, user_id, user_message, image_variant, video_variant, scheduled, user_message_persisted=True, ) try: run_task = asyncio.create_task(run_coro) except Exception as e: run_coro.close() err = f"background scheduling failed: {type(e).__name__}: {e}" with session_scope() as s: s.execute( update(Task).where(Task.task_id == task_id).values( run_status="error", run_error=err, ) ) broker.close(task_id) raise RunScheduleError(err) from e app.state.inflight[run_task] = task_id run_task.add_done_callback(lambda t: app.state.inflight.pop(t, None)) return run_task