375 lines
18 KiB
Python
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)
|