"""PostgreSQL implementation of the durable, user-owned session repository.""" from __future__ import annotations import hashlib import json import re import uuid from collections.abc import Iterator, Sequence 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 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\d+)_(?P[a-z0-9_]+)\.sql$") _SAFE_CTE_NAME = re.compile(r"[A-Za-z0-9_-]+\Z") _ARTIFACT_KEYS = { "question", "schema_linking", "evidence", "evidence_receipts", "sql_final", "memory_proposals", "validation_report", "retrieval_pack", "cte_tests", "cte_plan", "cte_plan_doc", } _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 list(self) -> list[SessionSnapshot]: with self._transaction() as connection: ids = [row[0] for row in connection.execute( text("SELECT id FROM thoth_sessions.sessions ORDER BY created_at DESC") ).all()] # psycopg2 materializes PostgreSQL UUID columns as ``uuid.UUID`` objects, # while the repository boundary intentionally accepts canonical UUIDv4 text. return [self.get(str(session_id)) for session_id in ids] 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 ensure_interaction_language( self, session_id: str, workspace_language: str ) -> SessionSnapshot: from tht.session.store import pin_interaction_language self._require_uuid4(session_id) with self._transaction() as connection: self._lock_session(connection, session_id) row = connection.execute( text("SELECT manifest FROM thoth_sessions.sessions WHERE id = :id FOR UPDATE"), {"id": session_id}, ).mappings().first() if row is None: raise SessionError(f"Sessione non trovata: {session_id}") manifest = SessionManifest.model_validate(row["manifest"]) if pin_interaction_language(manifest, workspace_language): connection.execute( text("UPDATE thoth_sessions.sessions " "SET manifest = CAST(:manifest AS jsonb), updated_at = pg_catalog.now() " "WHERE id = :id"), {"id": session_id, "manifest": json.dumps(manifest.model_dump(by_alias=True, mode="json"))}, ) return self.get(session_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 delete_artifact(self, session_id: str, key: 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("DELETE FROM thoth_sessions.session_artifacts WHERE session_id = :id AND artifact_key = :key"), {"id": session_id, "key": key}, ) 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) from tht.phase import current_phase manifest = SessionManifest.model_validate(connection.execute( text("SELECT manifest FROM thoth_sessions.sessions WHERE id = :id"), {"id": session_id} ).scalar_one()) ledger = [ 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() ] phase = current_phase(SessionSnapshot(manifest=manifest, decisions=ledger)) 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=phase, **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 finalize(self, manifest: SessionManifest, artifacts: dict[str, str]) -> SessionSnapshot: """Atomically publish DWH-verified artifacts and the final manifest.""" self._require_uuid4(manifest.id) for key in artifacts: self._require_artifact_key(key) with self._transaction() as connection: self._lock_session(connection, manifest.id) self._require_session(connection, manifest.id) for key, content in artifacts.items(): 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": manifest.id, "key": key, "content": content}, ) 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 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"]