tmux-amq/src/tamq/store.py
tegwick 74b2f27997
Some checks failed
tamq-ci / test (push) Failing after 6s
feat: make sessions terminal neutral
Assistant: codex
Assistant-Model: gpt-5.6-sol
Assistant-Session: 01a03397-4d51-7fd1-8ff2-946eb22ea2bc
2026-08-24 20:10:08 +02:00

201 lines
8.6 KiB
Python

from __future__ import annotations
import json
import sqlite3
import time
from pathlib import Path
from typing import Iterable
from uuid import uuid4
SCHEMA_VERSION = 2
class Store:
def __init__(self, path: Path):
path.parent.mkdir(parents=True, exist_ok=True)
self.db = sqlite3.connect(path, timeout=5)
self.db.row_factory = sqlite3.Row
self.db.execute("PRAGMA journal_mode=WAL")
self.db.execute("PRAGMA foreign_keys=ON")
self.db.execute("PRAGMA busy_timeout=5000")
self.db.executescript("""
CREATE TABLE IF NOT EXISTS metadata (key TEXT PRIMARY KEY, value TEXT NOT NULL);
CREATE TABLE IF NOT EXISTS messages (
message_id TEXT PRIMARY KEY,
sender_repo TEXT NOT NULL,
target_repo TEXT NOT NULL,
body TEXT NOT NULL,
created_at REAL NOT NULL,
state TEXT NOT NULL,
endpoint_id TEXT,
provenance TEXT,
injected_at REAL,
acknowledged_at REAL
);
CREATE INDEX IF NOT EXISTS messages_target_state ON messages(target_repo, state);
CREATE TABLE IF NOT EXISTS endpoints (
endpoint_id TEXT PRIMARY KEY,
pid INTEGER NOT NULL,
session TEXT NOT NULL,
repos TEXT NOT NULL,
delivery_mode TEXT NOT NULL DEFAULT 'manual',
connected_at REAL NOT NULL,
disconnected_at REAL
);
CREATE TABLE IF NOT EXISTS leases (
message_id TEXT PRIMARY KEY REFERENCES messages(message_id) ON DELETE CASCADE,
lease_id TEXT NOT NULL,
endpoint_id TEXT NOT NULL,
acquired_at REAL NOT NULL,
expires_at REAL NOT NULL
);
""")
endpoint_columns = {
row["name"] for row in self.db.execute("PRAGMA table_info(endpoints)")
}
if "delivery_mode" not in endpoint_columns:
self.db.execute(
"ALTER TABLE endpoints ADD COLUMN delivery_mode TEXT NOT NULL DEFAULT 'manual'"
)
self.db.execute(
"INSERT INTO metadata(key,value) VALUES('schema_version',?) "
"ON CONFLICT(key) DO UPDATE SET value=excluded.value",
(str(SCHEMA_VERSION),),
)
self.db.commit()
def close(self) -> None:
self.db.close()
def register_endpoint(
self,
endpoint_id: str,
pid: int,
session: str,
repos: list[str],
delivery_mode: str = "manual",
) -> None:
import json
if delivery_mode not in {"manual", "pane"}:
raise ValueError(f"invalid delivery mode: {delivery_mode}")
self.db.execute(
"UPDATE endpoints SET disconnected_at=strftime('%s','now') "
"WHERE pid=? AND session=? AND endpoint_id<>? AND disconnected_at IS NULL",
(pid, session, endpoint_id),
)
self.db.execute(
"INSERT INTO endpoints(endpoint_id,pid,session,repos,delivery_mode,connected_at,disconnected_at) VALUES(?,?,?,?,?,strftime('%s','now'),NULL) "
"ON CONFLICT(endpoint_id) DO UPDATE SET pid=excluded.pid, session=excluded.session, repos=excluded.repos, delivery_mode=excluded.delivery_mode, connected_at=excluded.connected_at, disconnected_at=NULL",
(endpoint_id, pid, session, json.dumps(repos), delivery_mode),
)
self.db.commit()
def disconnect_endpoint(self, endpoint_id: str) -> None:
self.db.execute("UPDATE endpoints SET disconnected_at=strftime('%s','now') WHERE endpoint_id=?", (endpoint_id,))
self.db.commit()
def disconnect_all(self) -> None:
self.db.execute("UPDATE endpoints SET disconnected_at=strftime('%s','now') WHERE disconnected_at IS NULL")
self.db.commit()
def endpoints(self) -> list[sqlite3.Row]:
return list(
self.db.execute(
"SELECT * FROM endpoints WHERE disconnected_at IS NULL ORDER BY endpoint_id"
)
)
def endpoint(self, endpoint_id: str) -> sqlite3.Row | None:
row = self.db.execute("SELECT * FROM endpoints WHERE endpoint_id=? AND disconnected_at IS NULL", (endpoint_id,)).fetchone()
if row is not None:
return row
if endpoint_id.startswith("tmux-amq-"):
visible_pid = endpoint_id.removeprefix("tmux-amq-")
if not visible_pid.isdigit():
return None
pid = int(visible_pid)
return self.db.execute("SELECT * FROM endpoints WHERE pid=? AND disconnected_at IS NULL ORDER BY connected_at DESC LIMIT 1", (pid,)).fetchone()
return None
def history_stats(self) -> tuple[int, float | None]:
size = 0
try:
size = self.db.execute("SELECT page_count * page_size FROM pragma_page_count(), pragma_page_size()").fetchone()[0]
except sqlite3.DatabaseError:
pass
oldest = self.db.execute("SELECT MIN(created_at) FROM messages").fetchone()[0]
return size, oldest
def add(self, sender: str, target: str, body: str, *, endpoint: str | None = None, provenance: str | None = None) -> str:
if len(body.encode("utf-8")) > 8192:
raise ValueError("message body exceeds 8 KiB limit")
message_id = f"m-{uuid4()}"
self.db.execute(
"INSERT INTO messages VALUES(?,?,?,?,?,?,?,?,?,?)",
(message_id, sender, target, body, time.time(), "pending", endpoint, provenance, None, None),
)
self.db.commit()
return message_id
def list(self, target: str | None = None, state: str | None = None) -> list[sqlite3.Row]:
clauses, values = [], []
if target:
clauses.append("target_repo=?"); values.append(target)
if state:
clauses.append("state=?"); values.append(state)
where = f" WHERE {' AND '.join(clauses)}" if clauses else ""
return list(self.db.execute(f"SELECT * FROM messages{where} ORDER BY created_at", values))
def set_state(self, message_id: str, state: str) -> None:
column = {"injected": "injected_at", "acknowledged": "acknowledged_at"}.get(state)
if column:
self.db.execute(f"UPDATE messages SET state=?, {column}=? WHERE message_id=?", (state, time.time(), message_id))
else:
self.db.execute("UPDATE messages SET state=? WHERE message_id=?", (state, message_id))
self.db.commit()
def acknowledge(self, message_id: str) -> bool:
row = self.db.execute("SELECT 1 FROM messages WHERE message_id=?", (message_id,)).fetchone()
if row is None:
return False
self.set_state(message_id, "acknowledged")
return True
def claim(self, message_id: str, endpoint_id: str, ttl: float = 30.0) -> str | None:
now = time.time()
lease_id = f"lease-{uuid4()}"
with self.db:
self.db.execute("DELETE FROM leases WHERE expires_at < ?", (now,))
try:
self.db.execute("INSERT INTO leases VALUES(?,?,?,?,?)", (message_id, lease_id, endpoint_id, now, now + ttl))
except sqlite3.IntegrityError:
return None
return lease_id
def renew(self, message_id: str, lease_id: str, ttl: float = 30.0) -> bool:
with self.db:
result = self.db.execute("UPDATE leases SET expires_at=? WHERE message_id=? AND lease_id=?", (time.time() + ttl, message_id, lease_id))
return result.rowcount == 1
def release(self, message_id: str, lease_id: str, state: str = "injected") -> bool:
with self.db:
result = self.db.execute("DELETE FROM leases WHERE message_id=? AND lease_id=?", (message_id, lease_id))
if result.rowcount:
self.set_state(message_id, state)
return result.rowcount == 1
def export(self, rows: Iterable[sqlite3.Row], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as stream:
for row in rows:
stream.write(json.dumps(dict(row), sort_keys=True) + "\n")
def purge(self, before: float | None = None, max_bytes: int | None = None) -> int:
rows = self.list()
ids = [r["message_id"] for r in rows if before is not None and r["created_at"] < before]
if max_bytes is not None and self.db.execute("SELECT page_count * page_size FROM pragma_page_count(), pragma_page_size()").fetchone()[0] > max_bytes:
ids += [r["message_id"] for r in rows if r["message_id"] not in ids][:max(0, len(rows) // 2)]
if ids:
self.db.executemany("DELETE FROM messages WHERE message_id=?", ((i,) for i in set(ids)))
self.db.commit()
return len(set(ids))