← Files Biohub ESMARCHIVED FILE
scripts/biohub_esm_lib/atlas_jobs.py
71.6 KB · Sep 30, 2026 · 23:14 UTC
"""Durable, credential-safe state for public Atlas batch jobs."""
from __future__ import annotations
import fcntl
import json
import math
import threading
import time
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, BinaryIO, Callable
from .atlas import validate_job_id
from .diagnostics import (
DIAGNOSTIC_MAX_COLLECTION_ITEMS,
bounded_provider_response_diagnostic,
)
from .errors import APIError, BiohubESMError, SchemaDriftError, ValidationError
from .provenance import input_digest, utc_now, write_json_atomic
from .security import redact
ATLAS_BATCH_STATE_SCHEMA = "1.1"
ATLAS_BATCH_LEGACY_SCHEMAS = {"1.0"}
ATLAS_BATCH_PROVIDER = "biohub-esm-atlas-v1alpha1"
ATLAS_BATCH_PROVIDER_CALL_HISTORY_LIMIT = 256
ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES = {400, 401, 402, 403, 404, 422, 429}
ATLAS_BATCH_STATE_MAX_BYTES = 4 * 1024 * 1024
ATLAS_BATCH_ARTIFACT_LIMIT = 2_048
ATLAS_BATCH_STATE_MAX_NUMBER_CHARACTERS = 128
ATLAS_BATCH_TERMINAL_STATUSES = {"completed", "cancelled", "failed", "expired"}
ATLAS_BATCH_STATUSES = {
"submitting",
"submission-rejected",
"submission-indeterminate",
"pending",
"completed",
"cancelled",
"cancellation-requested",
"failed",
"expired",
}
def _submission_reconciliation_guidance(provider_acceptance: str) -> str:
prefix = {
"accepted": "Provider acceptance was confirmed, but no job_id was durably captured. ",
"not-proven": "The provider returned an error, but remote side effects were not proven. ",
"unknown": "Provider acceptance is unknown and no job_id was durably captured. ",
}.get(provider_acceptance)
if prefix is None:
raise ValidationError("Atlas submission acceptance reconciliation is invalid")
return prefix + "Reconcile with Atlas operators before any new submission."
@contextmanager
def atlas_batch_sink_preparer(
client: Any,
*,
temporary_file: Callable[..., BinaryIO],
) -> Iterator[Callable[[], BinaryIO | None]]:
"""Lazily reserve an anonymous sink before any durable submit transition."""
handle: BinaryIO | None = None
def prepare() -> BinaryIO | None:
nonlocal handle
if getattr(client, "_supports_synchronous_sink", False) is not True:
return None
if handle is None:
try:
handle = temporary_file(mode="w+b")
except OSError as exc:
raise ValidationError(
"Atlas synchronous response spool could not be created"
) from exc
return handle
active_failure: BaseException | None = None
try:
yield prepare
except BaseException as exc:
active_failure = exc
raise
finally:
if handle is not None and not handle.closed:
try:
handle.close()
except (OSError, ValueError) as exc:
if active_failure is None:
raise ValidationError(
"Atlas synchronous response spool could not be closed"
) from exc
def _batch_provider_call(
*, endpoint: str, operation: str, http_status: int | None
) -> dict[str, Any]:
call = {
"endpoint": endpoint,
"operation": operation,
"http_status": http_status,
"timestamp": utc_now(),
}
if http_status is not None:
if http_status < 200:
call["outcome"] = "indeterminate"
call["operation_indeterminate"] = True
else:
call["outcome"] = "success"
call["operation_indeterminate"] = False
return call
def _atlas_exception_provider_call(
client: Any,
*,
endpoint: str,
operation: str,
exc: BaseException,
) -> dict[str, Any]:
error_http_status = getattr(exc, "status", None)
if error_http_status is None:
error_http_status = getattr(exc, "provider_status", None)
if (
isinstance(error_http_status, bool)
or not isinstance(error_http_status, int)
or not 100 <= error_http_status <= 599
):
error_http_status = None
calls = getattr(client, "provider_calls", None)
matching_call = None
if isinstance(calls, list):
matching_call = next(
(
candidate
for candidate in reversed(calls)
if isinstance(candidate, dict) and candidate.get("endpoint") == endpoint
),
None,
)
if matching_call is None:
call = _batch_provider_call(
endpoint=endpoint,
operation=operation,
http_status=error_http_status,
)
else:
call = dict(matching_call)
if error_http_status is None:
candidate_status = call.get("http_status")
if (
isinstance(candidate_status, int)
and not isinstance(candidate_status, bool)
and 100 <= candidate_status <= 599
):
error_http_status = candidate_status
if isinstance(exc, BiohubESMError) and error_http_status is not None:
exc.provider_status = error_http_status
if isinstance(exc, APIError) and error_http_status is not None:
if error_http_status < 200:
exc.operation_indeterminate = True
exc.response_body_partial = True
else:
exc.operation_indeterminate = False
call["endpoint"] = endpoint
call["operation"] = operation
call["error_kind"] = getattr(exc, "kind", exc.__class__.__name__)
if isinstance(exc, APIError):
if error_http_status is not None:
call["http_status"] = error_http_status
call.pop("retry_after", None)
call.update(
{
"outcome": ("indeterminate" if exc.operation_indeterminate else "error"),
"partial": exc.partial,
"response_body_partial": exc.response_body_partial,
"operation_indeterminate": exc.operation_indeterminate,
}
)
if exc.retry_after is not None:
call["retry_after"] = exc.retry_after
elif isinstance(exc, BiohubESMError):
if call.get("http_status") is None and error_http_status is not None:
call["http_status"] = error_http_status
operation_indeterminate = bool(getattr(exc, "operation_indeterminate", False))
call["outcome"] = "indeterminate" if operation_indeterminate else "error"
call["operation_indeterminate"] = operation_indeterminate
call["partial"] = bool(getattr(exc, "partial", False))
else:
operation_indeterminate = error_http_status is None or error_http_status < 200
if error_http_status is not None:
call["http_status"] = error_http_status
call["outcome"] = "indeterminate" if operation_indeterminate else "error"
call["operation_indeterminate"] = operation_indeterminate
return call
def _atlas_failure_record(exc: BiohubESMError) -> dict[str, Any]:
failure: dict[str, Any] = {
"kind": getattr(exc, "kind", exc.__class__.__name__),
"message": str(exc),
}
if isinstance(exc, APIError):
failure.update(
{
"status": exc.status,
"retry_after": exc.retry_after,
"partial": exc.partial,
"response_body_partial": exc.response_body_partial,
"operation_indeterminate": exc.operation_indeterminate,
}
)
provider_status = getattr(exc, "provider_status", None)
if (
isinstance(provider_status, int)
and not isinstance(provider_status, bool)
and 100 <= provider_status <= 599
):
failure["provider_status"] = provider_status
if isinstance(exc, SchemaDriftError) and exc.diagnostic_path:
failure["diagnostic_path"] = exc.diagnostic_path
return redact(failure)
def _utc_datetime(value: Any, field: str) -> datetime:
if not isinstance(value, str):
raise ValidationError(f"Atlas batch {field} must be a UTC timestamp")
try:
normalized = value[:-1] + "+00:00" if value.endswith("Z") else value
parsed = datetime.fromisoformat(normalized)
except ValueError as exc:
raise ValidationError(f"Atlas batch {field} must be a UTC timestamp") from exc
if parsed.tzinfo is None or parsed.utcoffset() != timezone.utc.utcoffset(parsed):
raise ValidationError(f"Atlas batch {field} must be a UTC timestamp")
return parsed
def _utc_string(value: datetime) -> str:
return value.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
def _reject_state_json_constant(value: str) -> None:
raise ValueError(f"non-finite JSON constant: {value}")
def _bounded_state_json_int(value: str) -> int:
if len(value) > ATLAS_BATCH_STATE_MAX_NUMBER_CHARACTERS:
raise ValueError("Atlas batch state integer exceeds the numeric safety bound")
return int(value)
def _finite_state_json_float(value: str) -> float:
if len(value) > ATLAS_BATCH_STATE_MAX_NUMBER_CHARACTERS:
raise ValueError("Atlas batch state float exceeds the numeric safety bound")
result = float(value)
if not math.isfinite(result):
raise ValueError("Atlas batch state contains a non-finite number")
mantissa = value.lower().split("e", 1)[0]
if result == 0.0 and any(character in "123456789" for character in mantissa):
raise ValueError("Atlas batch state number underflows the finite float range")
return result
_RESPONSE_PRIORITY_KEYS = (
"job_id",
"status",
"delivery",
"completed_count",
"total_count",
"download_url",
"artifact_materialization",
"download",
"provider_call_result",
"provider_acceptance",
"retry_safe",
"retry_after",
"retry_not_before",
"reason",
"retry",
"intent_persisted",
)
def _omit_download_capabilities(value: dict[str, Any]) -> None:
"""Remove capabilities from an already bounded diagnostic projection."""
stack: list[Any] = [value]
while stack:
item = stack.pop()
if isinstance(item, dict):
for key, child in item.items():
if key.lower() == "download_url":
item[key] = "[EPHEMERAL URL OMITTED]"
elif isinstance(child, (dict, list)):
stack.append(child)
elif isinstance(item, list):
stack.extend(child for child in item if isinstance(child, (dict, list)))
def _safe_response(value: dict[str, Any]) -> dict[str, Any]:
"""Keep bounded alpha-schema diagnostics and required local state semantics."""
prioritized: dict[str, Any] = {}
for key in _RESPONSE_PRIORITY_KEYS:
if key in value:
prioritized[key] = value[key]
for key, child in value.items():
text_key = str(key)
if text_key in prioritized:
continue
if len(prioritized) >= DIAGNOSTIC_MAX_COLLECTION_ITEMS:
break
prioritized[text_key] = child
diagnostic = bounded_provider_response_diagnostic(prioritized)
sanitized = diagnostic.get("projection")
if not isinstance(sanitized, dict): # pragma: no cover - defensive
raise ValidationError("Atlas batch response must be an object")
_omit_download_capabilities(sanitized)
source_truncated = len(value) > len(prioritized)
if diagnostic.get("truncated") or source_truncated:
reasons = set(diagnostic.get("truncation_reasons", []))
if source_truncated:
reasons.add("collection-item-limit")
sanitized["response_diagnostic"] = {
"representation": diagnostic.get("representation"),
"truncated": True,
"truncation_reasons": sorted(reasons),
"visited_nodes": diagnostic.get("visited_nodes"),
"retained_text_bytes": diagnostic.get("retained_text_bytes"),
"projection_sha256": diagnostic.get("projection_sha256"),
}
return sanitized
_PROVIDER_CALL_PRIORITY_KEYS = (
"endpoint",
"operation",
"method",
"outcome",
"http_status",
"started_at",
"finished_at",
"timestamp",
"error_kind",
"partial",
"response_body_partial",
"operation_indeterminate",
"retry_after",
"transport_attempts",
"resume_recovery",
"observed_status",
"state_status_before",
"state_status_after",
"transition_ignored",
"error",
)
_ARTIFACT_KEYS = (
"path",
"size_bytes",
"media_type",
"sha256",
"evidence_class",
)
def _safe_provider_call(value: dict[str, Any]) -> dict[str, Any]:
prioritized: dict[str, Any] = {}
for key in _PROVIDER_CALL_PRIORITY_KEYS:
if key in value:
prioritized[key] = value[key]
for key, child in value.items():
text_key = str(key)
if text_key in prioritized:
continue
if len(prioritized) >= DIAGNOSTIC_MAX_COLLECTION_ITEMS:
break
prioritized[text_key] = child
safe = bounded_provider_response_diagnostic(prioritized).get("projection")
if (
not isinstance(safe, dict)
or not isinstance(safe.get("endpoint"), str)
or not safe["endpoint"].strip()
or not isinstance(safe.get("operation"), str)
or not safe["operation"].strip()
):
raise ValidationError("Atlas batch provider call requires endpoint and operation")
http_status = safe.get("http_status")
if http_status is not None and (
isinstance(http_status, bool)
or not isinstance(http_status, int)
or not 100 <= http_status <= 599
):
raise ValidationError("Atlas batch provider call HTTP status is invalid")
for boolean_field in (
"partial",
"response_body_partial",
"operation_indeterminate",
):
if boolean_field in safe and not isinstance(safe[boolean_field], bool):
raise ValidationError(f"Atlas batch provider call {boolean_field} must be boolean")
operation_indeterminate = safe.get("operation_indeterminate")
outcome = safe.get("outcome")
if operation_indeterminate is True and outcome != "indeterminate":
raise ValidationError(
"Atlas batch indeterminate provider call must have indeterminate outcome"
)
if outcome == "indeterminate" and operation_indeterminate is not True:
raise ValidationError("Atlas batch indeterminate outcome requires operation indeterminacy")
if operation_indeterminate is True and http_status is not None and http_status >= 200:
raise ValidationError(
"Atlas batch provider call with final HTTP status cannot be indeterminate"
)
if (
isinstance(http_status, int)
and 100 <= http_status < 200
and (outcome != "indeterminate" or operation_indeterminate is not True)
):
raise ValidationError(
"Atlas batch informational status requires an indeterminate provider call"
)
return safe
def _state_value_changed(before: Any, after: Any) -> bool:
try:
return before != after
except (RecursionError, MemoryError):
return True
def _sanitize_state_collections(
value: dict[str, Any], *, normalize_missing_indeterminate: bool = False
) -> bool:
changed = False
calls = value.get("provider_calls")
if isinstance(calls, list):
safe_calls = []
certainty_migrated = False
for call in calls:
candidate = dict(call) if isinstance(call, dict) else call
if normalize_missing_indeterminate and isinstance(candidate, dict):
call_status = candidate.get("http_status")
operation = candidate.get("operation")
outcome = candidate.get("outcome")
informational = (
isinstance(call_status, int)
and not isinstance(call_status, bool)
and 100 <= call_status < 200
)
final_status = (
isinstance(call_status, int)
and not isinstance(call_status, bool)
and call_status >= 200
)
if informational and (
outcome != "indeterminate"
or candidate.get("operation_indeterminate") is not True
):
candidate["outcome"] = "indeterminate"
candidate["operation_indeterminate"] = True
candidate["certainty_normalized_from_legacy"] = True
certainty_migrated = True
elif outcome == "indeterminate" and final_status:
candidate["outcome"] = (
"accepted"
if operation in {"submit", "cancel"} and call_status < 300
else "error"
)
candidate["operation_indeterminate"] = False
candidate["certainty_normalized_from_legacy"] = True
certainty_migrated = True
elif (
outcome == "indeterminate"
and candidate.get("operation_indeterminate") is not True
):
candidate["operation_indeterminate"] = True
candidate["certainty_normalized_from_legacy"] = True
certainty_migrated = True
safe_calls.append(
_safe_provider_call(candidate) if isinstance(candidate, dict) else candidate
)
changed = _state_value_changed(calls, safe_calls) or changed
value["provider_calls"] = safe_calls
if (
normalize_missing_indeterminate
and certainty_migrated
and value.get("status") == "submission-indeterminate"
):
latest_submit = next(
(
call
for call in reversed(safe_calls)
if isinstance(call, dict) and call.get("operation") == "submit"
),
None,
)
reconciliation = value.get("manual_reconciliation")
if isinstance(latest_submit, dict) and isinstance(reconciliation, dict):
provider_acceptance = {
"accepted": "accepted",
"error": "not-proven",
"indeterminate": "unknown",
}.get(latest_submit.get("outcome"))
call_status = latest_submit.get("http_status")
legacy_definitive_rejection = (
latest_submit.get("outcome") == "error"
and call_status in ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES
)
if legacy_definitive_rejection:
normalized_reconciliation = dict(reconciliation)
normalized_reconciliation.pop("provider_acceptance", None)
normalized_reconciliation["guidance"] = (
"The provider returned a definitive rejection, but this legacy "
"indeterminate record remains do-not-resubmit until Atlas operators "
"reconcile it."
)
changed = (
_state_value_changed(reconciliation, normalized_reconciliation) or changed
)
value["manual_reconciliation"] = normalized_reconciliation
elif provider_acceptance is not None:
normalized_reconciliation = dict(reconciliation)
normalized_reconciliation["provider_acceptance"] = provider_acceptance
normalized_reconciliation["guidance"] = _submission_reconciliation_guidance(
provider_acceptance
)
changed = (
_state_value_changed(reconciliation, normalized_reconciliation) or changed
)
value["manual_reconciliation"] = normalized_reconciliation
normalized_status = latest_submit.get("http_status")
changed = (
_state_value_changed(value.get("last_http_status"), normalized_status)
or changed
)
value["last_http_status"] = normalized_status
if (
normalize_missing_indeterminate
and certainty_migrated
and value.get("status") == "cancellation-requested"
):
latest_cancel = next(
(
call
for call in reversed(safe_calls)
if isinstance(call, dict) and call.get("operation") == "cancel"
),
None,
)
latest_call = safe_calls[-1] if safe_calls else None
if (
latest_cancel is latest_call
and isinstance(latest_cancel, dict)
and latest_cancel.get("outcome")
in {
"accepted",
"error",
}
):
response = value.get("last_response")
if isinstance(response, dict):
normalized_response = dict(response)
normalized_response["provider_acceptance"] = (
"confirmed" if latest_cancel["outcome"] == "accepted" else "rejected"
)
normalized_response["retry_safe"] = latest_cancel["outcome"] != "accepted"
changed = _state_value_changed(response, normalized_response) or changed
value["last_response"] = normalized_response
normalized_status = latest_cancel.get("http_status")
changed = (
_state_value_changed(value.get("last_http_status"), normalized_status)
or changed
)
value["last_http_status"] = normalized_status
artifacts = value.get("artifacts")
if isinstance(artifacts, list):
if len(artifacts) > ATLAS_BATCH_ARTIFACT_LIMIT:
raise ValidationError("Atlas batch artifact history exceeds the safety limit")
safe_artifacts = [
{key: artifact[key] for key in _ARTIFACT_KEYS if key in artifact}
if isinstance(artifact, dict)
else artifact
for artifact in artifacts
]
changed = _state_value_changed(artifacts, safe_artifacts) or changed
value["artifacts"] = safe_artifacts
response = value.get("last_response")
if isinstance(response, dict):
safe_response = _safe_response(response)
changed = _state_value_changed(response, safe_response) or changed
value["last_response"] = safe_response
return changed
def _state_json_exceeds_limit(value: dict[str, Any]) -> bool:
encoder = json.JSONEncoder(
indent=2,
sort_keys=True,
ensure_ascii=False,
allow_nan=False,
)
size = 1
for chunk in encoder.iterencode(value):
size += len(chunk.encode("utf-8"))
if size > ATLAS_BATCH_STATE_MAX_BYTES:
return True
return False
def _increment_history_count(counts: dict[str, int], value: Any) -> None:
key = str(value) if value is not None else "unspecified"
counts[key] = counts.get(key, 0) + 1
def _validated_provider_call_history(
history: Any,
) -> tuple[int, str, dict[str, int], dict[str, int]]:
required = {
"schema_version",
"compacted_call_count",
"chain_sha256",
"operation_counts",
"outcome_counts",
}
if not isinstance(history, dict) or set(history) != required:
raise ValidationError("Atlas batch provider call history is invalid")
count = history.get("compacted_call_count")
chain_sha256 = history.get("chain_sha256")
if (
history.get("schema_version") != "1.0"
or isinstance(count, bool)
or not isinstance(count, int)
or count < 1
or not isinstance(chain_sha256, str)
or len(chain_sha256) != 64
or any(character not in "0123456789abcdef" for character in chain_sha256)
):
raise ValidationError("Atlas batch provider call history is invalid")
validated_counts: list[dict[str, int]] = []
for field in ("operation_counts", "outcome_counts"):
counts = history.get(field)
if not isinstance(counts, dict) or any(
not isinstance(key, str)
or not key
or isinstance(item_count, bool)
or not isinstance(item_count, int)
or item_count < 0
for key, item_count in counts.items()
):
raise ValidationError(f"Atlas batch provider call history {field} is invalid")
if sum(counts.values()) != count:
raise ValidationError(f"Atlas batch compacted {field} count is inconsistent")
validated_counts.append(dict(counts))
return count, chain_sha256, validated_counts[0], validated_counts[1]
def _semantic_provider_call_indices(value: dict[str, Any], calls: list[Any]) -> set[int]:
"""Retain calls still required to reconcile the durable control state."""
status = value.get("status")
required: set[int] = set()
if status in {"submitting", "submission-rejected", "submission-indeterminate"}:
for index in range(len(calls) - 1, -1, -1):
call = calls[index]
if isinstance(call, dict) and call.get("operation") == "submit":
required.add(index)
break
if status == "cancellation-requested":
for index in range(len(calls) - 1, -1, -1):
call = calls[index]
if not isinstance(call, dict) or call.get("operation") != "cancel":
continue
if call.get("outcome") in {"in-flight", "error", "indeterminate"}:
required.add(index)
break
return required
def _compact_provider_calls(value: dict[str, Any]) -> bool:
"""Bound retained attempts while chaining deterministic evidence for older calls."""
calls = value.get("provider_calls")
limit = ATLAS_BATCH_PROVIDER_CALL_HISTORY_LIMIT
if not isinstance(calls, list):
return False
if isinstance(limit, bool) or not isinstance(limit, int) or limit < 1:
raise ValidationError("Atlas batch provider call history limit is invalid")
if len(calls) <= limit:
return False
retained_indices = _semantic_provider_call_indices(value, calls)
for index in range(len(calls) - 1, -1, -1):
if len(retained_indices) >= limit:
break
retained_indices.add(index)
dropped = [call for index, call in enumerate(calls) if index not in retained_indices]
retained = [call for index, call in enumerate(calls) if index in retained_indices]
existing = value.get("provider_call_history")
if existing is None:
count = 0
previous_sha256 = "0" * 64
operation_counts: dict[str, int] = {}
outcome_counts: dict[str, int] = {}
else:
count, previous_sha256, operation_counts, outcome_counts = _validated_provider_call_history(
existing
)
for call in dropped:
if not isinstance(call, dict):
raise ValidationError("Atlas batch provider calls must be objects")
_increment_history_count(operation_counts, call.get("operation"))
_increment_history_count(outcome_counts, call.get("outcome"))
value["provider_call_history"] = {
"schema_version": "1.0",
"compacted_call_count": count + len(dropped),
"chain_sha256": input_digest(
{
"previous_chain_sha256": previous_sha256,
"compacted_calls": dropped,
}
),
"operation_counts": operation_counts,
"outcome_counts": outcome_counts,
}
value["provider_calls"] = retained
return True
class AtlasBatchStore:
"""Atomic, locked job-state persistence suitable across CLI invocations."""
def __init__(self, path: Path) -> None:
self.path = path
self._thread_lock = threading.RLock()
self._lock_depth = 0
@property
def provenance_path(self) -> Path:
return self.path.with_name(f"{self.path.stem}.provenance.json")
@property
def lock_path(self) -> Path:
return self.path.with_name(f".{self.path.name}.lock")
@contextmanager
def operation_lock(self, *, deadline: float | None = None) -> Iterator[None]:
"""Serialize one state/provider/provenance operation per durable path.
The thread lock makes the flock safely reentrant for methods called by a
command that already owns the operation lock. Separately constructed
stores and other processes still contend on the same lock file. When a
monotonic deadline is supplied, both contention points are bounded.
"""
if deadline is not None and (
isinstance(deadline, bool)
or not isinstance(deadline, (int, float))
or not math.isfinite(deadline)
):
raise ValidationError("Atlas batch operation deadline must be finite")
if deadline is None:
acquired = self._thread_lock.acquire()
else:
remaining = deadline - time.monotonic()
acquired = (
self._thread_lock.acquire(blocking=False)
if remaining <= 0
else self._thread_lock.acquire(timeout=min(remaining, threading.TIMEOUT_MAX))
)
if not acquired:
raise APIError(
status=None,
kind="timeout",
message="Atlas batch operation lock timed out",
)
try:
if self._lock_depth:
self._lock_depth += 1
try:
yield
finally:
self._lock_depth -= 1
return
self.path.parent.mkdir(parents=True, exist_ok=True)
with self.lock_path.open("a+b") as handle:
if deadline is None:
fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
else:
while True:
remaining = deadline - time.monotonic()
if remaining <= 0:
raise APIError(
status=None,
kind="timeout",
message="Atlas batch operation lock timed out",
)
try:
fcntl.flock(handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
break
except BlockingIOError:
time.sleep(min(0.01, remaining))
self._lock_depth = 1
try:
yield
finally:
self._lock_depth = 0
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
finally:
self._thread_lock.release()
def _load_unlocked(self) -> dict[str, Any]:
if not self.path.is_file():
raise ValidationError(f"Atlas batch state does not exist: {self.path}")
try:
with self.path.open("rb") as handle:
encoded = handle.read(ATLAS_BATCH_STATE_MAX_BYTES + 1)
if len(encoded) > ATLAS_BATCH_STATE_MAX_BYTES:
raise ValidationError("Atlas batch state exceeds the file-size safety limit")
value = json.loads(
encoded.decode("utf-8"),
parse_constant=_reject_state_json_constant,
parse_float=_finite_state_json_float,
parse_int=_bounded_state_json_int,
)
except ValidationError:
raise
except (OSError, UnicodeDecodeError, ValueError, RecursionError, MemoryError) as exc:
raise ValidationError("Atlas batch state is not readable JSON") from exc
if not isinstance(value, dict):
raise ValidationError("Atlas batch state must be a JSON object")
sanitized = _sanitize_state_collections(
value,
normalize_missing_indeterminate=True,
)
compacted = _compact_provider_calls(value)
self.validate(value)
if sanitized or compacted:
return self._save_unlocked(value)
return value
def load(self) -> dict[str, Any]:
with self.operation_lock():
value = self._load_unlocked()
if value["status"] == "submitting":
value = self._mark_submission_indeterminate_unlocked(
value, "process ended during Atlas batch submission"
)
elif value["status"] == "cancellation-requested" and any(
call.get("operation") == "cancel" and call.get("outcome") == "in-flight"
for call in reversed(value["provider_calls"])
):
value = self._mark_cancellation_indeterminate_unlocked(value)
return value
def _save_unlocked(self, value: dict[str, Any]) -> dict[str, Any]:
_sanitize_state_collections(value)
_compact_provider_calls(value)
safe = redact(value)
if not isinstance(safe, dict): # pragma: no cover - defensive
raise ValidationError("Atlas batch state must remain an object")
safe["schema_version"] = ATLAS_BATCH_STATE_SCHEMA
self.validate(safe)
if _state_json_exceeds_limit(safe):
raise ValidationError("Atlas batch state exceeds the file-size safety limit")
write_json_atomic(self.path, safe)
return safe
def save(self, value: dict[str, Any]) -> dict[str, Any]:
with self.operation_lock():
return self._save_unlocked(value)
@staticmethod
def validate(value: Any) -> None:
if not isinstance(value, dict):
raise ValidationError("Atlas batch state must be a JSON object")
required = {
"schema_version",
"provider",
"job_id",
"status",
"submitted_at",
"updated_at",
"endpoint",
"request",
"last_http_status",
"last_response",
"provider_calls",
"artifacts",
"provenance_path",
}
missing = sorted(required - set(value))
if missing:
raise ValidationError(f"Atlas batch state is missing: {', '.join(missing)}")
if value["schema_version"] not in {
ATLAS_BATCH_STATE_SCHEMA,
*ATLAS_BATCH_LEGACY_SCHEMAS,
}:
raise ValidationError("unsupported Atlas batch state schema")
if value["provider"] != ATLAS_BATCH_PROVIDER:
raise ValidationError("Atlas batch state provider is invalid")
if value["job_id"] is not None:
validate_job_id(value["job_id"])
if value["status"] not in ATLAS_BATCH_STATUSES:
raise ValidationError("Atlas batch state status is invalid")
last_http_status = value["last_http_status"]
if last_http_status is not None and (
isinstance(last_http_status, bool)
or not isinstance(last_http_status, int)
or not 100 <= last_http_status <= 599
):
raise ValidationError("Atlas batch last HTTP status is invalid")
if (
value["status"]
in {"pending", "cancellation-requested", "cancelled", "failed", "expired"}
and value["job_id"] is None
):
raise ValidationError("resumable Atlas batch state requires a job_id")
if not isinstance(value["endpoint"], str) or not value["endpoint"].strip():
raise ValidationError("Atlas batch state endpoint is invalid")
request = value["request"]
if (
not isinstance(request, dict)
or not isinstance(request.get("input_sha256"), str)
or len(request["input_sha256"]) != 64
):
raise ValidationError("Atlas batch state request digest is invalid")
if not isinstance(request.get("parameters"), dict):
raise ValidationError("Atlas batch state parameters must be an object")
calls = value["provider_calls"]
if not isinstance(calls, list) or not calls:
raise ValidationError("Atlas batch state requires provider calls")
if len(calls) > ATLAS_BATCH_PROVIDER_CALL_HISTORY_LIMIT:
raise ValidationError("Atlas batch retained provider call history exceeds its limit")
for call in calls:
if not isinstance(call, dict):
raise ValidationError("Atlas batch provider calls must be objects")
_safe_provider_call(call)
history = value.get("provider_call_history")
if history is not None:
_validated_provider_call_history(history)
if not isinstance(value["last_response"], dict):
raise ValidationError("Atlas batch state last_response must be an object")
if not isinstance(value["artifacts"], list):
raise ValidationError("Atlas batch state artifacts must be a list")
if len(value["artifacts"]) > ATLAS_BATCH_ARTIFACT_LIMIT:
raise ValidationError("Atlas batch artifact history exceeds the safety limit")
submission_call = None
if value["status"] in {
"submitting",
"submission-rejected",
"submission-indeterminate",
}:
submission_call = next(
(call for call in reversed(calls) if call.get("operation") == "submit"),
None,
)
if submission_call is None or (submission_call.get("http_status") != last_http_status):
raise ValidationError(
"Atlas batch last HTTP status is not bound to its submit call"
)
if value["status"] == "submitting" and (
calls[-1].get("operation") != "submit" or calls[-1].get("outcome") != "in-flight"
):
raise ValidationError("submitting Atlas batch state requires an in-flight call")
if value["status"] == "submission-rejected":
rejection = calls[-1].get("error")
if (
value["job_id"] is not None
or last_http_status not in ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES
or calls[-1].get("operation") != "submit"
or calls[-1].get("outcome") != "rejected"
or not isinstance(rejection, dict)
or rejection.get("provider_acceptance") != "rejected"
or rejection.get("retry_safe") is not True
or rejection.get("rejected_at") != calls[-1].get("finished_at")
):
raise ValidationError("rejected Atlas batch state requires a rejected submit call")
rejected_at = _utc_datetime(calls[-1].get("finished_at"), "submission rejection time")
retry_not_before = rejection.get("retry_not_before")
if retry_not_before is not None and (
_utc_datetime(retry_not_before, "retry-not-before time") < rejected_at
):
raise ValidationError("Atlas batch retry-not-before precedes its rejection")
if value["status"] == "submission-indeterminate":
reconciliation = value.get("manual_reconciliation")
if (
not isinstance(reconciliation, dict)
or reconciliation.get("input_sha256") != request["input_sha256"]
or reconciliation.get("possible_remote_acceptance") is not True
or reconciliation.get("do_not_resubmit") is not True
):
raise ValidationError(
"indeterminate Atlas batch state requires manual reconciliation"
)
provider_acceptance = reconciliation.get("provider_acceptance")
if provider_acceptance is not None and provider_acceptance not in {
"unknown",
"accepted",
"not-proven",
}:
raise ValidationError("indeterminate Atlas batch provider acceptance is invalid")
if provider_acceptance is not None:
if submission_call is None: # pragma: no cover - checked above
raise ValidationError("indeterminate Atlas batch submit call is missing")
call_status = submission_call.get("http_status")
call_outcome = submission_call.get("outcome")
operation_indeterminate = submission_call.get("operation_indeterminate")
accepted = (
call_outcome == "accepted"
and isinstance(call_status, int)
and not isinstance(call_status, bool)
and 200 <= call_status < 300
and operation_indeterminate is False
)
unknown = (
call_outcome == "indeterminate"
and (call_status is None or call_status < 200)
and operation_indeterminate is True
)
not_proven = (
call_outcome == "error"
and isinstance(call_status, int)
and not isinstance(call_status, bool)
and call_status >= 300
and call_status not in ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES
and operation_indeterminate is False
)
if not {
"accepted": accepted,
"unknown": unknown,
"not-proven": not_proven,
}[provider_acceptance]:
raise ValidationError(
"indeterminate Atlas batch provider acceptance contradicts its submit call"
)
if value["status"] == "cancellation-requested":
cancellation_response = value["last_response"]
provider_acceptance = cancellation_response.get("provider_acceptance")
if provider_acceptance is not None:
latest_cancel = next(
(call for call in reversed(calls) if call.get("operation") == "cancel"),
None,
)
if latest_cancel is None:
raise ValidationError(
"Atlas cancellation response requires a matching cancel call"
)
call_status = latest_cancel.get("http_status")
call_outcome = latest_cancel.get("outcome")
operation_indeterminate = latest_cancel.get("operation_indeterminate")
if call_status != last_http_status:
raise ValidationError(
"Atlas batch last HTTP status is not bound to its cancel call"
)
cancellation_contracts = {
"confirmed": (
call_outcome == "accepted"
and isinstance(call_status, int)
and not isinstance(call_status, bool)
and 200 <= call_status < 300
and operation_indeterminate is False
and cancellation_response.get("retry_safe") is False
),
"rejected": (
call_outcome == "error"
and isinstance(call_status, int)
and not isinstance(call_status, bool)
and call_status >= 200
and operation_indeterminate is False
and cancellation_response.get("retry_safe") is True
),
"unknown": (
call_outcome == "indeterminate"
and (call_status is None or call_status < 200)
and operation_indeterminate is True
and bool(cancellation_response.get("retry_safe"))
),
}
if provider_acceptance not in cancellation_contracts:
raise ValidationError("Atlas cancellation provider acceptance is invalid")
if not cancellation_contracts[provider_acceptance]:
raise ValidationError(
"Atlas cancellation provider acceptance contradicts its cancel call"
)
def initialize(
self,
*,
job_id: str | None,
status: str,
submitted_at: str,
endpoint: str,
input_sha256: str,
parameters: dict[str, Any],
http_status: int | None,
response: dict[str, Any],
provider_call: dict[str, Any],
artifacts: list[dict[str, Any]] | None = None,
) -> dict[str, Any]:
with self.operation_lock():
if self.path.exists():
raise ValidationError("refusing to overwrite an existing Atlas batch state")
if job_id is not None:
validate_job_id(job_id)
now = utc_now()
return self._save_unlocked(
{
"schema_version": ATLAS_BATCH_STATE_SCHEMA,
"provider": ATLAS_BATCH_PROVIDER,
"job_id": job_id,
"status": status,
"submitted_at": submitted_at,
"updated_at": now,
"endpoint": endpoint,
"request": {
"input_sha256": input_sha256,
"parameters": redact(parameters),
},
"last_http_status": http_status,
"last_response": _safe_response(response),
"provider_calls": [_safe_provider_call(provider_call)],
"artifacts": redact(artifacts or []),
"provenance_path": str(self.provenance_path),
}
)
def begin_submission(
self,
*,
submitted_at: str,
endpoint: str,
input_sha256: str,
parameters: dict[str, Any],
) -> dict[str, Any]:
return self.initialize(
job_id=None,
status="submitting",
submitted_at=submitted_at,
endpoint=endpoint,
input_sha256=input_sha256,
parameters=parameters,
http_status=None,
response={"status": "submitting"},
provider_call={
"endpoint": endpoint,
"operation": "submit",
"http_status": None,
"started_at": submitted_at,
"outcome": "in-flight",
},
)
@staticmethod
def _finish_in_flight_call(
state: dict[str, Any],
*,
operation: str,
outcome: str,
http_status: int | None,
) -> None:
call = next(
(
candidate
for candidate in reversed(state["provider_calls"])
if candidate.get("operation") == operation
and candidate.get("outcome") == "in-flight"
),
None,
)
if call is None:
raise ValidationError(f"Atlas batch {operation} has no in-flight provider call")
call["outcome"] = outcome
call["http_status"] = http_status
call["operation_indeterminate"] = outcome == "indeterminate"
call["finished_at"] = utc_now()
@staticmethod
def _merge_provider_call_details(call: dict[str, Any], details: dict[str, Any] | None) -> None:
if details is None:
return
for key, value in details.items():
if key not in {
"endpoint",
"operation",
"started_at",
"finished_at",
"outcome",
"http_status",
"operation_indeterminate",
}:
call[key] = value
def reconcile_submission(
self,
*,
job_id: str | None,
status: str,
endpoint: str,
http_status: int,
response: dict[str, Any],
artifacts: list[dict[str, Any]] | None = None,
provider_call: dict[str, Any] | None = None,
) -> dict[str, Any]:
with self.operation_lock():
state = self._load_unlocked()
if state["status"] != "submitting":
raise ValidationError("Atlas batch submission is not awaiting reconciliation")
if status not in {"pending", "completed"}:
raise ValidationError("Atlas batch submission result status is invalid")
if job_id is not None:
validate_job_id(job_id)
if (status == "pending") != (job_id is not None):
raise ValidationError("pending Atlas batch submission requires a job_id")
self._finish_in_flight_call(
state,
operation="submit",
outcome="accepted",
http_status=http_status,
)
details = _safe_provider_call(provider_call) if provider_call is not None else None
self._merge_provider_call_details(state["provider_calls"][-1], details)
state["job_id"] = job_id
state["status"] = status
state["endpoint"] = endpoint
state["updated_at"] = utc_now()
state["last_http_status"] = http_status
state["last_response"] = _safe_response(response)
if artifacts:
state["artifacts"].extend(redact(artifacts))
return self._save_unlocked(state)
def reconcile_submission_failure(
self,
*,
observed_status: int | None,
endpoint: str,
destination_path: str,
error_name: str,
failure: dict[str, Any],
provider_call: dict[str, Any],
) -> dict[str, Any]:
"""Persist a submit failure without conflating body loss with acceptance."""
if observed_status == 200:
call = dict(provider_call)
call["operation_indeterminate"] = False
call["partial"] = False
return self.reconcile_submission(
job_id=None,
status="completed",
endpoint=endpoint,
http_status=observed_status,
response={
"status": "completed",
"delivery": "synchronous",
"artifact_materialization": {
"status": "failed",
"destination_path": destination_path,
"partial_output": False,
"error": (
f"provider response artifact materialization ended with {error_name}"
),
"failure": failure,
},
},
provider_call=call,
)
if observed_status in ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES:
return self.mark_submission_rejected(
http_status=observed_status,
response=failure,
provider_call=provider_call,
)
provider_acceptance = "unknown"
if observed_status is not None and 200 <= observed_status < 300:
provider_acceptance = "accepted"
elif observed_status is not None and observed_status >= 300:
provider_acceptance = "not-proven"
return self.mark_submission_indeterminate(
f"Atlas submit ended with {error_name}",
provider_call=provider_call,
response=failure,
provider_acceptance=provider_acceptance,
)
def recover_synchronous_submission(
self,
*,
endpoint: str,
input_sha256: str,
parameters: dict[str, Any],
response: dict[str, Any],
artifacts: list[dict[str, Any]],
) -> dict[str, Any]:
"""Reconcile a sync 200 proven durable before submission state was updated."""
with self.operation_lock():
state = self._load_unlocked()
if state["status"] not in {"submitting", "submission-indeterminate"}:
raise ValidationError("Atlas batch submission is not awaiting synchronous recovery")
if state["job_id"] is not None or state["endpoint"] != endpoint:
raise ValidationError("Atlas synchronous recovery does not match submission state")
if state["request"]["input_sha256"] != input_sha256 or state["request"][
"parameters"
] != redact(parameters):
raise ValidationError("Atlas synchronous recovery request does not match state")
call = state["provider_calls"][-1]
previous_outcome = call.get("outcome")
migrated_accepted = (
previous_outcome == "accepted"
and call.get("http_status") == 200
and call.get("operation_indeterminate") is False
and call.get("certainty_normalized_from_legacy") is True
)
if call.get("operation") != "submit" or (
previous_outcome not in {"in-flight", "indeterminate"} and not migrated_accepted
):
raise ValidationError(
"Atlas synchronous recovery requires the unresolved submit call"
)
recovered_at = utc_now()
call["previous_outcome"] = previous_outcome
call["outcome"] = "accepted"
call["operation_indeterminate"] = False
call["http_status"] = 200
call.setdefault("finished_at", recovered_at)
call["recovered_at"] = recovered_at
call["recovered_from_artifact_provenance"] = True
state["status"] = "completed"
state["updated_at"] = recovered_at
state["last_http_status"] = 200
state["last_response"] = _safe_response(response)
state["artifacts"] = redact(artifacts)
state.pop("manual_reconciliation", None)
return self._save_unlocked(state)
def _mark_submission_indeterminate_unlocked(
self,
state: dict[str, Any],
reason: str,
*,
provider_call: dict[str, Any] | None = None,
response: dict[str, Any] | None = None,
provider_acceptance: str = "unknown",
) -> dict[str, Any]:
if provider_acceptance not in {"unknown", "accepted", "not-proven"}:
raise ValidationError("Atlas submission acceptance reconciliation is invalid")
details = _safe_provider_call(provider_call) if provider_call is not None else None
http_status = details.get("http_status") if details is not None else None
if state["status"] == "submitting":
call_outcome = "indeterminate"
if provider_acceptance == "accepted":
call_outcome = "accepted"
elif details is not None and details.get("outcome") == "error":
call_outcome = "error"
self._finish_in_flight_call(
state,
operation="submit",
outcome=call_outcome,
http_status=http_status,
)
self._merge_provider_call_details(state["provider_calls"][-1], details)
elif state["status"] != "submission-indeterminate":
raise ValidationError("Atlas batch submission is not indeterminate")
safe_reason = redact(reason)
state["status"] = "submission-indeterminate"
state["updated_at"] = utc_now()
state["last_http_status"] = http_status
state["last_response"] = {
"status": "submission-indeterminate",
"reason": safe_reason,
}
if response is not None:
state["last_response"]["error"] = _safe_response(response)
state["manual_reconciliation"] = {
"endpoint": state["endpoint"],
"input_sha256": state["request"]["input_sha256"],
"provider_acceptance": provider_acceptance,
"possible_remote_acceptance": True,
"do_not_resubmit": True,
"guidance": _submission_reconciliation_guidance(provider_acceptance),
}
return self._save_unlocked(state)
def mark_submission_indeterminate(
self,
reason: str,
*,
provider_call: dict[str, Any] | None = None,
response: dict[str, Any] | None = None,
provider_acceptance: str = "unknown",
) -> dict[str, Any]:
with self.operation_lock():
return self._mark_submission_indeterminate_unlocked(
self._load_unlocked(),
reason,
provider_call=provider_call,
response=response,
provider_acceptance=provider_acceptance,
)
def mark_submission_rejected(
self,
*,
http_status: int,
response: dict[str, Any],
provider_call: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Persist a definitive, retry-safe HTTP rejection of the POST."""
with self.operation_lock():
if http_status not in ATLAS_DEFINITIVE_SUBMISSION_REJECTION_STATUSES:
raise ValidationError(
"Atlas submission rejection status is not a definitive retry-safe rejection"
)
state = self._load_unlocked()
if state["status"] != "submitting":
raise ValidationError("Atlas batch submission is not awaiting rejection")
self._finish_in_flight_call(
state,
operation="submit",
outcome="rejected",
http_status=http_status,
)
details = _safe_provider_call(provider_call) if provider_call is not None else None
self._merge_provider_call_details(state["provider_calls"][-1], details)
rejected_at = state["provider_calls"][-1]["finished_at"]
rejection = {
**response,
"status": "submission-rejected",
"provider_acceptance": "rejected",
"retry_safe": True,
"rejected_at": rejected_at,
}
retry_after = response.get("retry_after")
if retry_after is not None:
if (
isinstance(retry_after, bool)
or not isinstance(retry_after, (int, float))
or not math.isfinite(retry_after)
or retry_after < 0
):
raise ValidationError("Atlas submission retry delay is invalid")
rejected_time = _utc_datetime(rejected_at, "submission rejection time")
try:
retry_deadline = rejected_time + timedelta(seconds=float(retry_after))
except OverflowError:
retry_deadline = datetime.max.replace(tzinfo=timezone.utc)
rejection["retry_not_before"] = _utc_string(retry_deadline)
safe_rejection = _safe_response(rejection)
state["provider_calls"][-1]["error"] = safe_rejection
state["status"] = "submission-rejected"
state["updated_at"] = rejected_at
state["last_http_status"] = http_status
state["last_response"] = safe_rejection
return self._save_unlocked(state)
def retry_rejected_submission(
self,
*,
endpoint: str,
input_sha256: str,
parameters: dict[str, Any],
invoked_at: str,
) -> dict[str, Any]:
"""Re-arm only the exact request that received a definitive rejection."""
with self.operation_lock():
state = self._load_unlocked()
if state["status"] != "submission-rejected":
raise ValidationError("Atlas batch submission is not retryable")
if state["request"]["input_sha256"] != input_sha256 or state["request"][
"parameters"
] != redact(parameters):
raise ValidationError(
"retry request does not match the rejected Atlas batch submission"
)
rejection = state["provider_calls"][-1]["error"]
rejected_at = _utc_datetime(rejection["rejected_at"], "submission rejection time")
if _utc_datetime(invoked_at, "retry invocation time") <= rejected_at:
raise ValidationError(
"Atlas batch retry must be invoked after the rejection is observed"
)
attempt_started_at = utc_now()
attempt_started = _utc_datetime(attempt_started_at, "retry attempt start time")
retry_not_before = rejection.get("retry_not_before")
if retry_not_before is not None:
deadline = _utc_datetime(retry_not_before, "retry-not-before time")
if attempt_started < deadline:
raise APIError(
status=429,
kind="rate-limit",
message="Atlas batch retry is not yet permitted; honor Retry-After",
retry_after=(deadline - attempt_started).total_seconds(),
)
state["status"] = "submitting"
state["endpoint"] = endpoint
state["updated_at"] = attempt_started_at
state["last_http_status"] = None
state["last_response"] = {"status": "submitting", "retry": True}
state["provider_calls"].append(
{
"endpoint": endpoint,
"operation": "submit",
"http_status": None,
"started_at": attempt_started_at,
"outcome": "in-flight",
"retry_after_rejection": True,
}
)
return self._save_unlocked(state)
def adopt(self, *, job_id: str, endpoint: str) -> dict[str, Any]:
job_id = validate_job_id(job_id)
now = utc_now()
return self.initialize(
job_id=job_id,
status="pending",
submitted_at=now,
endpoint=endpoint,
input_sha256=input_digest({"adopted_job_id": job_id}),
parameters={"adopted": True},
http_status=None,
response={"job_id": job_id, "status": "pending", "adopted": True},
provider_call={
"endpoint": endpoint,
"operation": "adopt",
"timestamp": now,
},
)
@staticmethod
def cancellation_retry_delay(state: dict[str, Any]) -> float | None:
"""Return seconds until an idempotent DELETE retry, or None when resolved."""
if state.get("status") != "cancellation-requested":
return None
calls = state.get("provider_calls")
if not isinstance(calls, list) or not calls:
return None
latest_cancel = next(
(call for call in reversed(calls) if call.get("operation") == "cancel"),
None,
)
if not isinstance(latest_cancel, dict) or latest_cancel.get("outcome") not in {
"error",
"indeterminate",
}:
return None
retry_after = latest_cancel.get("retry_after")
if retry_after is None:
return 0.0
if (
isinstance(retry_after, bool)
or not isinstance(retry_after, (int, float))
or not math.isfinite(retry_after)
or retry_after < 0
):
return None
finished_at = _utc_datetime(latest_cancel.get("finished_at"), "cancellation failure time")
try:
retry_not_before = finished_at + timedelta(seconds=float(retry_after))
except OverflowError:
retry_not_before = datetime.max.replace(tzinfo=timezone.utc)
return max(
0.0,
(retry_not_before - datetime.now(timezone.utc)).total_seconds(),
)
@staticmethod
def cancellation_retry_allowed(state: dict[str, Any]) -> bool:
return AtlasBatchStore.cancellation_retry_delay(state) == 0
def begin_cancellation(self, *, provider_call: dict[str, Any]) -> dict[str, Any]:
with self.operation_lock():
state = self._load_unlocked()
if state["status"] in ATLAS_BATCH_TERMINAL_STATUSES:
return state
if state["status"] == "cancellation-requested" and not self.cancellation_retry_allowed(
state
):
return state
if (
state["status"] not in {"pending", "cancellation-requested"}
or state["job_id"] is None
):
raise ValidationError("Atlas batch state cannot be cancelled")
call = _safe_provider_call(provider_call)
call["http_status"] = None
call["started_at"] = call.pop("timestamp", utc_now())
call["outcome"] = "in-flight"
state["provider_calls"].append(call)
state["status"] = "cancellation-requested"
state["updated_at"] = utc_now()
state["last_http_status"] = None
state["last_response"] = {
"job_id": state["job_id"],
"status": "cancellation-requested",
"intent_persisted": True,
}
return self._save_unlocked(state)
def _mark_cancellation_indeterminate_unlocked(self, state: dict[str, Any]) -> dict[str, Any]:
self._finish_in_flight_call(
state,
operation="cancel",
outcome="indeterminate",
http_status=None,
)
state["status"] = "cancellation-requested"
state["updated_at"] = utc_now()
state["last_http_status"] = None
state["last_response"] = {
"job_id": state["job_id"],
"status": "cancellation-requested",
"provider_acceptance": "unknown",
"retry_safe": "DELETE is idempotent",
}
return self._save_unlocked(state)
def reconcile_cancellation(
self,
*,
http_status: int | None,
accepted: bool,
provider_call: dict[str, Any] | None = None,
) -> dict[str, Any]:
with self.operation_lock():
state = self._load_unlocked()
details = _safe_provider_call(provider_call) if provider_call is not None else None
outcome = "accepted"
if not accepted:
outcome = (
details.get("outcome")
if details is not None and details.get("outcome") in {"error", "indeterminate"}
else "indeterminate"
)
self._finish_in_flight_call(
state,
operation="cancel",
outcome=outcome,
http_status=http_status,
)
call = state["provider_calls"][-1]
self._merge_provider_call_details(call, details)
if state["status"] not in ATLAS_BATCH_TERMINAL_STATUSES:
state["status"] = "cancellation-requested"
state["last_http_status"] = http_status
state["last_response"] = {
"job_id": state["job_id"],
"status": "cancellation-requested",
"provider_acceptance": (
"confirmed"
if accepted
else ("rejected" if outcome == "error" else "unknown")
),
"retry_safe": not accepted,
}
if not accepted and call.get("retry_after") is not None:
state["last_response"]["retry_after"] = call["retry_after"]
state["updated_at"] = utc_now()
return self._save_unlocked(state)
@staticmethod
def _merged_status(current: str, observed: str) -> tuple[str, bool]:
if current in ATLAS_BATCH_TERMINAL_STATUSES:
return current, observed != current
if current == "cancellation-requested" and observed == "pending":
return current, True
allowed = {
"pending": {
"pending",
"cancellation-requested",
*ATLAS_BATCH_TERMINAL_STATUSES,
},
"cancellation-requested": {
"cancellation-requested",
*ATLAS_BATCH_TERMINAL_STATUSES,
},
}
if observed not in allowed.get(current, set()):
raise ValidationError(f"invalid Atlas batch state transition: {current} -> {observed}")
return observed, False
def update(
self,
*,
response: dict[str, Any],
http_status: int | None,
provider_call: dict[str, Any] | None,
status: str | None = None,
artifacts: list[dict[str, Any]] | None = None,
) -> dict[str, Any]:
with self.operation_lock():
state = self._load_unlocked()
response_job_id = response.get("job_id")
if response_job_id is not None and response_job_id != state["job_id"]:
raise ValidationError("Atlas batch response job_id does not match state")
observed_status = status or response.get("status")
if observed_status not in ATLAS_BATCH_STATUSES:
raise ValidationError("Atlas batch update status is invalid")
resolved_status, ignored = self._merged_status(state["status"], observed_status)
if provider_call is not None:
call = _safe_provider_call(provider_call)
call["observed_status"] = observed_status
call["state_status_before"] = state["status"]
call["state_status_after"] = resolved_status
if ignored:
call["transition_ignored"] = "monotonic-state-preserved"
state["provider_calls"].append(call)
state["status"] = resolved_status
state["updated_at"] = utc_now()
if not (resolved_status in ATLAS_BATCH_TERMINAL_STATUSES and ignored):
state["last_http_status"] = http_status
state["last_response"] = _safe_response(response)
if artifacts:
state["artifacts"].extend(redact(artifacts))
return self._save_unlocked(state)
def record_provider_call(
self,
*,
provider_call: dict[str, Any],
response: dict[str, Any],
) -> dict[str, Any]:
"""Append one completed attempt without changing the durable job status."""
with self.operation_lock():
state = self._load_unlocked()
call = _safe_provider_call(provider_call)
state["provider_calls"].append(call)
state["updated_at"] = utc_now()
state["last_http_status"] = call.get("http_status")
state["last_response"] = _safe_response(
{
"status": state["status"],
"provider_call_result": response,
}
)
return self._save_unlocked(state)
def resolve_job_id(self, requested: str | None, *, state: dict[str, Any] | None = None) -> str:
state = self.load() if state is None else state
self.validate(state)
state_job_id = state["job_id"]
if state_job_id is None:
if state["status"] == "submission-indeterminate":
raise ValidationError(
"Atlas batch submission is indeterminate and has no durable job_id; "
"do not resubmit automatically, reconcile it manually"
)
if state["status"] == "submission-rejected":
raise ValidationError(
"Atlas batch submission was rejected and has no job_id; "
"remediate the provider error and retry batch-submit"
)
raise ValidationError("synchronous Atlas batch state has no resumable job_id")
if requested is not None and validate_job_id(requested) != state_job_id:
raise ValidationError("--job-id does not match the durable Atlas batch state")
return state_job_id
SHA-256: f52b72c3f36a1f18b427dfad1054f3e77b8d55af3d99d3aa7bf32523dbfa7174