feat(harness): persist server sessions in postgres
This commit is contained in:
@@ -0,0 +1,483 @@
|
||||
"""PostgreSQL implementation of the durable, user-owned session repository."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from importlib.resources import files
|
||||
from importlib.resources.abc import Traversable
|
||||
from pathlib import Path
|
||||
from typing import Iterator, Sequence
|
||||
|
||||
from sqlalchemy import Engine, create_engine, text
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from tht.decisions import DecisionInput, DecisionRecord
|
||||
from tht.session.models import PrincipalContext, SessionManifest, SessionSnapshot
|
||||
from tht.session.store import SessionError
|
||||
|
||||
MIGRATIONS_DIR = files("tht").joinpath("migrations", "sessions")
|
||||
_MIGRATION_NAME = re.compile(r"^(?P<version>\d+)_(?P<name>[a-z0-9_]+)\.sql$")
|
||||
_SAFE_CTE_NAME = re.compile(r"[A-Za-z0-9_-]+\Z")
|
||||
_ARTIFACT_KEYS = {
|
||||
"question",
|
||||
"schema_linking",
|
||||
"evidence",
|
||||
"sql_final",
|
||||
"validation_report",
|
||||
"retrieval_pack",
|
||||
"cte_tests",
|
||||
}
|
||||
_LOCK_KEY = 8_420_613_069_444_020_731
|
||||
_RUNTIME_ROLE = "thoth_sessions_runtime"
|
||||
|
||||
|
||||
class MigrationError(RuntimeError):
|
||||
"""Raised when session schema migration discovery or application is unsafe."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Migration:
|
||||
version: str
|
||||
name: str
|
||||
path: Traversable
|
||||
checksum: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MigrationStatus:
|
||||
applied: tuple[Migration, ...]
|
||||
pending: tuple[Migration, ...]
|
||||
drifted: tuple[Migration, ...]
|
||||
|
||||
|
||||
def _migration_source(directory: Traversable | Path | str) -> Traversable:
|
||||
return Path(directory) if isinstance(directory, (str, Path)) else directory
|
||||
|
||||
|
||||
def _discover(directory: Traversable | Path | str = MIGRATIONS_DIR) -> tuple[Migration, ...]:
|
||||
source = _migration_source(directory)
|
||||
seen_versions: set[int] = set()
|
||||
parsed = []
|
||||
for path in (item for item in source.iterdir() if item.name.endswith(".sql")):
|
||||
match = _MIGRATION_NAME.fullmatch(path.name)
|
||||
if match is None:
|
||||
raise MigrationError(f"Invalid migration filename: {path.name}")
|
||||
version = match.group("version")
|
||||
numeric_version = int(version)
|
||||
if numeric_version in seen_versions:
|
||||
raise MigrationError(f"Duplicate migration version: {numeric_version}")
|
||||
seen_versions.add(numeric_version)
|
||||
parsed.append((numeric_version, version, match.group("name"), path))
|
||||
if not parsed:
|
||||
raise MigrationError(f"No migrations found in {source}")
|
||||
return tuple(
|
||||
Migration(version, name, path, hashlib.sha256(path.read_bytes()).hexdigest())
|
||||
for _, version, name, path in sorted(parsed)
|
||||
)
|
||||
|
||||
|
||||
def _engine(database_url: str) -> Engine:
|
||||
url = make_url(database_url)
|
||||
if not url.drivername.startswith("postgresql"):
|
||||
raise ValueError("Session repository requires a direct PostgreSQL URL")
|
||||
return create_engine(url)
|
||||
|
||||
|
||||
def _applied(connection) -> dict[str, str]:
|
||||
exists = connection.execute(
|
||||
text("SELECT pg_catalog.to_regclass('public.tht_session_migrations')")
|
||||
).scalar()
|
||||
if exists is None:
|
||||
return {}
|
||||
return dict(
|
||||
connection.execute(text("SELECT version, checksum FROM public.tht_session_migrations")).all()
|
||||
)
|
||||
|
||||
|
||||
def _reject_unknown_versions(
|
||||
migrations: tuple[Migration, ...], applied_checksums: dict[str, str]
|
||||
) -> None:
|
||||
local_versions = {migration.version for migration in migrations}
|
||||
unknown = sorted(
|
||||
set(applied_checksums) - local_versions,
|
||||
key=lambda value: (0, int(value)) if value.isdigit() else (1, value),
|
||||
)
|
||||
if unknown:
|
||||
raise MigrationError("Database migration versions absent from local manifest: " + ", ".join(unknown))
|
||||
|
||||
|
||||
def migration_status(
|
||||
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
|
||||
) -> MigrationStatus:
|
||||
migrations = _discover(migrations_dir)
|
||||
engine = _engine(database_url)
|
||||
try:
|
||||
with engine.connect() as connection:
|
||||
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
|
||||
applied_checksums = _applied(connection)
|
||||
finally:
|
||||
engine.dispose()
|
||||
_reject_unknown_versions(migrations, applied_checksums)
|
||||
return MigrationStatus(
|
||||
applied=tuple(item for item in migrations if applied_checksums.get(item.version) == item.checksum),
|
||||
pending=tuple(item for item in migrations if item.version not in applied_checksums),
|
||||
drifted=tuple(
|
||||
item
|
||||
for item in migrations
|
||||
if item.version in applied_checksums and applied_checksums[item.version] != item.checksum
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def migrate(
|
||||
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
|
||||
) -> MigrationStatus:
|
||||
migrations = _discover(migrations_dir)
|
||||
engine = _engine(database_url)
|
||||
current: Migration | None = None
|
||||
try:
|
||||
with engine.begin() as connection:
|
||||
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
|
||||
connection.execute(text("SELECT pg_catalog.pg_advisory_xact_lock(:key)"), {"key": _LOCK_KEY})
|
||||
connection.exec_driver_sql(
|
||||
"""CREATE TABLE IF NOT EXISTS public.tht_session_migrations (
|
||||
version text PRIMARY KEY,
|
||||
name text NOT NULL,
|
||||
checksum text NOT NULL,
|
||||
applied_at timestamptz NOT NULL DEFAULT pg_catalog.now()
|
||||
)"""
|
||||
)
|
||||
connection.exec_driver_sql("REVOKE ALL ON public.tht_session_migrations FROM PUBLIC")
|
||||
applied_checksums = _applied(connection)
|
||||
_reject_unknown_versions(migrations, applied_checksums)
|
||||
drifted = [
|
||||
item
|
||||
for item in migrations
|
||||
if item.version in applied_checksums and applied_checksums[item.version] != item.checksum
|
||||
]
|
||||
if drifted:
|
||||
raise MigrationError("Migration checksum drift: " + ", ".join(item.version for item in drifted))
|
||||
for current in migrations:
|
||||
if current.version in applied_checksums:
|
||||
continue
|
||||
connection.exec_driver_sql(current.path.read_text())
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO public.tht_session_migrations (version, name, checksum) "
|
||||
"VALUES (:version, :name, :checksum)"
|
||||
),
|
||||
{"version": current.version, "name": current.name, "checksum": current.checksum},
|
||||
)
|
||||
except MigrationError:
|
||||
raise
|
||||
except SQLAlchemyError as exc:
|
||||
filename = current.path.name if current is not None else "migration setup"
|
||||
raise MigrationError(f"Failed to apply {filename}: {type(exc).__name__}") from exc
|
||||
finally:
|
||||
engine.dispose()
|
||||
return migration_status(database_url, migrations_dir)
|
||||
|
||||
|
||||
class PostgresSessionRepository:
|
||||
"""Session data stored in private PostgreSQL tables under transaction-local RLS context."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
principal: PrincipalContext,
|
||||
*,
|
||||
engine: Engine | None = None,
|
||||
runtime_role: str = _RUNTIME_ROLE,
|
||||
):
|
||||
self.principal = principal
|
||||
self._engine = engine or _engine(database_url)
|
||||
self._owns_engine = engine is None
|
||||
self._runtime_role = runtime_role
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config, principal: PrincipalContext) -> "PostgresSessionRepository":
|
||||
query = {"sslmode": config.sslmode}
|
||||
if config.sslrootcert is not None:
|
||||
query["sslrootcert"] = str(config.sslrootcert)
|
||||
url = URL.create(
|
||||
"postgresql+psycopg2",
|
||||
username=config.user,
|
||||
password=config.password,
|
||||
host=config.host,
|
||||
port=config.port,
|
||||
database=config.database,
|
||||
query=query,
|
||||
)
|
||||
return cls(url.render_as_string(hide_password=False), principal)
|
||||
|
||||
def close(self) -> None:
|
||||
if self._owns_engine:
|
||||
self._engine.dispose()
|
||||
self._owns_engine = False
|
||||
|
||||
def create(self, manifest: SessionManifest) -> SessionSnapshot:
|
||||
self._require_uuid4(manifest.id)
|
||||
with self._transaction() as connection:
|
||||
principal_id = self._upsert_principal(connection)
|
||||
try:
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.sessions (id, principal_id, manifest) "
|
||||
"VALUES (:id, :principal_id, CAST(:manifest AS jsonb))"
|
||||
),
|
||||
{
|
||||
"id": manifest.id,
|
||||
"principal_id": principal_id,
|
||||
"manifest": json.dumps(manifest.model_dump(by_alias=True, mode="json")),
|
||||
},
|
||||
)
|
||||
except SQLAlchemyError as exc:
|
||||
if "unique" in str(exc).lower():
|
||||
raise SessionError(f"Sessione esiste gia': {manifest.id}") from exc
|
||||
raise
|
||||
return self.get(manifest.id)
|
||||
|
||||
def get(self, session_id: str) -> SessionSnapshot:
|
||||
self._require_uuid4(session_id)
|
||||
with self._transaction() as connection:
|
||||
row = connection.execute(
|
||||
text("SELECT manifest FROM thoth_sessions.sessions WHERE id = :id"), {"id": session_id}
|
||||
).mappings().first()
|
||||
if row is None:
|
||||
raise SessionError(f"Sessione non trovata: {session_id}")
|
||||
artifacts = dict(
|
||||
connection.execute(
|
||||
text(
|
||||
"SELECT artifact_key, content FROM thoth_sessions.session_artifacts "
|
||||
"WHERE session_id = :id ORDER BY artifact_key"
|
||||
),
|
||||
{"id": session_id},
|
||||
).all()
|
||||
)
|
||||
decisions = [
|
||||
DecisionRecord.model_validate(dict(item))
|
||||
for item in connection.execute(
|
||||
text(
|
||||
"SELECT seq, ts, phase, type, subject, detail, rationale, retracts "
|
||||
"FROM thoth_sessions.review_decisions WHERE session_id = :id ORDER BY seq"
|
||||
),
|
||||
{"id": session_id},
|
||||
).mappings()
|
||||
]
|
||||
return SessionSnapshot(
|
||||
principal=self.principal,
|
||||
manifest=SessionManifest.model_validate(row["manifest"]),
|
||||
artifacts=artifacts,
|
||||
decisions=decisions,
|
||||
)
|
||||
|
||||
def save_manifest(self, manifest: SessionManifest) -> SessionSnapshot:
|
||||
self._require_uuid4(manifest.id)
|
||||
with self._transaction() as connection:
|
||||
result = connection.execute(
|
||||
text(
|
||||
"UPDATE thoth_sessions.sessions "
|
||||
"SET manifest = CAST(:manifest AS jsonb), updated_at = pg_catalog.now() "
|
||||
"WHERE id = :id"
|
||||
),
|
||||
{"id": manifest.id, "manifest": json.dumps(manifest.model_dump(by_alias=True, mode="json"))},
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
raise SessionError(f"Sessione non trovata: {manifest.id}")
|
||||
return self.get(manifest.id)
|
||||
|
||||
def read_artifact(self, session_id: str, key: str) -> str | None:
|
||||
self._require_uuid4(session_id)
|
||||
self._require_artifact_key(key)
|
||||
with self._transaction() as connection:
|
||||
self._require_session(connection, session_id)
|
||||
return connection.execute(
|
||||
text(
|
||||
"SELECT content FROM thoth_sessions.session_artifacts "
|
||||
"WHERE session_id = :id AND artifact_key = :key"
|
||||
),
|
||||
{"id": session_id, "key": key},
|
||||
).scalar()
|
||||
|
||||
def write_artifact(self, session_id: str, key: str, content: str) -> None:
|
||||
self._require_uuid4(session_id)
|
||||
self._require_artifact_key(key)
|
||||
with self._transaction() as connection:
|
||||
self._lock_session(connection, session_id)
|
||||
self._require_session(connection, session_id)
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.session_artifacts (session_id, artifact_key, content) "
|
||||
"VALUES (:id, :key, :content) "
|
||||
"ON CONFLICT (session_id, artifact_key) DO UPDATE "
|
||||
"SET content = EXCLUDED.content, updated_at = pg_catalog.now()"
|
||||
),
|
||||
{"id": session_id, "key": key, "content": content},
|
||||
)
|
||||
|
||||
def append_decisions(
|
||||
self, session_id: str, decisions: Sequence[DecisionInput | dict]
|
||||
) -> list[DecisionRecord]:
|
||||
self._require_uuid4(session_id)
|
||||
inputs = [DecisionInput.model_validate(item) for item in decisions]
|
||||
if not inputs:
|
||||
return []
|
||||
with self._transaction() as connection:
|
||||
self._lock_session(connection, session_id)
|
||||
self._require_session(connection, session_id)
|
||||
next_seq = connection.execute(
|
||||
text(
|
||||
"SELECT COALESCE(MAX(seq), 0) + 1 FROM thoth_sessions.review_decisions "
|
||||
"WHERE session_id = :id"
|
||||
),
|
||||
{"id": session_id},
|
||||
).scalar_one()
|
||||
now = datetime.now(UTC)
|
||||
records = [
|
||||
DecisionRecord(seq=next_seq + offset, ts=now, phase=None, **item.model_dump())
|
||||
for offset, item in enumerate(inputs)
|
||||
]
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.review_decisions "
|
||||
"(session_id, seq, ts, phase, type, subject, detail, rationale, retracts) "
|
||||
"VALUES (:session_id, :seq, :ts, :phase, :type, :subject, :detail, :rationale, :retracts)"
|
||||
),
|
||||
[
|
||||
{"session_id": session_id, **record.model_dump(mode="python")}
|
||||
for record in records
|
||||
],
|
||||
)
|
||||
return records
|
||||
|
||||
def get_preferences(self) -> dict:
|
||||
with self._transaction() as connection:
|
||||
principal_id = self._upsert_principal(connection)
|
||||
preferences = connection.execute(
|
||||
text(
|
||||
"SELECT preferences FROM thoth_sessions.principal_preferences "
|
||||
"WHERE principal_id = :principal_id"
|
||||
),
|
||||
{"principal_id": principal_id},
|
||||
).scalar()
|
||||
return preferences or {}
|
||||
|
||||
def set_preferences(self, preferences: dict) -> None:
|
||||
if not isinstance(preferences, dict):
|
||||
raise TypeError("preferences must be a dictionary")
|
||||
with self._transaction() as connection:
|
||||
principal_id = self._upsert_principal(connection)
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.principal_preferences (principal_id, preferences) "
|
||||
"VALUES (:principal_id, CAST(:preferences AS jsonb)) "
|
||||
"ON CONFLICT (principal_id) DO UPDATE SET preferences = EXCLUDED.preferences, "
|
||||
"updated_at = pg_catalog.now()"
|
||||
),
|
||||
{"principal_id": principal_id, "preferences": json.dumps(preferences)},
|
||||
)
|
||||
|
||||
def delete(self, session_id: str) -> None:
|
||||
self._require_uuid4(session_id)
|
||||
with self._transaction() as connection:
|
||||
self._lock_session(connection, session_id)
|
||||
owner = connection.execute(
|
||||
text(
|
||||
"SELECT p.issuer, p.subject FROM thoth_sessions.sessions s "
|
||||
"JOIN thoth_sessions.principals p ON p.id = s.principal_id WHERE s.id = :id"
|
||||
),
|
||||
{"id": session_id},
|
||||
).mappings().first()
|
||||
if owner is None:
|
||||
raise SessionError(f"Sessione non trovata: {session_id}")
|
||||
connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.audit_log "
|
||||
"(action, session_id, actor_issuer, actor_subject, owner_issuer, owner_subject) "
|
||||
"VALUES ('session_deleted', :id, :actor_issuer, :actor_subject, :owner_issuer, :owner_subject)"
|
||||
),
|
||||
{
|
||||
"id": session_id,
|
||||
"actor_issuer": self.principal.issuer,
|
||||
"actor_subject": self.principal.subject,
|
||||
"owner_issuer": owner["issuer"],
|
||||
"owner_subject": owner["subject"],
|
||||
},
|
||||
)
|
||||
connection.execute(text("DELETE FROM thoth_sessions.sessions WHERE id = :id"), {"id": session_id})
|
||||
|
||||
@contextmanager
|
||||
def _transaction(self) -> Iterator:
|
||||
try:
|
||||
with self._engine.begin() as connection:
|
||||
connection.exec_driver_sql("SET LOCAL search_path = thoth_sessions, pg_catalog, pg_temp")
|
||||
connection.exec_driver_sql(f"SET LOCAL ROLE {self._runtime_role}")
|
||||
connection.execute(
|
||||
text("SELECT pg_catalog.set_config('thoth_sessions.actor_issuer', :value, true)"),
|
||||
{"value": self.principal.issuer},
|
||||
)
|
||||
connection.execute(
|
||||
text("SELECT pg_catalog.set_config('thoth_sessions.actor_subject', :value, true)"),
|
||||
{"value": self.principal.subject},
|
||||
)
|
||||
connection.execute(
|
||||
text("SELECT pg_catalog.set_config('thoth_sessions.is_admin', :value, true)"),
|
||||
{"value": "true" if self.principal.is_admin else "false"},
|
||||
)
|
||||
yield connection
|
||||
except SessionError:
|
||||
raise
|
||||
except SQLAlchemyError as exc:
|
||||
raise SessionError("Session storage unavailable") from exc
|
||||
|
||||
def _upsert_principal(self, connection) -> int:
|
||||
return connection.execute(
|
||||
text(
|
||||
"INSERT INTO thoth_sessions.principals (issuer, subject, display_name) "
|
||||
"VALUES (:issuer, :subject, :display_name) "
|
||||
"ON CONFLICT (issuer, subject) DO UPDATE SET display_name = EXCLUDED.display_name, "
|
||||
"updated_at = pg_catalog.now() RETURNING id"
|
||||
),
|
||||
self.principal.model_dump(include={"issuer", "subject", "display_name"}),
|
||||
).scalar_one()
|
||||
|
||||
@staticmethod
|
||||
def _lock_session(connection, session_id: str) -> None:
|
||||
connection.execute(
|
||||
text("SELECT pg_catalog.pg_advisory_xact_lock(pg_catalog.hashtextextended(:id, 0))"),
|
||||
{"id": session_id},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _require_session(connection, session_id: str) -> None:
|
||||
if connection.execute(
|
||||
text("SELECT 1 FROM thoth_sessions.sessions WHERE id = :id"), {"id": session_id}
|
||||
).scalar() is None:
|
||||
raise SessionError(f"Sessione non trovata: {session_id}")
|
||||
|
||||
@staticmethod
|
||||
def _require_uuid4(session_id: str) -> None:
|
||||
try:
|
||||
parsed = uuid.UUID(session_id, version=4)
|
||||
except ValueError as exc:
|
||||
raise SessionError(f"ID sessione non UUIDv4: {session_id}") from exc
|
||||
if str(parsed) != session_id or parsed.version != 4:
|
||||
raise SessionError(f"ID sessione non UUIDv4: {session_id}")
|
||||
|
||||
@staticmethod
|
||||
def _require_artifact_key(key: str) -> None:
|
||||
if key in _ARTIFACT_KEYS:
|
||||
return
|
||||
if key.startswith("cte_sql:") and _SAFE_CTE_NAME.fullmatch(key.removeprefix("cte_sql:")):
|
||||
return
|
||||
raise ValueError(f"Unsupported artifact key: {key}")
|
||||
|
||||
|
||||
__all__ = ["MigrationError", "MigrationStatus", "PostgresSessionRepository", "migrate", "migration_status"]
|
||||
Reference in New Issue
Block a user