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

375 lines
18 KiB
Python

#!/usr/bin/env python3
"""Durable, restart-safe assignment state for the Hermes execution pool.
The store owns every transition of a claimed Kanban run through the pool. It
deliberately keeps two invariants that the coordinator depends on for recovery:
* ``lease_failed`` is a *retryable* state, not a terminal one. A row reaches it
when the pool has released the ordinal but has not yet confirmed the outcome
with Kanban, so the coordinator must keep retrying it.
* Only ``finalized`` and ``stale`` are terminal, and only terminal rows are ever
garbage-collected. A row is therefore never removed before the pool holds
authoritative evidence about what happened to its run.
"""
from __future__ import annotations
import json
import sqlite3
import threading
import time
from pathlib import Path
from typing import Any
from execution_pool_protocol import ProtocolError, canonical_json, payload_digest
LIVE_STATES = ("assigned", "running", "result")
TERMINAL_STATES = ("finalized", "stale")
LEASE_FAILED = "lease_failed"
SCHEMA = """
CREATE TABLE IF NOT EXISTS assignments (
board TEXT NOT NULL, task_id TEXT NOT NULL, run_id TEXT NOT NULL,
worker_ordinal INTEGER NOT NULL CHECK(worker_ordinal BETWEEN 0 AND 2),
attempt INTEGER NOT NULL, assignment_digest TEXT NOT NULL,
payload_json TEXT NOT NULL, state TEXT NOT NULL,
lease_until REAL NOT NULL DEFAULT 0, last_heartbeat REAL NOT NULL DEFAULT 0,
result_digest TEXT, result_json TEXT, created_at REAL NOT NULL,
updated_at REAL NOT NULL, PRIMARY KEY(board, task_id, run_id)
);
CREATE UNIQUE INDEX IF NOT EXISTS one_live_assignment_per_worker
ON assignments(worker_ordinal) WHERE state IN ('assigned','running','result');
CREATE TABLE IF NOT EXISTS deliveries (
delivery_id TEXT PRIMARY KEY, kind TEXT NOT NULL, digest TEXT NOT NULL,
received_at REAL NOT NULL
);
"""
class PoolStore:
"""Coordinator-owned durable assignments; never stores provider secrets."""
def __init__(self, path: Path, lease_seconds: int = 90):
self.path = path
self.lease_seconds = max(60, min(int(lease_seconds), 600))
self._lock = threading.RLock()
path.parent.mkdir(parents=True, exist_ok=True)
self._initialize()
def _connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(self.path, timeout=10, isolation_level=None)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA journal_mode=WAL")
connection.execute("PRAGMA synchronous=FULL")
connection.execute("PRAGMA busy_timeout=10000")
return connection
def _initialize(self) -> None:
with self._connect() as connection:
connection.executescript(SCHEMA)
@staticmethod
def _record(row: sqlite3.Row | None) -> dict[str, Any] | None:
if row is None:
return None
value = dict(row)
value["payload"] = json.loads(value.pop("payload_json"))
if value.get("result_json"):
value["result"] = json.loads(value["result_json"])
return value
def _select(self, sql: str) -> list[dict[str, Any]]:
with self._connect() as connection:
rows = connection.execute(sql).fetchall()
return [self._record(row) or {} for row in rows]
def add(self, binding: dict[str, Any], payload: dict[str, Any]) -> bool:
"""Create exactly one assignment for a claimed run and free ordinal."""
now = time.time()
digest = payload_digest(payload)
values = (
binding["board"], binding["task_id"], binding["run_id"],
binding["worker_ordinal"], binding["attempt"], digest,
canonical_json(payload).decode(), "assigned", now, now,
)
with self._lock, self._connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
connection.execute(
"""INSERT INTO assignments
(board,task_id,run_id,worker_ordinal,attempt,assignment_digest,
payload_json,state,created_at,updated_at)
VALUES (?,?,?,?,?,?,?,?,?,?)""",
values,
)
connection.commit()
return True
except sqlite3.IntegrityError as error:
connection.rollback()
existing = connection.execute(
"""SELECT assignment_digest,worker_ordinal,attempt FROM assignments
WHERE board=? AND task_id=? AND run_id=?""",
values[:3],
).fetchone()
if existing and tuple(existing) == (
digest, binding["worker_ordinal"], binding["attempt"]
):
return False
if existing:
raise ProtocolError("conflicting duplicate assignment") from error
raise ProtocolError("worker ordinal already has a live assignment") from error
def available_ordinals(self) -> list[int]:
with self._connect() as connection:
rows = connection.execute(
"SELECT worker_ordinal FROM assignments WHERE state IN ('assigned','running','result')"
).fetchall()
occupied = {int(row[0]) for row in rows}
return [ordinal for ordinal in range(3) if ordinal not in occupied]
def active_assignments(self) -> list[dict[str, Any]]:
return self._select(
"SELECT * FROM assignments WHERE state IN ('assigned','running','result')"
)
def terminal_record(self, board: str, task_id: str, run_id: str) -> dict[str, Any] | None:
"""Return one final pool record without exposing nonterminal work."""
with self._connect() as connection:
row = connection.execute(
"SELECT * FROM assignments WHERE board=? AND task_id=? AND run_id=? "
"AND state IN ('finalized','stale')", (board, task_id, run_id)
).fetchone()
return self._record(row)
def record_publication_lease_failure(self, binding: dict[str, Any], receipt: dict[str, Any], marker: str) -> bool:
"""Persist coordinator-owned OOM evidence without replacing a worker result."""
source = receipt.get("source") if isinstance(receipt, dict) else None
digest = receipt.get("result_digest") if isinstance(receipt, dict) else None
expected_marker = f"[hermes-publication-retry-fence:{binding['run_id']}:{digest}]"
if (
not isinstance(source, dict) or source.get("board") != binding["board"]
or source.get("task_id") != binding["task_id"]
or not isinstance(source.get("worker_ordinal"), int)
or source["worker_ordinal"] != binding["worker_ordinal"]
or not isinstance(digest, str) or len(digest) != 64 or marker != expected_marker
):
raise ProtocolError("publication lease evidence is invalid")
result = {
"structured": {"status": "blocked", "summary": "Publication retry worker lease expired."},
"capacity_failure": True,
"scm_submission": None,
"scm_resume": receipt,
"publication_retry_lease_fence": marker,
}
encoded = canonical_json(result).decode()
with self._lock, self._connect() as connection:
row = connection.execute(
"SELECT payload_json,result_json,state FROM assignments WHERE board=? AND task_id=? AND run_id=? "
"AND worker_ordinal=? AND attempt=?",
tuple(binding[name] for name in ("board", "task_id", "run_id", "worker_ordinal", "attempt")),
).fetchone()
if row is None or row["state"] not in {LEASE_FAILED, "finalized"}:
return False
try:
payload = json.loads(row["payload_json"])
except (TypeError, json.JSONDecodeError) as error:
raise ProtocolError("publication lease assignment is malformed") from error
if payload.get("scm_resume") != receipt:
return False
if row["result_json"]:
return row["result_json"] == encoded
changed = connection.execute(
"UPDATE assignments SET result_digest=?,result_json=?,updated_at=? WHERE board=? AND task_id=? "
"AND run_id=? AND worker_ordinal=? AND attempt=? AND result_json IS NULL",
(payload_digest(result), encoded, time.time(), *(binding[name] for name in ("board", "task_id", "run_id", "worker_ordinal", "attempt"))),
).rowcount
return bool(changed)
def known_runs(self) -> set[tuple[str, str, str]]:
"""Every run identity this store already owns a row for, in any state.
Adoption must consult this rather than only the live states: re-adding a
run that still has a retrying or not-yet-collected row would collide with
``PRIMARY KEY(board, task_id, run_id)`` under a different payload digest.
"""
with self._connect() as connection:
rows = connection.execute(
"SELECT board,task_id,run_id FROM assignments"
).fetchall()
return {(str(row[0]), str(row[1]), str(row[2])) for row in rows}
def offer(self, ordinal: int) -> dict[str, Any] | None:
"""Return the ordinal's durable assignment, preserving restart identity."""
now = time.time()
with self._lock, self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
row = connection.execute(
"""SELECT * FROM assignments WHERE worker_ordinal=?
AND state IN ('assigned','running') ORDER BY created_at LIMIT 1""",
(ordinal,),
).fetchone()
if row is not None:
connection.execute(
"""UPDATE assignments SET state='running',lease_until=?,last_heartbeat=?,updated_at=?
WHERE board=? AND task_id=? AND run_id=?""",
(now + self.lease_seconds, now, now, row["board"], row["task_id"], row["run_id"]),
)
connection.commit()
return self._record(row)
def _matching(self, connection: sqlite3.Connection, envelope: dict[str, Any]) -> sqlite3.Row:
row = connection.execute(
"SELECT * FROM assignments WHERE board=? AND task_id=? AND run_id=?",
(envelope["board"], envelope["task_id"], envelope["run_id"]),
).fetchone()
if row is None:
raise ProtocolError("assignment is unknown or stale")
if int(row["worker_ordinal"]) != int(envelope["worker_ordinal"]):
raise ProtocolError("worker ordinal does not own this assignment")
if int(row["attempt"]) != int(envelope["attempt"]):
raise ProtocolError("assignment attempt is stale")
return row
def heartbeat(self, envelope: dict[str, Any]) -> tuple[bool, bool]:
now = time.time()
delivery_digest = payload_digest(
{
name: envelope[name]
for name in (
"kind", "board", "task_id", "run_id", "worker_ordinal",
"attempt", "payload_digest",
)
}
)
with self._lock, self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
duplicate = connection.execute(
"SELECT digest FROM deliveries WHERE delivery_id=?",
(envelope["delivery_id"],),
).fetchone()
row = self._matching(connection, envelope)
if row["state"] not in {"assigned", "running"}:
raise ProtocolError("assignment is no longer running")
if row["state"] == "running" and float(row["lease_until"]) < now:
raise ProtocolError("assignment lease expired")
if duplicate and duplicate[0] != delivery_digest:
raise ProtocolError("delivery identifier was reused")
if not duplicate:
connection.execute(
"INSERT INTO deliveries VALUES (?,?,?,?)",
(envelope["delivery_id"], "heartbeat", delivery_digest, now),
)
connection.execute(
"UPDATE assignments SET state='running',lease_until=?,last_heartbeat=?,updated_at=? WHERE board=? AND task_id=? AND run_id=?",
(now + self.lease_seconds, now, now, envelope["board"], envelope["task_id"], envelope["run_id"]),
)
connection.commit()
return True, bool(duplicate)
def accept_result(self, envelope: dict[str, Any]) -> tuple[dict[str, Any], bool]:
now = time.time()
with self._lock, self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
row = self._matching(connection, envelope)
digest = envelope["payload_digest"]
if row["result_digest"]:
if row["result_digest"] != digest:
raise ProtocolError("conflicting result for completed delivery")
connection.rollback()
return self._record(row) or {}, True
if row["state"] not in {"assigned", "running"}:
raise ProtocolError("assignment cannot accept a result")
if row["state"] == "running" and float(row["lease_until"]) < now:
raise ProtocolError("assignment lease expired")
connection.execute(
"""UPDATE assignments SET state='result',result_digest=?,result_json=?,updated_at=?
WHERE board=? AND task_id=? AND run_id=?""",
(digest, canonical_json(envelope["payload"]).decode(), now,
envelope["board"], envelope["task_id"], envelope["run_id"]),
)
connection.commit()
row = connection.execute(
"SELECT * FROM assignments WHERE board=? AND task_id=? AND run_id=?",
(envelope["board"], envelope["task_id"], envelope["run_id"]),
).fetchone()
return self._record(row) or {}, False
def pending_results(self) -> list[dict[str, Any]]:
return self._select(
"SELECT * FROM assignments WHERE state='result' ORDER BY updated_at"
)
def failed_leases(self) -> list[dict[str, Any]]:
"""Rows whose ordinal is released but whose Kanban outcome is unconfirmed.
The coordinator drains this on every maintenance pass, so a Kanban write
that failed at expiry time is retried instead of stranding the run.
"""
return self._select(
"SELECT * FROM assignments WHERE state='lease_failed' ORDER BY updated_at"
)
def expire_leases(
self, *, now: float | None = None, max_attempts: int = 3
) -> list[dict[str, Any]]:
"""Fence expired attempts and re-offer or terminally release their ordinals."""
current = time.time() if now is None else float(now)
maximum = max(1, min(int(max_attempts), 10))
changed: list[dict[str, Any]] = []
with self._lock, self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
rows = connection.execute(
"""SELECT * FROM assignments WHERE state='running'
AND lease_until > 0 AND lease_until < ? ORDER BY updated_at""",
(current,),
).fetchall()
for row in rows:
if int(row["attempt"]) >= maximum:
state, attempt = LEASE_FAILED, int(row["attempt"])
else:
state, attempt = "assigned", int(row["attempt"]) + 1
connection.execute(
"""UPDATE assignments SET state=?,attempt=?,lease_until=0,
last_heartbeat=0,updated_at=? WHERE board=? AND task_id=? AND run_id=?
AND attempt=? AND state='running'""",
(
state, attempt, current, row["board"], row["task_id"],
row["run_id"], row["attempt"],
),
)
updated = connection.execute(
"SELECT * FROM assignments WHERE board=? AND task_id=? AND run_id=?",
(row["board"], row["task_id"], row["run_id"]),
).fetchone()
if updated is not None:
changed.append(self._record(updated) or {})
connection.commit()
return changed
def finalize(self, binding: dict[str, Any], state: str) -> bool:
"""Apply one terminal state to the exact run, ordinal, and attempt."""
if state not in TERMINAL_STATES:
raise ProtocolError("invalid terminal assignment state")
with self._lock, self._connect() as connection:
cursor = connection.execute(
"UPDATE assignments SET state=?,updated_at=? WHERE board=? AND task_id=? AND run_id=? AND worker_ordinal=? AND attempt=?",
(state, time.time(), *(binding[name] for name in ("board", "task_id", "run_id", "worker_ordinal", "attempt"))),
)
return bool(cursor.rowcount)
def garbage_collect(self, retention_seconds: int) -> int:
"""Remove aged rows, and only ones with authoritative terminal evidence.
``lease_failed`` is excluded on purpose: dropping a row whose Kanban
outcome was never confirmed would lose the pool's only record that the
run still needs to be surfaced.
"""
cutoff = time.time() - max(3600, retention_seconds)
with self._lock, self._connect() as connection:
cursor = connection.execute(
"DELETE FROM assignments WHERE state IN ('finalized','stale') AND updated_at < ?",
(cutoff,),
)
connection.execute("DELETE FROM deliveries WHERE received_at < ?", (cutoff,))
return int(cursor.rowcount)