zcbot/web/run_lifecycle.py

142 lines
4.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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,
attachment_refs: Optional[list[dict[str, Any]]] = 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},
attachment_refs=attachment_refs,
))
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