"""Web-only PostgreSQL saver. CLI sessions remain in sessions.py. The saver also owns PostgreSQL turn fencing for Web-hosted agents. There is no SQLite fallback. Schema setup is performed only by the reviewed Gateway and LangGraph migration paths. """ from __future__ import annotations import asyncio import hashlib import os from contextlib import asynccontextmanager from dataclasses import dataclass def _web_saver_type(): from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver class WebPostgresSaver(AsyncPostgresSaver): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._turn_fence_lock = asyncio.Lock() @dataclass(frozen=True, slots=True) class TurnLease: thread_id: str owner_id: str fencing_token: int expires_at_ms: int checkpoint_snapshot_id: str @asynccontextmanager async def _turn_lock(self, thread_id): lock_key = str(thread_id or "") async with self._turn_fence_lock: async with self._cursor() as cursor: await cursor.execute( "SELECT pg_advisory_lock(hashtextextended(%s, 90421066))", (lock_key,), ) try: yield finally: async with self._cursor() as cursor: await cursor.execute( "SELECT pg_advisory_unlock(hashtextextended(%s, 90421066))", (lock_key,), ) async def acquire_turn_lease(self, thread_id, owner_id, *, ttl_seconds): clean_thread = str(thread_id or "").strip() clean_owner = str(owner_id or "").strip() if not clean_thread or not clean_owner or ttl_seconds < 1: raise ValueError("thread, owner and positive TTL are required") async with self._turn_lock(clean_thread): async with self._cursor(pipeline=True) as cursor: await cursor.execute( """INSERT INTO evomemory_checkpoint_turn_fences (thread_id,generation,owner_id,expires_at,updated_at) VALUES (%s,1,%s,NOW()+(%s*INTERVAL '1 second'),NOW()) ON CONFLICT(thread_id) DO UPDATE SET generation=evomemory_checkpoint_turn_fences.generation+1, owner_id=EXCLUDED.owner_id, expires_at=EXCLUDED.expires_at, updated_at=NOW() WHERE evomemory_checkpoint_turn_fences.owner_id IS NULL OR evomemory_checkpoint_turn_fences.owner_id=EXCLUDED.owner_id OR evomemory_checkpoint_turn_fences.expires_at=NOW() RETURNING EXTRACT(EPOCH FROM expires_at)*1000 AS expires_at_ms""", ( ttl_seconds, lease.thread_id, lease.fencing_token, lease.owner_id, ), ) row = await cursor.fetchone() if row is None: raise RuntimeError("TURN_LEASE_LOST") return self.TurnLease( lease.thread_id, lease.owner_id, lease.fencing_token, int(row["expires_at_ms"]), lease.checkpoint_snapshot_id, ) async def release_turn_lease(self, lease): async with self._turn_lock(lease.thread_id): async with self._cursor(pipeline=True) as cursor: await cursor.execute( """UPDATE evomemory_checkpoint_turn_fences SET owner_id=NULL,expires_at=NULL,updated_at=NOW() WHERE thread_id=%s AND generation=%s AND owner_id=%s RETURNING 1""", (lease.thread_id, lease.fencing_token, lease.owner_id), ) return await cursor.fetchone() is not None async def _require_write_lease(self, config): configurable = dict(config.get("configurable") or {}) thread_id = str(configurable.get("thread_id") or "") if not thread_id.startswith("web:"): return thread_id, 0 owner_id = str(configurable.get("turn_lease_owner") or "") token = int(configurable.get("turn_fencing_token") or 0) async with self._cursor() as cursor: await cursor.execute( """SELECT 1 FROM evomemory_checkpoint_turn_fences WHERE thread_id=%s AND generation=%s AND owner_id=%s AND expires_at>=NOW()""", (thread_id, token, owner_id), ) if await cursor.fetchone() is None: raise RuntimeError("TURN_FENCED") return thread_id, token async def aput(self, config, checkpoint, metadata, new_versions): thread_id = str(config.get("configurable", {}).get("thread_id") or "") async with self._turn_lock(thread_id): thread_id, token = await self._require_write_lease(config) result = await super().aput(config, checkpoint, metadata, new_versions) if token: checkpoint_id = str( result.get("configurable", {}).get("checkpoint_id") or "" ) async with self._cursor(pipeline=True) as cursor: await cursor.execute( """INSERT INTO evomemory_checkpoint_versions (thread_id,sequence,checkpoint_id,updated_at) VALUES (%s,1,%s,NOW()) ON CONFLICT(thread_id) DO UPDATE SET sequence=evomemory_checkpoint_versions.sequence+1, checkpoint_id=EXCLUDED.checkpoint_id, updated_at=NOW()""", (thread_id, checkpoint_id), ) return result async def aput_writes(self, config, writes, task_id, task_path=""): thread_id = str(config.get("configurable", {}).get("thread_id") or "") async with self._turn_lock(thread_id): await self._require_write_lease(config) await super().aput_writes(config, writes, task_id, task_path) async def adelete_for_runs(self, run_ids): values = tuple(dict.fromkeys(str(run_id) for run_id in run_ids)) if not values: return async with self._cursor(pipeline=True) as cursor: await cursor.execute( """DELETE FROM checkpoint_writes w USING checkpoints c WHERE w.thread_id=c.thread_id AND w.checkpoint_ns=c.checkpoint_ns AND w.checkpoint_id=c.checkpoint_id AND c.metadata->>'run_id'=ANY(%s)""", (list(values),), ) await cursor.execute( "DELETE FROM checkpoints WHERE metadata->>'run_id'=ANY(%s)", (list(values),), ) return WebPostgresSaver @asynccontextmanager async def create_web_checkpointer(): from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer from psycopg import AsyncConnection from psycopg.rows import dict_row dsn = os.environ.get("EVOSCIENTIST_WEB_CHECKPOINT_DSN", "") if not dsn.startswith(("postgresql://", "postgres://")): raise RuntimeError("Web checkpointer requires an explicit PostgreSQL DSN") serde = JsonPlusSerializer( allowed_msgpack_modules=frozenset( { ("EvoScientist.llm.errors", "AgentControlError"), ("EvoScientist.llm.errors", "ModelToolProtocolError"), ("EvoScientist.llm.errors", "ProviderStreamError"), } ) ) try: conn = await AsyncConnection.connect( dsn, autocommit=True, prepare_threshold=0, row_factory=dict_row, connect_timeout=5, options="-c search_path=public", ) except Exception: raise RuntimeError("Web checkpoint PostgreSQL connection failed") from None async with conn: saver = _web_saver_type()(conn, serde=serde) try: async with conn.cursor() as cursor: await cursor.execute("SELECT v FROM checkpoint_migrations ORDER BY v") versions = [row["v"] for row in await cursor.fetchall()] if versions != list(range(len(saver.MIGRATIONS))): raise RuntimeError("migration version mismatch") for table in ( "checkpoints", "checkpoint_blobs", "checkpoint_writes", "evomemory_checkpoint_turn_fences", "evomemory_checkpoint_versions", ): await cursor.execute(f"SELECT 1 FROM {table} LIMIT 0") except Exception: raise RuntimeError( "Web checkpoint schema is not ready; run reviewed setup separately" ) from None yield saver