125 lines
5.4 KiB
Python
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
|