deerflow-code/offline-backend-20260512/backend/app/report_collaboration/execution/recovery.py
2026-09-07 18:24:55 +08:00

158 lines
6.2 KiB
Python

"""Crash recovery: interrupt leftover attempts; keep validated work."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from app.report_collaboration.contracts.node_machine import require_transition
from deerflow.persistence.report_collaboration.base import ReportCollaborationStore
from deerflow.persistence.report_collaboration.codec import iso_now, make_node_run_id
_ORPHAN_TO_INTERRUPTED = frozenset({"queued", "running"})
_ORPHAN_TO_SUPERSEDED = frozenset({"validating"})
@dataclass
class RecoveryReport:
interrupted_node_run_ids: list[str] = field(default_factory=list)
new_attempt_ids: list[str] = field(default_factory=list)
preserved_completed: list[str] = field(default_factory=list)
def _max_attempts(node: dict[str, Any], spec: dict[str, Any] | None) -> int:
contract = node.get("contract_json") if isinstance(node.get("contract_json"), dict) else {}
return int(contract.get("max_attempts") or (spec or {}).get("max_attempts") or 3)
async def recover_orphaned_attempts(store: ReportCollaborationStore, run_id: str) -> RecoveryReport:
"""Mark leftover in-flight attempts and open a new planned attempt when allowed.
Completed nodes and ``awaiting_input`` pauses are left untouched so a
takeover does not re-search finished work or drop a human question.
"""
run = await store.get_run(run_id)
if run is None:
raise KeyError(run_id)
plan = await store.get_plan(str(run.get("plan_id") or ""))
specs = {str(item.get("id")): item for item in (plan or {}).get("nodes") or [] if item.get("id")}
nodes = await store.list_node_runs(run_id)
latest: dict[str, dict[str, Any]] = {}
for item in nodes:
current = latest.get(item["node_id"])
if current is None or int(item["attempt"]) >= int(current["attempt"]):
latest[item["node_id"]] = item
report = RecoveryReport()
for node_id, node in latest.items():
status = str(node.get("status") or "")
if status == "completed":
report.preserved_completed.append(str(node["node_run_id"]))
continue
if status == "awaiting_input":
continue
if status in _ORPHAN_TO_INTERRUPTED:
require_transition(status, "interrupted")
node["status"] = "interrupted"
elif status in _ORPHAN_TO_SUPERSEDED:
require_transition(status, "superseded")
node["status"] = "superseded"
else:
continue
node["ended_at"] = iso_now()
node["error"] = node.get("error") or "lease_reclaimed"
await store.save_node_run(node)
report.interrupted_node_run_ids.append(str(node["node_run_id"]))
await store.append_event(
session_id=str(run["session_id"]),
run_id=run_id,
event_type="node.status.changed",
data={"node_id": node_id, "node_run_id": node["node_run_id"], "status": node["status"], "recovered": True},
)
await _maybe_open_next_attempt(store, run, node, specs.get(node_id), report)
return report
async def cancel_in_flight_nodes(store: ReportCollaborationStore, run_id: str) -> list[str]:
"""Best-effort cancel leftover work after ``finalize_cancel``."""
run = await store.get_run(run_id)
if run is None:
return []
nodes = await store.list_node_runs(run_id)
latest: dict[str, dict[str, Any]] = {}
for item in nodes:
current = latest.get(item["node_id"])
if current is None or int(item["attempt"]) >= int(current["attempt"]):
latest[item["node_id"]] = item
changed: list[str] = []
for node in latest.values():
status = str(node.get("status") or "")
target: str | None = None
if status in {"planned", "queued", "awaiting_input"}:
target = "cancelled"
elif status == "running":
require_transition("running", "cancel_requested")
node["status"] = "cancel_requested"
await store.save_node_run(node)
target = "cancelled"
elif status in {"validating", "interrupted", "failed"}:
target = "superseded"
elif status == "cancel_requested":
target = "cancelled"
if target is None:
continue
require_transition(node["status"], target)
node["status"] = target
node["ended_at"] = iso_now()
await store.save_node_run(node)
changed.append(str(node["node_run_id"]))
await store.append_event(
session_id=str(run["session_id"]),
run_id=run_id,
event_type="node.status.changed",
data={"node_id": node["node_id"], "node_run_id": node["node_run_id"], "status": target, "cancelled": True},
)
return changed
async def _maybe_open_next_attempt(
store: ReportCollaborationStore,
run: dict[str, Any],
node: dict[str, Any],
spec: dict[str, Any] | None,
report: RecoveryReport,
) -> None:
attempt = int(node.get("attempt") or 1)
if attempt >= _max_attempts(node, spec):
return
nxt = attempt + 1
node_run_id = make_node_run_id(str(run["id"]), str(node["node_id"]), nxt)
await store.save_node_run(
{
"node_run_id": node_run_id,
"session_id": run["session_id"],
"run_id": run["id"],
"node_id": node["node_id"],
"attempt": nxt,
"status": "planned",
"label": node.get("label") or node["node_id"],
"role_display_name": node.get("role_display_name"),
"angle": node.get("angle"),
"contract_json": node.get("contract_json"),
"input_artifact_ids": list(node.get("input_artifact_ids") or []),
"command_overlay_revision": node.get("command_overlay_revision") or 0,
"agent_snapshot": node.get("agent_snapshot"),
"error": None,
"started_at": None,
"ended_at": None,
}
)
report.new_attempt_ids.append(node_run_id)
await store.append_event(
session_id=str(run["session_id"]),
run_id=str(run["id"]),
event_type="node.attempt.created",
data={"node_id": node["node_id"], "node_run_id": node_run_id, "attempt": nxt},
)