deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/runtime/checkpointer/checkpoint_migration.py
2026-09-07 18:24:55 +08:00

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