#!/usr/bin/env python3 """Merge one Google session delta into an isolated recovery candidate. The source connection is forced read-only. The target is accepted only when it matches the known isolated recovery-candidate shape. No identifiers or user content are printed. The command defaults to a dry run; ``--apply`` is required to write the candidate. """ from __future__ import annotations import argparse import hashlib import json import os import sys import traceback from collections.abc import Iterable from dataclasses import dataclass from typing import Any import psycopg from psycopg import sql from psycopg.types.json import Jsonb EXPECTED_SOURCE_ROWS = { "sessions": 1, "turns": 0, "case_profiles": 1, "session_states": 1, "alliance_pulses": 1, "pulse_status_events": 2, "self_assessments": 1, "measurement_events": 9, "session_summaries": 1, "model_runs": 2, } EXPECTED_TARGET_BEFORE = {"users": 84, "sessions": 29, "turns": 705} EXPECTED_TARGET_AFTER = { "users": 84, "sessions": 30, "turns": 705, "google_sessions": 30, } COPIED_TABLES = ( ("app", "case_profile"), ("app", "sessions"), ("app", "session_state"), ("audit", "model_run"), ("app", "alliance_pulse"), ("audit", "alliance_pulse_status_event"), ("app", "self_assessment"), ("app", "measurement_instrument"), ("app", "measurement_event"), ("app", "session_summary"), ) class RecoveryGuardError(RuntimeError): """A fail-closed recovery precondition or postcondition failed.""" @dataclass(frozen=True) class DbSettings: host: str port: int dbname: str user: str password: str def kwargs(self, *, read_only: bool) -> dict[str, Any]: kwargs: dict[str, Any] = { "host": self.host, "port": self.port, "dbname": self.dbname, "user": self.user, "password": self.password, "connect_timeout": 10, "application_name": "vignette_isolated_recovery_delta", } if read_only: kwargs["options"] = "-c default_transaction_read_only=on" return kwargs def _required_env(name: str) -> str: value = os.environ.get(name, "") if not value: raise RecoveryGuardError(f"required environment variable is missing: {name}") return value def _settings(prefix: str, expected_port: int) -> DbSettings: port = int(_required_env(f"{prefix}_PORT")) if port != expected_port: raise RecoveryGuardError(f"{prefix} port guard failed") dbname = _required_env(f"{prefix}_DBNAME") if dbname != "vignette": raise RecoveryGuardError(f"{prefix} database-name guard failed") return DbSettings( host=_required_env(f"{prefix}_HOST"), port=port, dbname=dbname, user=_required_env(f"{prefix}_USER"), password=_required_env(f"{prefix}_PASSWORD"), ) def _stable_hash(value: Any) -> str: payload = json.dumps( value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str, ).encode("utf-8") return hashlib.sha256(payload).hexdigest() def _scalar(cur: psycopg.Cursor[Any], query: str, params: Iterable[Any] = ()) -> Any: cur.execute(query, tuple(params)) row = cur.fetchone() if row is None: raise RecoveryGuardError("expected scalar query returned no row") return row[0] def _rows_json( cur: psycopg.Cursor[Any], schema: str, table: str, where_sql: str, params: Iterable[Any], order_sql: str = "", ) -> list[dict[str, Any]]: query = sql.SQL("SELECT to_jsonb(t) FROM {}.{} AS t WHERE ").format( sql.Identifier(schema), sql.Identifier(table) ) + sql.SQL(where_sql) if order_sql: query += sql.SQL(" ORDER BY ") + sql.SQL(order_sql) cur.execute(query, tuple(params)) return [row[0] for row in cur.fetchall()] def _insert_json_record( cur: psycopg.Cursor[Any], schema: str, table: str, record: dict[str, Any] ) -> None: query = sql.SQL( "INSERT INTO {}.{} SELECT (jsonb_populate_record(NULL::{}.{}, %s)).*" ).format( sql.Identifier(schema), sql.Identifier(table), sql.Identifier(schema), sql.Identifier(table), ) cur.execute(query, (Jsonb(record),)) if cur.rowcount != 1: raise RecoveryGuardError(f"insert count mismatch for {schema}.{table}") def _table_columns( cur: psycopg.Cursor[Any], schema: str, table: str ) -> list[tuple[str, str, str, str]]: cur.execute( """ SELECT column_name, udt_schema, udt_name, is_nullable FROM information_schema.columns WHERE table_schema = %s AND table_name = %s ORDER BY column_name """, (schema, table), ) return [tuple(row) for row in cur.fetchall()] def _assert_schema_compatibility( source_cur: psycopg.Cursor[Any], target_cur: psycopg.Cursor[Any] ) -> None: for schema, table in COPIED_TABLES: source_columns = _table_columns(source_cur, schema, table) target_columns = _table_columns(target_cur, schema, table) if not source_columns or source_columns != target_columns: raise RecoveryGuardError(f"schema mismatch for {schema}.{table}") source_constraints = _table_constraints(source_cur, schema, table) target_constraints = _table_constraints(target_cur, schema, table) if source_constraints != target_constraints: raise RecoveryGuardError(f"constraint mismatch for {schema}.{table}") def _table_constraints( cur: psycopg.Cursor[Any], schema: str, table: str ) -> list[tuple[str, str, str]]: cur.execute( r""" SELECT c.contype::text, c.conname, regexp_replace(pg_get_constraintdef(c.oid), '\\s+', ' ', 'g') FROM pg_constraint c JOIN pg_class t ON t.oid = c.conrelid JOIN pg_namespace n ON n.oid = t.relnamespace WHERE n.nspname = %s AND t.relname = %s ORDER BY c.contype, c.conname, pg_get_constraintdef(c.oid) """, (schema, table), ) return [tuple(row) for row in cur.fetchall()] def _target_counts(cur: psycopg.Cursor[Any]) -> dict[str, int]: return { "users": int(_scalar(cur, "SELECT count(*) FROM app.app_user")), "sessions": int(_scalar(cur, "SELECT count(*) FROM app.sessions")), "turns": int(_scalar(cur, "SELECT count(*) FROM app.turns")), } def _topological_measurements(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: pending = {row["measurement_id"]: row for row in rows} ordered: list[dict[str, Any]] = [] emitted: set[Any] = set() while pending: ready = [ row for row in pending.values() if row.get("supersedes_id") is None or row.get("supersedes_id") in emitted or row.get("supersedes_id") not in pending ] if not ready: raise RecoveryGuardError("measurement supersession graph is cyclic") ready.sort(key=lambda row: (str(row.get("created_at", "")), str(row["measurement_id"]))) for row in ready: ordered.append(row) emitted.add(row["measurement_id"]) pending.pop(row["measurement_id"]) return ordered def _source_bundle( source_cur: psycopg.Cursor[Any], target_cur: psycopg.Cursor[Any] ) -> dict[str, Any]: source_sessions = _rows_json( source_cur, "app", "sessions", "learner_id IN (SELECT user_id FROM app.app_user WHERE external_id LIKE 'google:%%')", (), ) target_session_ids = set( str(row[0]) for row in target_cur.execute("SELECT id FROM app.sessions").fetchall() ) if len(source_sessions) != EXPECTED_SOURCE_ROWS["sessions"]: raise RecoveryGuardError("expected exactly one source Google session") session = source_sessions[0] already_present = session["id"] in target_session_ids session_id = session["id"] source_user_id = session["learner_id"] source_external_rows = source_cur.execute( "SELECT external_id FROM app.app_user WHERE user_id = %s AND external_id LIKE 'google:%%'", (source_user_id,), ).fetchall() if len(source_external_rows) != 1: raise RecoveryGuardError("source Google identity lookup was not unique") external_id = source_external_rows[0][0] target_user_rows = target_cur.execute( "SELECT user_id FROM app.app_user WHERE external_id = %s", (external_id,) ).fetchall() if len(target_user_rows) != 1: raise RecoveryGuardError("candidate Google identity lookup was not unique") target_user_id = target_user_rows[0][0] turns = _rows_json(source_cur, "app", "turns", "session_id = %s", (session_id,)) session_states = _rows_json( source_cur, "app", "session_state", "session_id = %s", (session_id,) ) pulses = _rows_json( source_cur, "app", "alliance_pulse", "session_id = %s", (session_id,) ) pulse_ids = [row["pulse_id"] for row in pulses] pulse_status_events = ( _rows_json( source_cur, "audit", "alliance_pulse_status_event", "pulse_id = ANY(%s)", (pulse_ids,), "changed_at, created_at", ) if pulse_ids else [] ) self_assessments = _rows_json( source_cur, "app", "self_assessment", "session_id = %s", (session_id,) ) measurements = _rows_json( source_cur, "app", "measurement_event", "session_id = %s", (session_id,) ) summaries = _rows_json( source_cur, "app", "session_summary", "session_id = %s", (session_id,) ) model_runs = _rows_json( source_cur, "audit", "model_run", "session_id = %s", (session_id,) ) case_profiles = ( _rows_json( source_cur, "app", "case_profile", "case_id = %s", (session["case_id"],) ) if session.get("case_id") is not None else [] ) actual = { "sessions": len(source_sessions), "turns": len(turns), "case_profiles": len(case_profiles), "session_states": len(session_states), "alliance_pulses": len(pulses), "pulse_status_events": len(pulse_status_events), "self_assessments": len(self_assessments), "measurement_events": len(measurements), "session_summaries": len(summaries), "model_runs": len(model_runs), } if actual != EXPECTED_SOURCE_ROWS: raise RecoveryGuardError("source row-count contract changed") if any(row.get("turn_id") is not None for row in measurements + model_runs): raise RecoveryGuardError("zero-turn session unexpectedly references a turn") if any(row.get("evidence_turn_ids") for row in measurements + self_assessments): raise RecoveryGuardError("zero-turn session unexpectedly has turn evidence") model_run_ids = {row["model_run_id"] for row in model_runs} for row in measurements: referenced_model_run = row.get("model_run_id") if referenced_model_run is not None and referenced_model_run not in model_run_ids: raise RecoveryGuardError("measurement references an uncaptured model run") for row in pulse_status_events: actor = row.get("changed_by_uid") if actor is not None and actor != source_user_id: raise RecoveryGuardError("pulse status event has an unexpected actor") source_fingerprint = _stable_hash( { "session": session, "case_profiles": case_profiles, "session_states": session_states, "pulses": pulses, "pulse_status_events": pulse_status_events, "self_assessments": self_assessments, "measurements": measurements, "summaries": summaries, "model_runs": model_runs, } ) return { "session": session, "session_id": session_id, "already_present": already_present, "source_fingerprint": source_fingerprint, "source_user_id": source_user_id, "target_user_id": target_user_id, "case_profiles": case_profiles, "session_states": session_states, "pulses": pulses, "pulse_status_events": pulse_status_events, "self_assessments": self_assessments, "measurements": measurements, "summaries": summaries, "model_runs": model_runs, } def _resolve_target_context( source_cur: psycopg.Cursor[Any], target_cur: psycopg.Cursor[Any], bundle: dict[str, Any], ) -> dict[str, Any]: session = bundle["session"] target_user_id = bundle["target_user_id"] target_persona_id = None target_persona_version = session.get("persona_version") if session.get("persona_id") is not None: source_persona = source_cur.execute( "SELECT code, version FROM app.persona_card WHERE persona_id = %s AND version = %s", (session["persona_id"], session["persona_version"]), ).fetchall() if len(source_persona) != 1: raise RecoveryGuardError("source session persona lookup was not unique") target_persona = target_cur.execute( "SELECT persona_id, version FROM app.persona_card WHERE code = %s AND version = %s", source_persona[0], ).fetchall() if len(target_persona) != 1: raise RecoveryGuardError("candidate session persona lookup was not unique") target_persona_id, target_persona_version = target_persona[0] case_profile_insert: dict[str, Any] | None = None target_case_id = None case_reused = False if bundle["case_profiles"]: source_case = dict(bundle["case_profiles"][0]) matching_cases = target_cur.execute( "SELECT case_id FROM app.case_profile WHERE persona_id = %s AND learner_id = %s FOR UPDATE", (target_persona_id, target_user_id), ).fetchall() if len(matching_cases) > 1: raise RecoveryGuardError("candidate case-profile mapping was not unique") if matching_cases: target_case_id = matching_cases[0][0] case_reused = True else: collision = int( _scalar( target_cur, "SELECT count(*) FROM app.case_profile WHERE case_id = %s", (source_case["case_id"],), ) ) if collision: raise RecoveryGuardError("source case identifier collides in candidate") source_case["learner_id"] = str(target_user_id) source_case["persona_id"] = str(target_persona_id) target_case_id = source_case["case_id"] case_profile_insert = source_case requested_session_no = session.get("session_no") session_no_collision = False resolved_session_no = requested_session_no if bundle["already_present"]: existing_session = target_cur.execute( "SELECT learner_id, case_id, persona_id, persona_version, session_no FROM app.sessions WHERE id = %s", (bundle["session_id"],), ).fetchall() if len(existing_session) != 1: raise RecoveryGuardError("idempotency session lookup was not unique") existing = existing_session[0] if ( existing[0] != target_user_id or existing[1] != target_case_id or existing[2] != target_persona_id or existing[3] != target_persona_version ): raise RecoveryGuardError("idempotency session mapping changed") resolved_session_no = existing[4] session_no_collision = resolved_session_no != requested_session_no elif target_case_id is not None and requested_session_no is not None: session_no_collision = bool( _scalar( target_cur, "SELECT EXISTS(SELECT 1 FROM app.sessions WHERE case_id = %s AND session_no = %s)", (target_case_id, requested_session_no), ) or _scalar( target_cur, "SELECT EXISTS(SELECT 1 FROM app.session_summary WHERE case_id = %s AND session_no = %s)", (target_case_id, requested_session_no), ) ) if session_no_collision: resolved_session_no = int( _scalar( target_cur, """ SELECT GREATEST( COALESCE((SELECT max(session_no) FROM app.sessions WHERE case_id = %s), 0), COALESCE((SELECT max(session_no) FROM app.session_summary WHERE case_id = %s), 0) ) + 1 """, (target_case_id, target_case_id), ) ) return { "target_persona_id": target_persona_id, "target_persona_version": target_persona_version, "target_case_id": target_case_id, "case_profile_insert": case_profile_insert, "case_reused": case_reused, "session_no_collision": session_no_collision, "resolved_session_no": resolved_session_no, } def _protected_fingerprint( cur: psycopg.Cursor[Any], target_user_id: Any, target_case_id: Any ) -> str: protected = _scalar( cur, """ SELECT jsonb_build_object( 'user_profile', (SELECT to_jsonb(u) FROM app.app_user u WHERE u.user_id = %s), 'case_profile', (SELECT to_jsonb(c) FROM app.case_profile c WHERE c.case_id = %s), 'auth_sessions', COALESCE(( SELECT jsonb_agg(to_jsonb(a) ORDER BY a.sid_hash) FROM app.auth_session a WHERE a.user_id = %s ), '[]'::jsonb) ) """, (target_user_id, target_case_id, target_user_id), ) return _stable_hash(protected) def _copy_bundle( source_cur: psycopg.Cursor[Any], target_cur: psycopg.Cursor[Any], bundle: dict[str, Any], context: dict[str, Any], *, allow_session_renumber: bool, ) -> None: if context["session_no_collision"] and not allow_session_renumber: raise RecoveryGuardError("session-number collision requires explicit --allow-session-renumber") target_user_id = bundle["target_user_id"] source_user_id = bundle["source_user_id"] session_id = bundle["session_id"] case_profile = context["case_profile_insert"] if case_profile is not None: _insert_json_record(target_cur, "app", "case_profile", case_profile) session = dict(bundle["session"]) session["learner_id"] = str(target_user_id) session["persona_id"] = ( str(context["target_persona_id"]) if context["target_persona_id"] is not None else None ) session["persona_version"] = context["target_persona_version"] session["case_id"] = ( str(context["target_case_id"]) if context["target_case_id"] is not None else None ) session["session_no"] = context["resolved_session_no"] _insert_json_record(target_cur, "app", "sessions", session) for row in bundle["session_states"]: _insert_json_record(target_cur, "app", "session_state", dict(row)) for row in bundle["model_runs"]: _insert_json_record(target_cur, "audit", "model_run", dict(row)) pulse_audit_trigger_enabled = bool( _scalar( target_cur, """ SELECT tgenabled = 'O' FROM pg_trigger WHERE tgrelid = 'app.alliance_pulse'::regclass AND tgname = 'trg_alliance_pulse_status_audit' """, ) ) if not pulse_audit_trigger_enabled: raise RecoveryGuardError("candidate pulse audit trigger is not enabled") # Recovery copies the two immutable source audit events with their original # timestamps. Suppress only the synthetic INSERT audit row that would be # generated from the already-terminal snapshot. ALTER TABLE is transactional; # any later failure restores the trigger state with the data rollback. target_cur.execute( "ALTER TABLE app.alliance_pulse DISABLE TRIGGER trg_alliance_pulse_status_audit" ) for row in bundle["pulses"]: _insert_json_record(target_cur, "app", "alliance_pulse", dict(row)) target_cur.execute( "ALTER TABLE app.alliance_pulse ENABLE TRIGGER trg_alliance_pulse_status_audit" ) if not bool( _scalar( target_cur, """ SELECT tgenabled = 'O' FROM pg_trigger WHERE tgrelid = 'app.alliance_pulse'::regclass AND tgname = 'trg_alliance_pulse_status_audit' """, ) ): raise RecoveryGuardError("candidate pulse audit trigger was not restored") for row in bundle["pulse_status_events"]: copied = dict(row) if copied.get("changed_by_uid") == source_user_id: copied["changed_by_uid"] = str(target_user_id) _insert_json_record(target_cur, "audit", "alliance_pulse_status_event", copied) for row in bundle["self_assessments"]: copied = dict(row) copied["learner_id"] = str(target_user_id) _insert_json_record(target_cur, "app", "self_assessment", copied) instrument_keys = sorted( {(row["instrument_id"], row["instrument_version"]) for row in bundle["measurements"]} ) for instrument_id, instrument_version in instrument_keys: source_instruments = _rows_json( source_cur, "app", "measurement_instrument", "instrument_id = %s AND instrument_version = %s", (instrument_id, instrument_version), ) if len(source_instruments) != 1: raise RecoveryGuardError("source measurement instrument lookup was not unique") target_instruments = _rows_json( target_cur, "app", "measurement_instrument", "instrument_id = %s AND instrument_version = %s", (instrument_id, instrument_version), ) if not target_instruments: _insert_json_record(target_cur, "app", "measurement_instrument", source_instruments[0]) else: source_comparable = dict(source_instruments[0]) target_comparable = dict(target_instruments[0]) source_comparable.pop("created_at", None) target_comparable.pop("created_at", None) if source_comparable != target_comparable: raise RecoveryGuardError("candidate measurement instrument conflicts with source") for row in _topological_measurements(bundle["measurements"]): _insert_json_record(target_cur, "app", "measurement_event", dict(row)) for row in bundle["summaries"]: copied = dict(row) copied["case_id"] = ( str(context["target_case_id"]) if context["target_case_id"] is not None else None ) copied["session_no"] = context["resolved_session_no"] _insert_json_record(target_cur, "app", "session_summary", copied) if int(_scalar(target_cur, "SELECT count(*) FROM app.sessions WHERE id = %s", (session_id,))) != 1: raise RecoveryGuardError("merged session is not unique in candidate") def _assert_postconditions( cur: psycopg.Cursor[Any], bundle: dict[str, Any], context: dict[str, Any] ) -> dict[str, int]: counts = _target_counts(cur) google_sessions = int( _scalar( cur, """ SELECT count(*) FROM app.sessions s JOIN app.app_user u ON u.user_id = s.learner_id WHERE u.external_id LIKE 'google:%%' """, ) ) actual = {**counts, "google_sessions": google_sessions} if actual != EXPECTED_TARGET_AFTER: raise RecoveryGuardError("candidate total-count postcondition failed") session_id = bundle["session_id"] expected_direct = { "session_states": EXPECTED_SOURCE_ROWS["session_states"], "alliance_pulses": EXPECTED_SOURCE_ROWS["alliance_pulses"], "self_assessments": EXPECTED_SOURCE_ROWS["self_assessments"], "measurement_events": EXPECTED_SOURCE_ROWS["measurement_events"], "session_summaries": EXPECTED_SOURCE_ROWS["session_summaries"], "model_runs": EXPECTED_SOURCE_ROWS["model_runs"], } direct = { "session_states": int( _scalar(cur, "SELECT count(*) FROM app.session_state WHERE session_id = %s", (session_id,)) ), "alliance_pulses": int( _scalar(cur, "SELECT count(*) FROM app.alliance_pulse WHERE session_id = %s", (session_id,)) ), "self_assessments": int( _scalar(cur, "SELECT count(*) FROM app.self_assessment WHERE session_id = %s", (session_id,)) ), "measurement_events": int( _scalar(cur, "SELECT count(*) FROM app.measurement_event WHERE session_id = %s", (session_id,)) ), "session_summaries": int( _scalar(cur, "SELECT count(*) FROM app.session_summary WHERE session_id = %s", (session_id,)) ), "model_runs": int( _scalar(cur, "SELECT count(*) FROM audit.model_run WHERE session_id = %s", (session_id,)) ), } if direct != expected_direct: raise RecoveryGuardError("candidate direct-row postcondition failed") pulse_events = int( _scalar( cur, """ SELECT count(*) FROM audit.alliance_pulse_status_event e JOIN app.alliance_pulse p ON p.pulse_id = e.pulse_id WHERE p.session_id = %s """, (session_id,), ) ) if pulse_events != EXPECTED_SOURCE_ROWS["pulse_status_events"]: raise RecoveryGuardError( "candidate pulse-audit postcondition failed: " f"expected={EXPECTED_SOURCE_ROWS['pulse_status_events']} actual={pulse_events}" ) orphan_queries = ( "SELECT count(*) FROM app.sessions s LEFT JOIN app.app_user u ON u.user_id=s.learner_id WHERE u.user_id IS NULL", "SELECT count(*) FROM app.sessions s LEFT JOIN app.case_profile c ON c.case_id=s.case_id WHERE s.case_id IS NOT NULL AND c.case_id IS NULL", "SELECT count(*) FROM app.session_state x LEFT JOIN app.sessions s ON s.id=x.session_id WHERE s.id IS NULL", "SELECT count(*) FROM app.alliance_pulse x LEFT JOIN app.sessions s ON s.id=x.session_id WHERE s.id IS NULL", "SELECT count(*) FROM app.self_assessment x LEFT JOIN app.sessions s ON s.id=x.session_id LEFT JOIN app.app_user u ON u.user_id=x.learner_id LEFT JOIN app.alliance_pulse p ON p.pulse_id=x.pulse_id WHERE s.id IS NULL OR u.user_id IS NULL OR p.pulse_id IS NULL", "SELECT count(*) FROM app.measurement_event x LEFT JOIN app.sessions s ON s.id=x.session_id LEFT JOIN app.measurement_instrument i ON i.instrument_id=x.instrument_id AND i.instrument_version=x.instrument_version LEFT JOIN audit.model_run m ON m.model_run_id=x.model_run_id WHERE s.id IS NULL OR i.instrument_id IS NULL OR (x.model_run_id IS NOT NULL AND m.model_run_id IS NULL)", "SELECT count(*) FROM app.session_summary x LEFT JOIN app.sessions s ON s.id=x.session_id WHERE s.id IS NULL", "SELECT count(*) FROM audit.alliance_pulse_status_event e LEFT JOIN app.alliance_pulse p ON p.pulse_id=e.pulse_id WHERE p.pulse_id IS NULL", ) orphan_count = sum(int(_scalar(cur, query)) for query in orphan_queries) if orphan_count != 0: raise RecoveryGuardError("candidate orphan postcondition failed") duplicate_external = int( _scalar( cur, "SELECT count(*) FROM (SELECT external_id FROM app.app_user WHERE external_id IS NOT NULL GROUP BY external_id HAVING count(*) > 1) d", ) ) if duplicate_external != 0: raise RecoveryGuardError("candidate duplicate-external-id postcondition failed") app_owned_tables = int( _scalar( cur, """ SELECT count(*) FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace JOIN pg_roles r ON r.oid = c.relowner WHERE n.nspname IN ('app','audit','eval','kb','ds') AND c.relkind IN ('r','p') AND r.rolname = 'vignette' """, ) ) if app_owned_tables != 0: raise RecoveryGuardError("application role owns tables") merged_session_rows = _rows_json( cur, "app", "sessions", "id = %s", (session_id,) ) if len(merged_session_rows) != 1: raise RecoveryGuardError("merged session projection was not unique") merged_session = merged_session_rows[0] source_session = bundle["session"] if ( merged_session.get("started_at") != source_session.get("started_at") or merged_session.get("ended_at") != source_session.get("ended_at") ): raise RecoveryGuardError("merged session timestamps changed") if merged_session.get("session_no") != context["resolved_session_no"]: raise RecoveryGuardError("merged session number changed unexpectedly") case_id = context["target_case_id"] case_session_no_duplicates = int( _scalar( cur, """ SELECT count(*) FROM ( SELECT session_no FROM app.sessions WHERE case_id = %s AND session_no IS NOT NULL GROUP BY session_no HAVING count(*) > 1 ) d """, (case_id,), ) ) chronology_inversions = int( _scalar( cur, """ SELECT count(*) FROM app.sessions earlier JOIN app.sessions later ON later.case_id = earlier.case_id WHERE earlier.case_id = %s AND earlier.session_no < later.session_no AND earlier.started_at > later.started_at """, (case_id,), ) ) is_latest_projection = bool( _scalar( cur, """ SELECT s.session_no = (SELECT max(x.session_no) FROM app.sessions x WHERE x.case_id=s.case_id) AND s.started_at = (SELECT max(x.started_at) FROM app.sessions x WHERE x.case_id=s.case_id) FROM app.sessions s WHERE s.id = %s """, (session_id,), ) ) if case_session_no_duplicates or chronology_inversions or not is_latest_projection: raise RecoveryGuardError("candidate chronological projection postcondition failed") return { **actual, "orphans": orphan_count, "duplicate_external_ids": duplicate_external, "app_owned_tables": app_owned_tables, "case_session_no_duplicates": case_session_no_duplicates, "chronology_inversions": chronology_inversions, "latest_projection": int(is_latest_projection), } def _assert_learner_rls(cur: psycopg.Cursor[Any], bundle: dict[str, Any]) -> None: cur.execute("SET LOCAL ROLE vignette") cur.execute("SELECT set_config('app.current_uid', %s, true)", (str(bundle["target_user_id"]),)) cur.execute("SELECT set_config('app.current_role', 'learner', true)") cur.execute("SELECT set_config('app.ai_context', '0', true)") visible = int( _scalar(cur, "SELECT count(*) FROM app.sessions WHERE id = %s", (bundle["session_id"],)) ) if visible != 1: raise RecoveryGuardError("merged session is not visible through learner RLS") cur.execute("RESET ROLE") def run( *, apply: bool, allow_session_renumber: bool, expected_source_fingerprint: str | None, ) -> dict[str, Any]: source = _settings("VIGNETTE_RECOVERY_SOURCE", 55432) target = _settings("VIGNETTE_RECOVERY_TARGET", 55433) if source.host != "127.0.0.1" or target.host != "127.0.0.1": raise RecoveryGuardError("both recovery database hosts must be 127.0.0.1") with psycopg.connect(**source.kwargs(read_only=True)) as source_conn: with psycopg.connect(**target.kwargs(read_only=False)) as target_conn: with source_conn.transaction(): with target_conn.transaction(): source_cur = source_conn.cursor() target_cur = target_conn.cursor() source_cur.execute("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY") target_cur.execute("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE") target_counts = _target_counts(target_cur) valid_target_counts = ( EXPECTED_TARGET_BEFORE, { "users": EXPECTED_TARGET_AFTER["users"], "sessions": EXPECTED_TARGET_AFTER["sessions"], "turns": EXPECTED_TARGET_AFTER["turns"], }, ) if target_counts not in valid_target_counts: raise RecoveryGuardError("candidate count guard failed") _assert_schema_compatibility(source_cur, target_cur) bundle = _source_bundle(source_cur, target_cur) if ( expected_source_fingerprint is not None and bundle["source_fingerprint"] != expected_source_fingerprint ): raise RecoveryGuardError("source fingerprint changed") expected_already_present = target_counts["sessions"] == EXPECTED_TARGET_AFTER["sessions"] if bundle["already_present"] != expected_already_present: raise RecoveryGuardError("candidate count and idempotency key disagree") context = _resolve_target_context(source_cur, target_cur, bundle) protected_before = _protected_fingerprint( target_cur, bundle["target_user_id"], context["target_case_id"], ) plan = { "source_read_only": bool( _scalar(source_cur, "SELECT current_setting('transaction_read_only')::boolean") ), "target_guard": target_counts, "source_rows": EXPECTED_SOURCE_ROWS, "source_fingerprint": bundle["source_fingerprint"], "session_key_verified": True, "case_profile_reused": context["case_reused"], "case_profile_inserted": context["case_profile_insert"] is not None, "session_number_collision": context["session_no_collision"], "session_renumbered": bool(context["session_no_collision"]), "original_session_no": bundle["session"].get("session_no"), "resolved_session_no": context["resolved_session_no"], } if not plan["source_read_only"]: raise RecoveryGuardError("source transaction is not read-only") if not apply: return { "mode": "already-applied" if bundle["already_present"] else "dry-run", "plan": plan, } if bundle["already_present"]: post = _assert_postconditions(target_cur, bundle, context) protected_after = _protected_fingerprint( target_cur, bundle["target_user_id"], context["target_case_id"], ) if protected_before != protected_after: raise RecoveryGuardError("protected candidate data changed") _assert_learner_rls(target_cur, bundle) return { "mode": "already-applied", "plan": plan, "post": post, "duplicate_insertions": 0, "protected_data_unchanged": True, "rls_visible": True, } _copy_bundle( source_cur, target_cur, bundle, context, allow_session_renumber=allow_session_renumber, ) post = _assert_postconditions(target_cur, bundle, context) protected_after = _protected_fingerprint( target_cur, bundle["target_user_id"], context["target_case_id"], ) if protected_before != protected_after: raise RecoveryGuardError("protected candidate data changed") _assert_learner_rls(target_cur, bundle) return { "mode": "applied", "plan": plan, "post": post, "duplicate_insertions": 0, "protected_data_unchanged": True, "rls_visible": True, } def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--apply", action="store_true", help="write only the guarded 55433 candidate") parser.add_argument( "--allow-session-renumber", action="store_true", help="append with the next case-local number when the recovered case already owns the source number", ) parser.add_argument( "--expected-source-fingerprint", help="optional SHA-256 from the first guarded run for idempotency verification", ) args = parser.parse_args() if args.allow_session_renumber and not args.apply: raise RecoveryGuardError("--allow-session-renumber requires --apply") if args.expected_source_fingerprint is not None and ( len(args.expected_source_fingerprint) != 64 or any(ch not in "0123456789abcdef" for ch in args.expected_source_fingerprint) ): raise RecoveryGuardError("expected source fingerprint must be lowercase SHA-256") result = run( apply=args.apply, allow_session_renumber=args.allow_session_renumber, expected_source_fingerprint=args.expected_source_fingerprint, ) print(result) return 0 if __name__ == "__main__": try: raise SystemExit(main()) except RecoveryGuardError as exc: print(f"RECOVERY_GUARD_FAILED: {exc}", file=sys.stderr) raise SystemExit(2) except psycopg.Error as exc: # Database exception details can include row identifiers or content. frames = [ f"{frame.name}:{frame.lineno}" for frame in traceback.extract_tb(exc.__traceback__) if frame.filename.endswith("merge-recovery-google-session-delta.py") ] print( "RECOVERY_DATABASE_FAILED: " f"type={type(exc).__name__} sqlstate={exc.sqlstate or 'unknown'} " f"location={'/'.join(frames)}", file=sys.stderr, ) raise SystemExit(3) except Exception as exc: # pragma: no cover - final secret/PII-safe boundary frames = [ f"{frame.name}:{frame.lineno}" for frame in traceback.extract_tb(exc.__traceback__) if frame.filename.endswith("merge-recovery-google-session-delta.py") ] print( f"RECOVERY_FAILED: {type(exc).__name__} location={'/'.join(frames)}", file=sys.stderr, ) raise SystemExit(4)