257 lines
12 KiB
Python
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
|