"""Admin checkpoint migration endpoints. These endpoints support the SQLite -> Postgres transition for LangGraph checkpoints. The live checkpointer should be Postgres with SQLite dual-read enabled; this router only starts manual, bounded migrations for selected time ranges. """ from __future__ import annotations import asyncio import contextlib import hashlib import json import logging import os from datetime import datetime, timezone from pathlib import Path from typing import Any from fastapi import APIRouter, HTTPException, Request from pydantic import BaseModel, Field from app.gateway.auth.disabled_mode import is_auth_disabled from app.gateway.deps import get_current_user_from_request from deerflow.config.runtime_paths import runtime_home from deerflow.runtime.checkpointer.checkpoint_migration import ( SQLiteCheckpointSource, copy_thread_state, sqlite_recent_threads, target_has_thread, ) from deerflow.runtime.store._sqlite_utils import resolve_sqlite_conn_str logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/admin/checkpoint-migration", tags=["checkpoint-migration"]) _TASKS: dict[str, dict[str, Any]] = {} class MigrationStartRequest(BaseModel): start_at: str = Field(description="Inclusive start datetime, ISO string.") end_at: str = Field(description="Inclusive end datetime, ISO string.") source_paths: list[str] | None = Field(default=None, description="Optional SQLite sources. Defaults to configured dual-read paths.") class MigrationTaskResponse(BaseModel): task_id: str status: str start_at: str end_at: str source_paths: list[str] total_threads: int = 0 copied_threads: int = 0 skipped_threads: int = 0 failed_threads: int = 0 skipped_ranges: int = 0 message: str = "" error: str | None = None created_at: str updated_at: str async def _require_admin(request: Request) -> None: if is_auth_disabled(): return user = await get_current_user_from_request(request) if getattr(user, "system_role", None) != "admin": raise HTTPException(status_code=403, detail="Checkpoint migration requires admin role") def _now_iso() -> str: return datetime.now(timezone.utc).isoformat() def _parse_dt(value: str) -> datetime: try: normalized = value.replace("Z", "+00:00") dt = datetime.fromisoformat(normalized) except ValueError as exc: raise HTTPException(status_code=400, detail=f"Invalid datetime: {value}") from exc if dt.tzinfo is None: dt = dt.replace(tzinfo=timezone.utc) return dt def _sqlite_datetime(dt: datetime) -> str: # Existing SQLite rows store naive wall-clock strings like # "2026-05-18 06:41:10.989427"; match that format for comparisons. return dt.astimezone(timezone.utc).replace(tzinfo=None).isoformat(sep=" ") def _configured_sources(request: Request) -> list[str]: config = getattr(request.app.state, "config", None) cp = getattr(config, "checkpointer", None) paths = list(getattr(cp, "dual_read_sqlite_paths", []) or []) raw_env = os.getenv("DEER_FLOW_CHECKPOINT_SQLITE_PATHS", "") paths.extend(part.strip() for part in raw_env.split(",") if part.strip()) resolved: list[str] = [] seen: set[str] = set() for path in paths: conn_str = resolve_sqlite_conn_str(path) if conn_str == ":memory:" or conn_str.startswith("file:"): continue if conn_str not in seen: resolved.append(conn_str) seen.add(conn_str) return resolved def _target_checkpointer(request: Request): checkpointer = getattr(request.app.state, "checkpointer", None) if checkpointer is None: raise HTTPException(status_code=503, detail="Checkpointer is not available") return getattr(checkpointer, "primary", checkpointer) def _state_path() -> Path: return runtime_home() / "checkpoint_migration_state.json" def _load_completed_ranges() -> set[str]: try: payload = json.loads(_state_path().read_text(encoding="utf-8")) except FileNotFoundError: return set() except Exception: logger.warning("Ignoring unreadable checkpoint migration state file", exc_info=True) return set() return {str(item) for item in payload.get("completed_ranges", [])} def _save_completed_ranges(items: set[str]) -> None: path = _state_path() path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(".tmp") tmp.write_text(json.dumps({"completed_ranges": sorted(items)}, ensure_ascii=False, indent=2), encoding="utf-8") tmp.replace(path) def _range_key(source: str, start_at: str, end_at: str) -> str: raw = json.dumps({"source": str(Path(source).resolve()), "start_at": start_at, "end_at": end_at}, sort_keys=True) return hashlib.sha256(raw.encode("utf-8")).hexdigest() async def _open_sqlite_source(path: str) -> SQLiteCheckpointSource: from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver ctx = AsyncSqliteSaver.from_conn_string(path) saver = await ctx.__aenter__() await saver.setup() return SQLiteCheckpointSource( path=path, saver=saver, context=ctx, progress_path=Path(path).with_suffix(Path(path).suffix + ".pg-migration.json"), ) def _task_response(task: dict[str, Any]) -> MigrationTaskResponse: return MigrationTaskResponse(**task) async def _run_task(task_id: str, target, start_at: datetime, end_at: datetime) -> None: task = _TASKS[task_id] completed = await asyncio.to_thread(_load_completed_ranges) start_sql = _sqlite_datetime(start_at) end_sql = _sqlite_datetime(end_at) range_start = start_at.isoformat() range_end = end_at.isoformat() try: task["status"] = "running" task["updated_at"] = _now_iso() for source_path in task["source_paths"]: if not Path(source_path).exists(): task["message"] = f"SQLite source not found: {source_path}" task["skipped_ranges"] += 1 continue key = _range_key(source_path, range_start, range_end) if key in completed: task["skipped_ranges"] += 1 task["message"] = "Selected range was already migrated; skipped" task["updated_at"] = _now_iso() continue thread_ids = await asyncio.to_thread(sqlite_recent_threads, source_path, since=start_sql, until=end_sql) if not thread_ids: task["message"] = f"No threads found in selected range for {Path(source_path).name}" completed.add(key) await asyncio.to_thread(_save_completed_ranges, completed) task["updated_at"] = _now_iso() continue task["total_threads"] += len(thread_ids) source = await _open_sqlite_source(source_path) try: for thread_id in thread_ids: if await target_has_thread(target, thread_id): task["skipped_threads"] += 1 else: copied = await copy_thread_state(source.saver, target, thread_id) if copied: task["copied_threads"] += 1 else: task["skipped_threads"] += 1 task["updated_at"] = _now_iso() await asyncio.sleep(0) finally: with contextlib.suppress(Exception): await source.close() completed.add(key) await asyncio.to_thread(_save_completed_ranges, completed) task["status"] = "completed" task["message"] = task["message"] or "Migration completed" except Exception as exc: logger.exception("Checkpoint migration task failed: %s", task_id) task["status"] = "failed" task["error"] = str(exc) finally: task["updated_at"] = _now_iso() @router.post("/tasks", response_model=MigrationTaskResponse) async def start_migration(body: MigrationStartRequest, request: Request): await _require_admin(request) start_at = _parse_dt(body.start_at) end_at = _parse_dt(body.end_at) if end_at <= start_at: raise HTTPException(status_code=400, detail="end_at must be later than start_at") sources = [resolve_sqlite_conn_str(p) for p in (body.source_paths or _configured_sources(request))] sources = [p for p in dict.fromkeys(sources) if p != ":memory:" and not p.startswith("file:")] if not sources: raise HTTPException(status_code=400, detail="No SQLite source paths configured") target = _target_checkpointer(request) task_id = hashlib.sha256(f"{_now_iso()}:{sources}:{body.start_at}:{body.end_at}".encode("utf-8")).hexdigest()[:16] now = _now_iso() _TASKS[task_id] = { "task_id": task_id, "status": "queued", "start_at": start_at.isoformat(), "end_at": end_at.isoformat(), "source_paths": sources, "total_threads": 0, "copied_threads": 0, "skipped_threads": 0, "failed_threads": 0, "skipped_ranges": 0, "message": "", "error": None, "created_at": now, "updated_at": now, } asyncio.create_task(_run_task(task_id, target, start_at, end_at)) return _task_response(_TASKS[task_id]) @router.get("/tasks/{task_id}", response_model=MigrationTaskResponse) async def get_migration_task(task_id: str, request: Request): await _require_admin(request) task = _TASKS.get(task_id) if task is None: raise HTTPException(status_code=404, detail="Migration task not found") return _task_response(task) @router.get("/sources") async def list_sources(request: Request): await _require_admin(request) return {"source_paths": _configured_sources(request)}