548 lines
23 KiB
Python
548 lines
23 KiB
Python
"""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",
|
|
"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()]
|
|
return [self.get(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 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"]
|