240 lines
7.8 KiB
Python
240 lines
7.8 KiB
Python
"""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")
|