158 lines
6.2 KiB
Python
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},
|
|
)
|