deerflow-code/offline-backend-20260512/backend/app/gateway/routers/checkpoint_migration.py
2026-09-07 18:24:55 +08:00

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)}