zcbot/web/software_followups.py

240 lines
7.8 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.

"""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
@dataclass(frozen=True)
class FollowupClaim:
job_id: UUID
task_id: UUID
user_id: UUID
action: str
prompt: str = ""
def _artifact_refs(manifest: list) -> list[dict]:
refs: list[dict] = []
for item in manifest:
if not isinstance(item, dict) or not item.get("artifact_id") or not item.get("path"):
continue
refs.append({
"path": item["path"],
"label": item.get("filename") or item["path"].rsplit("/", 1)[-1],
"artifact_id": item["artifact_id"],
"version": 2,
})
return refs
def _report_text(job: SoftwareJob) -> str:
contract = get_contract(job.capability)
output_dir = f"{contract.output_namespace}/{job.job_id}"
names = [
str(item.get("filename"))
for item in job.artifact_manifest
if isinstance(item, dict) and item.get("artifact_id") and item.get("filename")
]
files = "".join(names) if names else "无可发布文件"
return (
f"专业软件任务已完成:{contract.display_name}\n\n"
f"- Job ID`{job.job_id}`\n"
f"- 输出目录:`{output_dir}`\n"
f"- 结果文件:{files}"
)
def _analysis_prompt(job: SoftwareJob) -> str:
contract = get_contract(job.capability)
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 == "succeeded",
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:
"""task 空闲时原子领取回调report 在事务内直接落固定 assistant 消息。"""
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 or task.run_status in {"running", "cancelling"}:
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 != "succeeded"
or job.followup_status != "pending"
):
return None
next_idx = allocate_message_idx(session, task.task_id, locked_task=task)
if job.completion_action == "report":
session.add(Message(
task_id=task.task_id,
idx=next_idx,
payload={"role": "assistant", "content": _report_text(job)},
artifact_refs=_artifact_refs(job.artifact_manifest),
kind="software_job_report",
))
job.followup_status = "completed"
return FollowupClaim(
job_id=job.job_id,
task_id=task.task_id,
user_id=job.user_id,
action="report",
)
prompt = _analysis_prompt(job)
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")