refactor(gateway/hosted_room_*): _control_event/_append_control_event fold claim+lost paths, _text() helper for payload fields, inline single-use fetchone locals

This commit is contained in:
Teknium
2026-09-02 23:06:15 -07:00
parent 380df1bfcb
commit 6edd2c4be9
3 changed files with 63 additions and 87 deletions
+10 -21
View File
@@ -102,19 +102,12 @@ _GENERATION_TRANSITIONS = {
class DriverStateError(ValueError): """Base class for invalid or conflicting driver-state operations."""
class DriverValidationError(DriverStateError): """Raised when an identifier, clock, TTL, or payload is invalid."""
class RoomUnavailableError(DriverStateError): """Raised when the hosted room does not exist or was disbanded."""
class LeaseHeldError(DriverStateError): """Raised when another unexpired driver generation owns the room."""
class StaleLeaseError(DriverStateError): """Raised when a lease generation can no longer mutate room state."""
class TaskConflictError(DriverStateError): """Raised when an idempotency key is reused for different task state."""
class StaleTaskError(DriverStateError): """Raised when an obsolete task attempt or cancellation tries to commit."""
class InvalidTaskTransitionError(DriverStateError): """Raised when a requested task transition is not allowed."""
@@ -290,9 +283,9 @@ def _connect(db_path: DbPath) -> sqlite3.Connection:
def ready(conn: sqlite3.Connection) -> bool:
existing.append(_schema_objects_exist(conn))
return existing[0] and _task_schema_supports_current_statuses(conn)
def initialize(conn: sqlite3.Connection) -> None:
(_migrate_task_status_constraint if existing[0] else _initialize_schema)(conn)
conn = connect(db_path, db_label="state.db (hosted_room_driver)", ready=ready, initialize=initialize)
conn = connect(
db_path, db_label="state.db (hosted_room_driver)", ready=ready,
initialize=lambda conn: (_migrate_task_status_constraint if existing[0] else _initialize_schema)(conn))
if existing[0]:
try:
_validate_schema(conn)
@@ -599,10 +592,9 @@ def release_lease(db_path: DbPath, lease: DriverLease, *, clock: Clock) -> dict[
return {"lease": _lease_from_row(row), "idempotent": True}
if float(row["expires_at"]) <= now:
raise StaleLeaseError("driver lease expired before release")
running = conn.execute(
if conn.execute(
"SELECT 1 FROM hosted_room_driver_tasks WHERE room_id=? AND status='running' LIMIT 1", (lease.room_id,)
).fetchone()
if running is not None:
).fetchone() is not None:
raise InvalidTaskTransitionError("cannot release a room lease while tasks are running")
conn.execute("""UPDATE hosted_room_driver_leases SET expires_at=?, updated_at=?, released_at=?
WHERE room_id=? AND lease_generation=?""",
@@ -625,10 +617,9 @@ def admit_task(db_path: DbPath, identity: TaskIdentity, *, payload: Any, clock:
if existing["payload_digest"] != payload_digest or existing["payload_json"] != payload_json:
raise TaskConflictError("task_id is already bound to a different payload")
return _task_from_row(existing, idempotent=True)
turn = conn.execute(
if conn.execute(
"SELECT * FROM hosted_room_driver_tasks WHERE room_id=? AND thread_id=? AND turn_id=?",
(identity.room_id, identity.thread_id, identity.turn_id)).fetchone()
if turn is not None:
(identity.room_id, identity.thread_id, identity.turn_id)).fetchone() is not None:
raise TaskConflictError("thread_id and turn_id are already bound to a task")
conn.execute("""INSERT INTO hosted_room_driver_tasks (
room_id, task_id, thread_id, turn_id, source_event_seq, payload_json, payload_digest,
@@ -653,11 +644,10 @@ def start_task(
_require_cancel_generation(row, expected_cancel_generation)
if row["status"] != "queued":
raise InvalidTaskTransitionError(f"cannot start task in state '{row['status']}'")
unresolved = conn.execute(
if conn.execute(
f"""SELECT task_id, status FROM hosted_room_driver_tasks
WHERE room_id=? AND status IN ('running', 'indeterminate', 'stopping') {_TASK_ORDER} LIMIT 1""",
(identity.room_id,)).fetchone()
if unresolved is not None:
(identity.room_id,)).fetchone() is not None:
raise InvalidTaskTransitionError("room recovery must resolve the prior task before starting new work")
next_queued = conn.execute(
f"SELECT task_id FROM hosted_room_driver_tasks WHERE room_id=? AND status='queued' {_TASK_ORDER} LIMIT 1",
@@ -895,8 +885,7 @@ def prune_published_terminal_tasks(
][:MAX_TASK_PRUNE_BATCH]
if not candidates:
return 0
placeholders = ",".join("?" for _ in candidates)
deleted = conn.execute(
f"DELETE FROM hosted_room_driver_tasks WHERE room_id=? AND task_id IN ({placeholders})",
f"DELETE FROM hosted_room_driver_tasks WHERE room_id=? AND task_id IN ({','.join('?' * len(candidates))})",
(room_id, *candidates))
return max(0, int(deleted.rowcount))
+34 -40
View File
@@ -79,6 +79,10 @@ class PolicySnapshot:
_event_from_room_row = hosted_rooms._event_from_row
def _text(mapping: Mapping[str, Any], key: str) -> str:
return str(mapping.get(key) or "")
def _require_room(conn: sqlite3.Connection, room_id: str) -> None:
if conn.execute("SELECT 1 FROM hosted_rooms WHERE room_id=?", (room_id,)).fetchone() is None:
raise hosted_rooms.RoomNotFoundError("hosted room not found")
@@ -145,7 +149,7 @@ class HostedRoomPolicyCheckpoint:
settled_seq_by_message: dict[str, int] = {}
for row in conn.execute("""SELECT seq, payload_json FROM hosted_room_events
WHERE room_id=? AND seq<=? AND kind='turn.settled' ORDER BY seq""", (room_id, through_seq)):
message_event_id = str(json.loads(row["payload_json"]).get("message_event_id") or "")
message_event_id = _text(json.loads(row["payload_json"]), "message_event_id")
if message_event_id:
settled_seq_by_message[message_event_id] = int(row["seq"])
rows = conn.execute(
@@ -156,7 +160,7 @@ class HostedRoomPolicyCheckpoint:
if row["kind"] == "message.member" and row["event_id"] not in settled_seq_by_message:
continue
event = _event_from_room_row(row)
thread_id = str(event["payload"].get("thread_id") or "")
thread_id = _text(event["payload"], "thread_id")
if thread_id:
self._store_transcript_event(
conn, event=event, thread_id=thread_id, settled_seq=settled_seq_by_message.get(str(row["event_id"]))
@@ -183,8 +187,7 @@ class HostedRoomPolicyCheckpoint:
def _apply_user_message(
self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None:
room_id = str(event["room_id"])
thread_id = str(payload.get("thread_id") or "")
event_id = str(event.get("event_id") or "")
thread_id, event_id = _text(payload, "thread_id"), _text(event, "event_id")
if not thread_id or not event_id:
return
conn.execute("""INSERT INTO hosted_room_policy_threads(
@@ -200,27 +203,23 @@ class HostedRoomPolicyCheckpoint:
def _apply_discussion_event(
self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None:
"""Index member messages and terminal turn outcomes of a known discussion."""
room_id = str(event["room_id"])
seq = int(event["seq"])
kind = str(event.get("kind") or "")
thread_id = str(payload.get("thread_id") or "")
discussion_event_id = str(payload.get("discussion_event_id") or "")
source = conn.execute(
room_id, seq, kind = str(event["room_id"]), int(event["seq"]), _text(event, "kind")
thread_id, discussion_event_id = _text(payload, "thread_id"), _text(payload, "discussion_event_id")
if conn.execute(
"SELECT 1 FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=? LIMIT 1",
(room_id, discussion_event_id)).fetchone()
if source is None:
(room_id, discussion_event_id)).fetchone() is None:
return
self._store_active_event(conn, event=event, thread_id=thread_id, discussion_event_id=discussion_event_id)
if kind not in _TERMINAL_KINDS:
return
task_id = str(payload.get("task_id") or "")
execution_generation = (int(payload.get("execution_generation") or 0) if kind == "turn.deferred" else 0)
task_id = _text(payload, "task_id")
execution_generation = int(payload.get("execution_generation") or 0) if kind == "turn.deferred" else 0
if task_id:
conn.execute("""INSERT OR IGNORE INTO hosted_room_policy_publications(
room_id, task_id, kind, execution_generation, seq
) VALUES (?, ?, ?, ?, ?)""",
(room_id, task_id, kind, execution_generation, seq))
member_id = str(payload.get("member_id") or "")
member_id = _text(payload, "member_id")
seen_through_seq = int(payload.get("seen_through_seq") or 0)
if kind == "turn.settled" and payload.get("message_event_id"):
committed = _settled_message(conn, room_id, discussion_event_id, payload["message_event_id"])
@@ -237,10 +236,8 @@ class HostedRoomPolicyCheckpoint:
def _apply_room_activity(
self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None:
room_id = str(event["room_id"])
thread_id = str(payload.get("thread_id") or "")
discussion_event_id = str(payload.get("discussion_event_id") or "")
conn.execute(_DELETE_ACTIVE_EVENTS_SQL, (room_id, discussion_event_id))
room_id, thread_id = str(event["room_id"]), _text(payload, "thread_id")
conn.execute(_DELETE_ACTIVE_EVENTS_SQL, (room_id, _text(payload, "discussion_event_id")))
conn.execute("DELETE FROM hosted_room_policy_threads WHERE room_id=? AND thread_id=?", (room_id, thread_id))
def _apply_stop_requested(
@@ -255,11 +252,10 @@ class HostedRoomPolicyCheckpoint:
"room.stop_requested": _apply_stop_requested}
def _apply_event(self, conn: sqlite3.Connection, event: Mapping[str, Any]) -> None:
handler = self._APPLY_BY_KIND.get(str(event.get("kind") or ""))
if handler is None:
return
payload = event.get("payload")
handler(self, conn, event, payload if isinstance(payload, Mapping) else {})
handler = self._APPLY_BY_KIND.get(_text(event, "kind"))
if handler is not None:
payload = event.get("payload")
handler(self, conn, event, payload if isinstance(payload, Mapping) else {})
def _ensure_cursor_and_transcript(self, conn: sqlite3.Connection, room_id: str) -> int:
"""Create the room cursor if absent, backfill the transcript once, return through_seq."""
@@ -267,8 +263,9 @@ class HostedRoomPolicyCheckpoint:
conn.execute("""INSERT OR IGNORE INTO hosted_room_policy_cursors(
room_id, through_seq, stopped_through_seq, updated_at
) VALUES (?, 0, 0, 0)""", (room_id,))
row = conn.execute("SELECT through_seq FROM hosted_room_policy_cursors WHERE room_id=?", (room_id,)).fetchone()
cursor = int(row["through_seq"])
cursor = int(
conn.execute("SELECT through_seq FROM hosted_room_policy_cursors WHERE room_id=?", (room_id,)).fetchone()[
"through_seq"])
transcript_state = conn.execute(
"SELECT schema_version FROM hosted_room_policy_transcript_state WHERE room_id=?", (room_id,)).fetchone()
if transcript_state is None or int(transcript_state["schema_version"]) < _TRANSCRIPT_SCHEMA_VERSION:
@@ -331,14 +328,14 @@ class HostedRoomPolicyCheckpoint:
def publication_exists(self, *, room_id: str, task_id: str, status: str, execution_generation: int) -> bool:
"""Return whether one exact driver outcome is already in the room log."""
if status == "deferred":
sql = """SELECT 1 FROM hosted_room_policy_publications
WHERE room_id=? AND task_id=? AND kind=? AND execution_generation=?"""
params = (room_id, task_id, f"turn.{status}", execution_generation)
else:
sql = """SELECT 1 FROM hosted_room_policy_publications
WHERE room_id=? AND task_id=? AND kind IN ('turn.settled', 'turn.failed', 'turn.cancelled')"""
params = (room_id, task_id)
sql, params = (
("""SELECT 1 FROM hosted_room_policy_publications
WHERE room_id=? AND task_id=? AND kind=? AND execution_generation=?""",
(room_id, task_id, f"turn.{status}", execution_generation))
if status == "deferred" else
("""SELECT 1 FROM hosted_room_policy_publications
WHERE room_id=? AND task_id=? AND kind IN ('turn.settled', 'turn.failed', 'turn.cancelled')""",
(room_id, task_id)))
with self._connect() as conn:
return conn.execute(sql, params).fetchone() is not None
@@ -348,9 +345,7 @@ class HostedRoomPolicyCheckpoint:
source = conn.execute(
"SELECT discussion_event_id, thread_id FROM hosted_room_policy_events WHERE room_id=? AND seq=?",
(room_id, source_event_seq)).fetchone()
if source is None:
return []
return self._discussion_events(
return [] if source is None else self._discussion_events(
conn, room_id=room_id, thread_id=str(source["thread_id"]),
discussion_event_id=str(source["discussion_event_id"]),
bound_error="task policy projection exceeded its bound")
@@ -358,9 +353,8 @@ class HostedRoomPolicyCheckpoint:
def compact_completed(self, *, room_id: str) -> None:
"""Drop any completed projections left by an interrupted sync."""
with self._connect() as conn:
completed = conn.execute(
for row in conn.execute(
"SELECT discussion_event_id FROM hosted_room_policy_threads WHERE room_id=? AND completed=1", (room_id,)
).fetchall()
for row in completed:
).fetchall():
conn.execute(_DELETE_ACTIVE_EVENTS_SQL, (room_id, str(row["discussion_event_id"])))
conn.execute("DELETE FROM hosted_room_policy_threads WHERE room_id=? AND completed=1", (room_id,))
+19 -26
View File
@@ -30,9 +30,7 @@ _SELECT_REPLICA = "SELECT * FROM hosted_room_replicas WHERE room_id=?"
class ReplicaError(HostedRoomError): """Base class for invalid or conflicting replica operations."""
class ReplicaGapError(ReplicaError): """A page does not start at the replica's next expected sequence."""
class ReplicaEpochRegressionError(ReplicaError): """A page or demotion carries an older authority epoch than stored."""
@@ -73,9 +71,16 @@ def _replica_transaction(db_path: DbPath) -> Iterator[sqlite3.Connection]:
_positive_int = partial(bounded_int, error=ReplicaError, low=1)
def _control_event_json(payload: dict[str, Any]) -> tuple[str, str]:
"""Canonical (actor_json, payload_json) for a system authority-control event."""
return _actor_json(_SYSTEM_ACTOR), _payload_json(payload)
def _control_event(kind: str, epoch: int, payload: dict[str, Any]) -> tuple[str, str, str, str]:
"""(event_id, kind, actor_json, payload_json) of the system ``authority.<kind>`` control event for ``epoch``."""
return f"system:authority-{kind}:{epoch}", f"authority.{kind}", _actor_json(_SYSTEM_ACTOR), _payload_json(payload)
def _append_control_event(
conn: sqlite3.Connection, room_id: str, seq: int, epoch: int, event: tuple[str, str, str, str], now: float
) -> None:
event_id, kind, actor_json, payload_json = event
conn.execute(_INSERT_ROOM_EVENT, (room_id, seq, event_id, kind, actor_json, epoch, payload_json, now))
def _event_bytes(event: dict[str, Any]) -> int:
@@ -88,8 +93,7 @@ def _event_bytes(event: dict[str, Any]) -> int:
def _validate_page(page: Any) -> tuple[list[dict[str, Any]], dict[str, Any]]:
if not isinstance(page, dict):
raise ReplicaError("page must be an object")
events = page.get("events")
authority = page.get("authority")
events, authority = page.get("events"), page.get("authority")
if not isinstance(events, list):
raise ReplicaError("page.events must be a list")
if not isinstance(authority, dict):
@@ -164,13 +168,12 @@ def ingest_page(
size = _event_bytes(event)
if stored_bytes + added_bytes + size > MAX_REPLICA_EVENT_BYTES:
raise ReplicaError("replica event storage exhausted")
actor_json = _actor_json(event["actor"])
payload_json = _payload_json(event["payload"])
conn.execute(
_INSERT_REPLICA_EVENT,
(
room_id, int(event["seq"]), event["event_id"], event["kind"], actor_json,
event.get("authority_epoch"), payload_json, float(event.get("created_at") or now)))
room_id, int(event["seq"]), event["event_id"], event["kind"], _actor_json(event["actor"]),
event.get("authority_epoch"), _payload_json(event["payload"]),
float(event.get("created_at") or now)))
added_bytes += size
new_last = int(new_events[-1]["seq"]) if new_events else last_seq
latest_seq = page.get("latest_seq")
@@ -226,27 +229,21 @@ def promote_replica(
previous_epoch = int(replica["authority_epoch"])
target_epoch = previous_epoch + 1
claim_seq = int(replica["last_seq"]) + 1
claim_event_id = f"system:authority-claimed:{target_epoch}"
claim_actor_json, claim_payload_json = _control_event_json({
claim = _control_event("claimed", target_epoch, {
"previous_gateway_id": previous_gateway, "authority_gateway_id": local_gateway,
"authority_epoch": target_epoch, "promoted_from_replica": True, "reason": reason})
claim_bytes = utf8_len(claim_event_id, "authority.claimed", claim_actor_json, claim_payload_json)
conn.execute("""INSERT INTO hosted_rooms
(room_id, name, members_json, authority_gateway_id, authority_epoch, next_seq, event_bytes,
revision, created_at, updated_at, disbanded_at)
VALUES (?, ?, ?, ?, ?, ?, ?, 1, ?, ?, NULL)""",
(
room_id, replica["name"], replica["members_json"], local_gateway, target_epoch, claim_seq + 1,
int(replica["event_bytes"]) + claim_bytes, now, now))
int(replica["event_bytes"]) + utf8_len(*claim), now, now))
conn.execute(
f"""INSERT INTO hosted_room_events {_EVENT_COLUMNS}
SELECT room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at
FROM hosted_room_replica_events WHERE room_id=?""", (room_id,))
conn.execute(
_INSERT_ROOM_EVENT,
(
room_id, claim_seq, claim_event_id, "authority.claimed", claim_actor_json, target_epoch,
claim_payload_json, now))
_append_control_event(conn, room_id, claim_seq, target_epoch, claim, now)
conn.execute("DELETE FROM hosted_room_replica_events WHERE room_id=?", (room_id,))
conn.execute("DELETE FROM hosted_room_replicas WHERE room_id=?", (room_id,))
return {
@@ -285,14 +282,10 @@ def demote_room(
raise ReplicaEpochRegressionError("observed epoch does not supersede the stored authority")
if current_gateway != local_gateway:
raise ReplicaError("room is not locally authoritative; nothing to demote")
lost_actor_json, lost_payload_json = _control_event_json({
lost = _control_event("lost", observed_epoch, {
"previous_gateway_id": current_gateway, "authority_gateway_id": observed_gateway_id,
"authority_epoch": observed_epoch})
conn.execute(
_INSERT_ROOM_EVENT,
(
room_id, int(row["next_seq"]), f"system:authority-lost:{observed_epoch}", "authority.lost",
lost_actor_json, observed_epoch, lost_payload_json, now))
_append_control_event(conn, room_id, int(row["next_seq"]), observed_epoch, lost, now)
conn.execute("""UPDATE hosted_rooms
SET authority_gateway_id=?, authority_epoch=?, next_seq=next_seq+1, revision=revision+1, updated_at=?
WHERE room_id=?""",