"""SQLite-to-Postgres checkpoint migration and read fallback helpers. The live gateway can wrap a Postgres checkpointer with ``DualReadCheckpointer``. New writes go to Postgres, reads prefer Postgres and fall back to legacy SQLite. Manual migration scripts reuse the copy helpers in this module. """ from __future__ import annotations import asyncio import contextlib import json import logging import sqlite3 from dataclasses import dataclass from pathlib import Path from typing import Any logger = logging.getLogger(__name__) @dataclass class SQLiteCheckpointSource: path: str saver: Any context: Any progress_path: Path closed: bool = False async def close(self) -> None: if self.closed: return self.closed = True try: await self.context.__aexit__(None, None, None) except Exception: logger.warning("Failed to close legacy SQLite checkpointer: %s", self.path, exc_info=True) def _sqlite_db_exists(path: str) -> bool: return Path(path).exists() def _sqlite_threads(path: str) -> list[str]: with sqlite3.connect(path) as conn: rows = conn.execute("SELECT DISTINCT thread_id FROM checkpoints ORDER BY thread_id").fetchall() return [str(row[0]) for row in rows if row and row[0]] def _sqlite_has_table(path: str, table: str) -> bool: with sqlite3.connect(path) as conn: row = conn.execute( "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?", (table,), ).fetchone() return row is not None def sqlite_recent_threads(path: str, *, since: str | None = None, until: str | None = None) -> list[str]: """Return threads updated since an ISO timestamp when ``threads_meta`` exists.""" if not _sqlite_has_table(path, "threads_meta"): return [] query = "SELECT thread_id FROM threads_meta" where: list[str] = [] params: list[str] = [] if since: where.append("updated_at >= ?") params.append(since) if until: where.append("updated_at <= ?") params.append(until) if where: query += " WHERE " + " AND ".join(where) query += " ORDER BY updated_at DESC" with sqlite3.connect(path) as conn: rows = conn.execute(query, tuple(params)).fetchall() return [str(row[0]) for row in rows if row and row[0]] def _sqlite_thread_checkpoint_count(path: str, thread_id: str) -> int: with sqlite3.connect(path) as conn: row = conn.execute( "SELECT COUNT(*) FROM checkpoints WHERE thread_id = ?", (thread_id,), ).fetchone() return int(row[0] if row else 0) def _load_progress(path: Path) -> set[str]: try: payload = json.loads(path.read_text(encoding="utf-8")) except FileNotFoundError: return set() except Exception: logger.warning("Ignoring unreadable checkpoint migration progress file: %s", path, exc_info=True) return set() migrated = payload.get("migrated_threads", []) return {str(x) for x in migrated if x} def _save_progress(path: Path, migrated: set[str]) -> None: path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(path.suffix + ".tmp") tmp.write_text( json.dumps({"migrated_threads": sorted(migrated)}, ensure_ascii=False, indent=2), encoding="utf-8", ) tmp.replace(path) async def _copy_thread_state(source, target, thread_id: str) -> bool: """Copy one thread through the public LangGraph checkpointer API.""" config = {"configurable": {"thread_id": thread_id}} tuples = [t async for t in source.alist(config)] if not tuples: return False for tup in reversed(tuples): configurable = (tup.config or {}).get("configurable", {}) or {} checkpoint_ns = configurable.get("checkpoint_ns", "") or "" parent_id = ((tup.parent_config or {}).get("configurable", {}) or {}).get("checkpoint_id") put_config: dict[str, Any] = {"configurable": {"thread_id": thread_id, "checkpoint_ns": checkpoint_ns}} if parent_id: put_config["configurable"]["checkpoint_id"] = parent_id new_versions = (tup.checkpoint or {}).get("channel_versions", {}) or {} await target.aput(put_config, tup.checkpoint, tup.metadata or {}, new_versions) if tup.pending_writes: by_task: dict[str, list] = {} for task_id, channel, value in tup.pending_writes: by_task.setdefault(task_id, []).append((channel, value)) checkpoint_id = (tup.checkpoint or {}).get("id") for task_id, writes in by_task.items(): writes_config = { "configurable": { "thread_id": thread_id, "checkpoint_ns": checkpoint_ns, "checkpoint_id": checkpoint_id, } } await target.aput_writes(writes_config, writes, task_id) return True async def copy_thread_state(source, target, thread_id: str) -> bool: return await _copy_thread_state(source, target, thread_id) async def _target_thread_checkpoint_count(target, thread_id: str) -> int: config = {"configurable": {"thread_id": thread_id}} count = 0 async for _ in target.alist(config): count += 1 return count async def target_has_thread(target, thread_id: str) -> bool: config = {"configurable": {"thread_id": thread_id}} async for _ in target.alist(config, limit=1): return True return False class DualReadCheckpointer: """Postgres-first checkpointer with SQLite read fallback.""" def __init__(self, primary, sources: list[SQLiteCheckpointSource], *, sleep_seconds: float = 0.05): self.primary = primary self.sources = sources self.sleep_seconds = max(0.0, sleep_seconds) self._closed = False def __getattr__(self, name: str): return getattr(self.primary, name) async def aclose_sources(self) -> None: self._closed = True for source in self.sources: await source.close() async def aget_tuple(self, config): found = await self.primary.aget_tuple(config) if found is not None: return found for source in self.sources: if source.closed: continue found = await source.saver.aget_tuple(config) if found is not None: return found return None async def aget(self, config): found = await self.primary.aget(config) if found is not None: return found for source in self.sources: if source.closed: continue found = await source.saver.aget(config) if found is not None: return found return None async def alist(self, config, *, filter=None, before=None, limit=None): seen: set[tuple[str, str, str]] = set() yielded = 0 async for tup in self.primary.alist(config, filter=filter, before=before, limit=limit): key = self._tuple_key(tup) seen.add(key) yield tup yielded += 1 if limit is not None and yielded >= limit: return remaining = None if limit is None else max(0, limit - yielded) for source in self.sources: if source.closed: continue async for tup in source.saver.alist(config, filter=filter, before=before, limit=remaining): key = self._tuple_key(tup) if key in seen: continue seen.add(key) yield tup yielded += 1 if limit is not None: remaining = max(0, limit - yielded) if remaining <= 0: return async def aput(self, *args, **kwargs): return await self.primary.aput(*args, **kwargs) async def aput_writes(self, *args, **kwargs): return await self.primary.aput_writes(*args, **kwargs) async def adelete_thread(self, thread_id: str) -> None: delete = getattr(self.primary, "adelete_thread", None) if delete is not None: await delete(thread_id) for source in self.sources: if source.closed: continue delete = getattr(source.saver, "adelete_thread", None) if delete is not None: await delete(thread_id) @staticmethod def _tuple_key(tup) -> tuple[str, str, str]: configurable = (tup.config or {}).get("configurable", {}) or {} return ( str(configurable.get("thread_id", "")), str(configurable.get("checkpoint_ns", "") or ""), str(configurable.get("checkpoint_id", "") or (tup.checkpoint or {}).get("id", "")), ) async def _verify_source(self, source: SQLiteCheckpointSource, threads: list[str]) -> bool: for thread_id in threads: source_count = await asyncio.to_thread(_sqlite_thread_checkpoint_count, source.path, thread_id) target_count = await _target_thread_checkpoint_count(self.primary, thread_id) if target_count < source_count: logger.warning( "Checkpoint migration verification failed: source=%s thread_id=%s source_count=%d target_count=%d", source.path, thread_id, source_count, target_count, ) return False return True