277 lines
9.7 KiB
Python
277 lines
9.7 KiB
Python
"""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)}
|