"""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")