atlas-iac/services/hermes/scripts/suite_recovery.py

125 lines
5.4 KiB
Python

"""Bounded isolated retries and volatile, identity-bound validated checkpoints."""
from __future__ import annotations
import copy
import time
from collections import Counter
from suite_contract import (EXECUTION_REVISION, MODELS, PROMPT_REVISION, REVISION,
Problem, digest)
from suite_policy import POLICY_REVISION
MAX_ATTEMPTS = 3
MAX_INVOCATIONS = 128
RETRIABLE = {
"incomplete_generation", "pass_missing_terminal_event", "pass_turn_limit",
"pass_schema_repair", "pass_provider_transient_failure", "rate_limit",
"backend_unavailable", "backend_timeout_or_unavailable", "pass_timeout", "timeout",
}
REPAIRABLE = {"invalid_json_result", "invalid_case_assignments"}
NONREPAIR_STAGES = {"runtime_model", "runtime_model_limits", "cli_initialization", "cli_usage"}
class Checkpoints:
"""Keep validated derived outputs only in this job's memory, until its deadline.
No disk persistence, cross-job reuse, restart recovery or status content access.
Complete call/context hashes include relevant prior-pass outputs.
"""
def __init__(self, request, provider, scope, deadline):
self.identity = {"request": digest(request), "provider": provider,
"model": MODELS[provider]["model"], "routing": request["routing"],
"credential_scope": scope, "configuration": REVISION,
"prompt": PROMPT_REVISION, "policy": POLICY_REVISION,
"execution": EXECUTION_REVISION}
self.deadline, self.values, self.reused = deadline, {}, []
def key(self, call):
"""Bind the whole source, schema, objective and prior results to each pass."""
return digest({**self.identity, "call": call})
def get(self, key, stage):
"""Return an isolated copy; expired checkpoints can never be restored."""
if time.monotonic() >= self.deadline:
self.values.clear()
raise Problem("job_time_budget_exhausted", 504)
if key not in self.values:
return None
self.reused.append(stage)
return copy.deepcopy(self.values[key][1])
def put(self, key, stage, value):
self.values[key] = (stage, copy.deepcopy(value))
def names(self):
return sorted({stage for stage, _ in self.values.values()})
def clear(self):
self.values.clear()
def recover(workflow, call, source, limits, validate, report):
"""Retry only eligible failures; every attempt is fresh and fully accounted."""
stage = call["stage"]
key = workflow.checkpoints.key(call)
cached = workflow.checkpoints.get(key, stage)
if cached is not None:
validate(cached)
return cached
repair_used = False
for attempt in range(1, MAX_ATTEMPTS + 1):
workflow.checkpoint()
turns = min(limits["max_turns"], 3 + attempt)
actual = {**call, "max_turns": turns}
if attempt > 1:
# No prior generated content is fed back; only the unchanged objective
# and complete source enter a new isolated CLI session.
report({"retry_attempt": attempt, "retry_reason": previous.code, "cli_running": False})
if workflow.cancel.wait(min(10, 2 ** (attempt - 1))):
raise Problem("cancelled", 409)
try:
value = workflow.invoke(actual, source, {**limits, "max_turns": turns}, attempt, report)
validate(value)
except Problem as exc:
previous = exc
if workflow.attempts and workflow.attempts[-1]["stage"] == stage:
workflow.attempts[-1].update(status="failed", error_code=exc.code)
can_repair = (exc.code in REPAIRABLE and not repair_used
and exc.details.get("failure_stage") not in NONREPAIR_STAGES)
retry = exc.code in RETRIABLE or can_repair
if not retry:
raise
if can_repair:
repair_used = True
if exc.code == "pass_turn_limit" and turns >= limits["max_turns"]:
code = "pass_turn_limit_exhausted"
break
if attempt == MAX_ATTEMPTS:
code = ("pass_turn_limit_exhausted" if exc.code == "pass_turn_limit" else
"pass_provider_transient_failure" if exc.code in {
"rate_limit", "backend_unavailable", "backend_timeout_or_unavailable",
"pass_provider_transient_failure"} else
"pass_timeout" if exc.code in {"timeout", "pass_timeout"} else
"pass_generation_retries_exhausted")
break
workflow.retries[stage] += 1
continue
workflow.checkpoints.put(key, stage, value)
report({"cli_running": False, "pass_state": "validated"})
return value
raise Problem(code, 502, **{**previous.details, "attempts": attempt,
"last_error_code": previous.code, "review_pass": stage}) from None
def allocate(workflow):
"""Protect explicit later-pass time floors within the single absolute deadline."""
remaining = workflow.deadline - time.monotonic()
reserved = workflow.future_seconds
if remaining <= 0:
raise Problem("job_time_budget_exhausted", 504)
if remaining - reserved < 1:
raise Problem("insufficient_remaining_pass_budget", 422,
remaining_seconds=round(remaining, 3), reserved_future_seconds=reserved)
return remaining - reserved