zcbot/web/software_followups.py

215 lines
7.2 KiB
Python
Raw Permalink 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.

"""Software Job 终态后的结果账本收尾与 Agent 自动续跑。"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from uuid import UUID
from sqlalchemy import select, update
from core.software_contracts import get_contract
from core.storage import session_scope
from core.storage.message_index import allocate_message_idx
from core.storage.models import Message, SoftwareJob, Task
from .common import INSTANCE
from .run_lifecycle import RunScheduleError, schedule_claimed_run
TERMINAL_STATUSES = {"succeeded", "failed", "cancelled"}
@dataclass(frozen=True)
class FollowupClaim:
job_id: UUID
task_id: UUID
user_id: UUID
action: str
prompt: str = ""
def _analysis_prompt(job: SoftwareJob) -> str:
contract = get_contract(job.capability)
if contract.workspace is not None and getattr(job, "workspace_id", None) is not None:
return (
"[专业软件任务完成事件]\n"
f"任务 {job.job_id}{contract.display_name})已成功完成,工程保存在 "
f"Workspace {job.workspace_id}。请调用 software_job_status 获取输入、"
"本地输出清单和预览信息,结合原始数据向用户报告结果;只有用户明确要求"
"取回文件时才请求导出。"
)
output_dir = f"{contract.output_namespace}/{job.job_id}"
return (
"[专业软件任务完成事件]\n"
f"任务 {job.job_id}{contract.display_name})已成功完成,"
f"输出目录为 {output_dir}。请调用 software_job_status 获取完整产物清单,"
"读取输出图表和输入数据,向用户报告结果并分析主要趋势、结论及必要的限制。"
)
def pending_followup_ids(limit: int = 20) -> list[UUID]:
with session_scope() as session:
return list(session.execute(
select(SoftwareJob.job_id)
.where(
SoftwareJob.status.in_(TERMINAL_STATUSES),
SoftwareJob.followup_status == "pending",
)
.order_by(SoftwareJob.terminal_at, SoftwareJob.job_id)
.limit(limit)
).scalars())
def requeue_stale_analysis_followups() -> int:
"""启动 reaper 已确认 run 随进程丢失时,把自动分析退回可重试状态。"""
with session_scope() as session:
job_ids = list(session.execute(
select(SoftwareJob.job_id)
.join(Task, Task.task_id == SoftwareJob.task_id)
.where(
SoftwareJob.followup_status == "running",
Task.run_status == "error",
Task.run_error == "server restarted before run finished",
)
).scalars())
if job_ids:
session.execute(
update(SoftwareJob)
.where(SoftwareJob.job_id.in_(job_ids))
.values(followup_status="pending")
)
return len(job_ids)
def claim_followup(job_id: UUID) -> FollowupClaim | None:
"""原子领取回调report 只完成账本analyze 等 task 空闲后进入对话。"""
with session_scope() as session:
candidate = session.execute(
select(SoftwareJob.task_id).where(SoftwareJob.job_id == job_id)
).scalar_one_or_none()
if candidate is None:
return None
task = session.execute(
select(Task).where(Task.task_id == candidate).with_for_update()
).scalar_one_or_none()
if task is None:
return None
job = session.execute(
select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update()
).scalar_one_or_none()
if (
job is None
or job.status not in TERMINAL_STATUSES
or job.followup_status != "pending"
):
return None
if job.status != "succeeded" or job.completion_action == "report":
job.followup_status = "completed"
return FollowupClaim(
job_id=job.job_id,
task_id=task.task_id,
user_id=job.user_id,
action="report",
)
if task.run_status in {"running", "cancelling"}:
return None
prompt = _analysis_prompt(job)
next_idx = allocate_message_idx(session, task.task_id, locked_task=task)
session.add(Message(
task_id=task.task_id,
idx=next_idx,
payload={"role": "user", "content": prompt},
kind="software_job_followup",
))
session.execute(update(Task).where(Task.task_id == task.task_id).values(
run_status="running",
run_error=None,
run_owner=INSTANCE or None,
))
job.followup_status = "running"
return FollowupClaim(
job_id=job.job_id,
task_id=task.task_id,
user_id=job.user_id,
action="analyze",
prompt=prompt,
)
def mark_followup_failed(job_id: UUID) -> None:
with session_scope() as session:
session.execute(
update(SoftwareJob)
.where(
SoftwareJob.job_id == job_id,
SoftwareJob.followup_status.in_({"pending", "running"}),
)
.values(followup_status="failed")
)
def finish_analysis_followup(job_id: UUID) -> None:
with session_scope() as session:
job = session.execute(
select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update()
).scalar_one_or_none()
if job is None or job.followup_status != "running":
return
task_status = session.execute(
select(Task.run_status).where(Task.task_id == job.task_id)
).scalar_one_or_none()
job.followup_status = (
"failed" if task_status in {"error", "cancelled"} else "completed"
)
def _track(app, task: asyncio.Task) -> None:
app.state.aux_tasks.add(task)
task.add_done_callback(app.state.aux_tasks.discard)
async def dispatch_followup(app, job_id: UUID) -> bool:
claim = await asyncio.to_thread(claim_followup, job_id)
if claim is None:
return False
if claim.action == "report":
return True
try:
run_task = schedule_claimed_run(
app,
claim.task_id,
claim.user_id,
claim.prompt,
scheduled=False,
)
except RunScheduleError:
await asyncio.to_thread(mark_followup_failed, claim.job_id)
return False
async def finish() -> None:
try:
await run_task
finally:
await asyncio.to_thread(finish_analysis_followup, claim.job_id)
_track(app, asyncio.create_task(finish()))
return True
async def dispatch_pending_followups(app) -> None:
for job_id in await asyncio.to_thread(pending_followup_ids):
if app.state.draining.is_set():
return
await dispatch_followup(app, job_id)
def start_software_followup_dispatcher(app) -> asyncio.Task:
async def loop() -> None:
while True:
await dispatch_pending_followups(app)
await asyncio.sleep(3)
return asyncio.create_task(loop(), name="software-followups")