277 lines
9.5 KiB
Python
277 lines
9.5 KiB
Python
"""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
|