Files
EvoScientist-Multi/EvoScientist/web_checkpointer.py
T

257 lines
12 KiB
Python

"""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 generation,
EXTRACT(EPOCH FROM expires_at)*1000 AS expires_at_ms""",
(clean_thread, clean_owner, ttl_seconds),
)
row = await cursor.fetchone()
if row is None:
raise RuntimeError("TURN_LEASE_BUSY")
await cursor.execute(
"""INSERT INTO evomemory_checkpoint_versions
(thread_id,sequence,checkpoint_id)
VALUES (%s,0,NULL) ON CONFLICT(thread_id) DO NOTHING""",
(clean_thread,),
)
await cursor.execute(
"""SELECT sequence,checkpoint_id
FROM evomemory_checkpoint_versions WHERE thread_id=%s""",
(clean_thread,),
)
version = await cursor.fetchone()
sequence = int(version["sequence"] if version else 0)
checkpoint_id = str((version or {}).get("checkpoint_id") or "root")
snapshot = hashlib.sha256(
f"{clean_thread}\0{sequence}\0{checkpoint_id}".encode()
).hexdigest()
return self.TurnLease(
clean_thread,
clean_owner,
int(row["generation"]),
int(row["expires_at_ms"]),
f"sha256:{snapshot}",
)
async def renew_turn_lease(self, lease, *, ttl_seconds):
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 expires_at=NOW()+(%s*INTERVAL '1 second'),updated_at=NOW()
WHERE thread_id=%s AND generation=%s AND owner_id=%s
AND 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