142 lines
4.4 KiB
Python
142 lines
4.4 KiB
Python
"""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 func, select, update
|
||
|
||
from core.storage import session_scope
|
||
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 = s.execute(
|
||
select(func.coalesce(func.max(Message.idx), -1) + 1)
|
||
.where(Message.task_id == task_id)
|
||
).scalar_one()
|
||
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
|