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