"""Migrate DeerFlow business tables from SQLite to MySQL. This intentionally excludes LangGraph checkpoint tables (``checkpoints`` and ``writes``). Keep those on SQLite unless/until a MySQL checkpointer is added. """ from __future__ import annotations import argparse import asyncio import json import os import sqlite3 import sys from collections.abc import Iterable from datetime import datetime, timezone from pathlib import Path from typing import Any import sqlalchemy as sa from sqlalchemy import MetaData, text from sqlalchemy.engine import make_url from sqlalchemy.ext.asyncio import create_async_engine BUSINESS_TABLES = [ "users", "threads_meta", "runs", "run_events", "feedback", "agents", "skills", "scheduled_tasks", "scheduled_task_subscriptions", "scheduled_task_runs", "scheduled_task_delivery_profiles", "scheduler_threads", "llm_call_metrics", "notifications", "recommended_questions", "tags", "tag_assignments", "thread_shares", ] JSON_COLUMNS = { ("threads_meta", "metadata_json"), ("runs", "metadata_json"), ("runs", "kwargs_json"), ("run_events", "event_metadata"), ("scheduled_tasks", "execution_context_json"), ("scheduled_task_runs", "result_json"), ("notifications", "payload"), } def _json_serializer(obj: object) -> str: return json.dumps(obj, ensure_ascii=False) def _add_vendor_path(vendor_dir: str | None) -> None: candidates: list[Path] = [] if vendor_dir: candidates.append(Path(vendor_dir)) if env_home := os.getenv("DEER_FLOW_HOME"): candidates.append(Path(env_home) / "python-vendor") candidates.append(Path(".deer-flow/python-vendor")) for candidate in candidates: resolved = candidate.resolve() if resolved.exists(): value = str(resolved) if value not in sys.path: sys.path.insert(0, value) def _normalize_mysql_url(raw: str) -> str: if raw.startswith("mysql://"): return raw.replace("mysql://", "mysql+asyncmy://", 1) return raw def _quote_sqlite_identifier(name: str) -> str: return '"' + name.replace('"', '""') + '"' def _sqlite_tables(conn: sqlite3.Connection) -> set[str]: rows = conn.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall() return {str(row[0]) for row in rows} def _sqlite_columns(conn: sqlite3.Connection, table_name: str) -> list[str]: rows = conn.execute(f"PRAGMA table_info({_quote_sqlite_identifier(table_name)})").fetchall() return [str(row[1]) for row in rows] def _sqlite_count(conn: sqlite3.Connection, table_name: str) -> int: return int(conn.execute(f"SELECT COUNT(*) FROM {_quote_sqlite_identifier(table_name)}").fetchone()[0]) def _parse_datetime(value: Any) -> Any: if value is None or isinstance(value, datetime): return value if isinstance(value, (int, float)): return datetime.fromtimestamp(value, tz=timezone.utc).replace(tzinfo=None) if not isinstance(value, str): return value raw = value.strip() if not raw: return None if raw.endswith("Z"): raw = raw[:-1] + "+00:00" parsed = datetime.fromisoformat(raw) if parsed.tzinfo is not None: parsed = parsed.astimezone(timezone.utc).replace(tzinfo=None) return parsed def _parse_json(value: Any) -> Any: if value is None or isinstance(value, (dict, list, int, float, bool)): return value if isinstance(value, bytes): value = value.decode("utf-8") if isinstance(value, str): raw = value.strip() if not raw: return None return json.loads(raw) return value def _convert_value(table_name: str, column: sa.Column, value: Any) -> Any: if value is None: return None if (table_name, column.name) in JSON_COLUMNS or isinstance(column.type, sa.JSON): parsed = _parse_json(value) if not isinstance(column.type, sa.JSON): return json.dumps(parsed, ensure_ascii=False) return parsed if isinstance(column.type, sa.DateTime): return _parse_datetime(value) if isinstance(column.type, sa.Boolean): return bool(value) return value def _build_rows( *, table_name: str, table: sa.Table, rows: Iterable[sqlite3.Row], columns: list[str], ) -> list[dict[str, Any]]: result: list[dict[str, Any]] = [] for row in rows: item: dict[str, Any] = {} for column_name in columns: column = table.c[column_name] item[column_name] = _convert_value(table_name, column, row[column_name]) result.append(item) return result async def _create_database_if_needed(mysql_url: str) -> None: parsed = make_url(mysql_url) db_name = parsed.database if not db_name: raise ValueError("MySQL URL must include a database name") maintenance = parsed.set(database="mysql") engine = create_async_engine( maintenance.render_as_string(hide_password=False), isolation_level="AUTOCOMMIT", pool_pre_ping=True, ) quoted = "`" + db_name.replace("`", "``") + "`" try: async with engine.connect() as conn: await conn.execute(text(f"CREATE DATABASE IF NOT EXISTS {quoted} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci")) finally: await engine.dispose() async def _create_schema(mysql_url: str) -> None: from deerflow.persistence.base import Base import deerflow.persistence.models # noqa: F401 engine = create_async_engine(mysql_url, pool_pre_ping=True, json_serializer=_json_serializer) try: async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) finally: await engine.dispose() def _quote_mysql_identifier(name: str) -> str: return "`" + name.replace("`", "``") + "`" async def _drop_empty_target_schema(mysql_url: str, table_names: list[str]) -> None: """Drop known target tables only when all existing known tables are empty. This is intentionally conservative. It lets a failed first attempt clean up half-created MySQL tables before re-running with fixed schema definitions, while refusing to touch a database that already contains migrated rows. """ engine = create_async_engine(mysql_url, pool_pre_ping=True) try: async with engine.begin() as conn: result = await conn.execute( text( "SELECT TABLE_NAME FROM information_schema.TABLES " "WHERE TABLE_SCHEMA = DATABASE()" ) ) existing = {str(row[0]) for row in result} candidates = [name for name in table_names if name in existing] if not candidates: return non_empty: list[str] = [] for table_name in candidates: count = int( ( await conn.execute( text(f"SELECT COUNT(*) FROM {_quote_mysql_identifier(table_name)}") ) ).scalar_one() ) if count > 0: non_empty.append(table_name) if non_empty: joined = ", ".join(non_empty) print(f"Target MySQL has non-empty tables; skip empty-schema cleanup: {joined}") return await conn.execute(text("SET FOREIGN_KEY_CHECKS=0")) try: for table_name in reversed(candidates): await conn.execute(text(f"DROP TABLE {_quote_mysql_identifier(table_name)}")) print(f"dropped empty target table: {table_name}") finally: await conn.execute(text("SET FOREIGN_KEY_CHECKS=1")) finally: await engine.dispose() def _run_alembic(mysql_url: str) -> None: from alembic import command from alembic.config import Config import deerflow.persistence as persistence_pkg migrations_dir = str(Path(persistence_pkg.__file__).parent / "migrations") cfg = Config() cfg.set_main_option("script_location", migrations_dir) cfg.set_main_option("sqlalchemy.url", mysql_url.replace("%", "%%")) command.upgrade(cfg, "head") async def _reflect_target(mysql_url: str) -> tuple[sa.ext.asyncio.AsyncEngine, MetaData]: engine = create_async_engine( mysql_url, pool_pre_ping=True, pool_recycle=1800, json_serializer=_json_serializer, ) metadata = MetaData() async with engine.begin() as conn: await conn.run_sync(metadata.reflect) return engine, metadata async def _target_count(conn, table: sa.Table) -> int: return int((await conn.execute(sa.select(sa.func.count()).select_from(table))).scalar_one()) async def _ensure_target_empty(engine, metadata: MetaData, table_names: list[str]) -> None: non_empty: list[str] = [] async with engine.connect() as conn: for table_name in table_names: table = metadata.tables.get(table_name) if table is not None and await _target_count(conn, table) > 0: non_empty.append(table_name) if non_empty: joined = ", ".join(non_empty) raise RuntimeError(f"Target MySQL tables are not empty: {joined}. Re-run with --truncate-target if this is intentional.") async def _truncate_target(engine, metadata: MetaData, table_names: list[str]) -> None: async with engine.begin() as conn: await conn.execute(text("SET FOREIGN_KEY_CHECKS=0")) try: for table_name in reversed(table_names): table = metadata.tables.get(table_name) if table is not None: await conn.execute(table.delete()) finally: await conn.execute(text("SET FOREIGN_KEY_CHECKS=1")) async def _copy_table( *, source: sqlite3.Connection, engine, metadata: MetaData, table_name: str, batch_size: int, ) -> tuple[int, int]: table = metadata.tables.get(table_name) if table is None: print(f"skip {table_name}: not present in MySQL schema") return 0, 0 source_columns = _sqlite_columns(source, table_name) columns = [name for name in source_columns if name in table.c] if not columns: print(f"skip {table_name}: no shared columns") return 0, 0 selected = ", ".join(_quote_sqlite_identifier(name) for name in columns) query = f"SELECT {selected} FROM {_quote_sqlite_identifier(table_name)}" cursor = source.execute(query) copied = 0 async with engine.begin() as conn: while True: batch = cursor.fetchmany(batch_size) if not batch: break payload = _build_rows(table_name=table_name, table=table, rows=batch, columns=columns) if payload: await conn.execute(table.insert(), payload) copied += len(payload) expected = _sqlite_count(source, table_name) return expected, copied async def _verify_counts(source: sqlite3.Connection, engine, metadata: MetaData, table_names: list[str]) -> None: failures: list[str] = [] async with engine.connect() as conn: for table_name in table_names: table = metadata.tables.get(table_name) if table is None: continue source_count = _sqlite_count(source, table_name) target_count = await _target_count(conn, table) if source_count != target_count: failures.append(f"{table_name}: sqlite={source_count}, mysql={target_count}") if failures: raise RuntimeError("Row count verification failed:\n " + "\n ".join(failures)) async def _migrate_data(args: argparse.Namespace, mysql_url: str) -> None: source_path = Path(args.sqlite).resolve() if not source_path.is_file(): raise FileNotFoundError(f"SQLite database not found: {source_path}") source = sqlite3.connect(str(source_path)) source.row_factory = sqlite3.Row try: source_tables = _sqlite_tables(source) table_names = [name for name in BUSINESS_TABLES if name in source_tables] missing = [name for name in BUSINESS_TABLES if name not in source_tables] if missing: print("SQLite tables not present, skipped: " + ", ".join(missing)) engine, metadata = await _reflect_target(mysql_url) try: if args.truncate_target: await _truncate_target(engine, metadata, table_names) else: await _ensure_target_empty(engine, metadata, table_names) for table_name in table_names: expected, copied = await _copy_table( source=source, engine=engine, metadata=metadata, table_name=table_name, batch_size=args.batch_size, ) print(f"{table_name}: copied {copied}/{expected}") await _verify_counts(source, engine, metadata, table_names) finally: await engine.dispose() finally: source.close() def _parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Migrate DeerFlow business tables from SQLite to MySQL") parser.add_argument("--sqlite", required=True, help="Source SQLite deerflow.db path") parser.add_argument("--mysql-url", default=os.getenv("MYSQL_DATABASE_URL", ""), help="Target MySQL SQLAlchemy URL") parser.add_argument("--vendor-dir", default="", help="Optional Python vendor directory containing asyncmy") parser.add_argument("--batch-size", type=int, default=1000) parser.add_argument("--create-database", action="store_true", help="Create the target MySQL database if it does not exist") parser.add_argument("--drop-empty-target-schema", action="store_true", help="Drop known target tables first, but only when all are empty") parser.add_argument("--skip-schema", action="store_true", help="Do not create/update target schema before copying") parser.add_argument("--truncate-target", action="store_true", help="Delete existing target rows before copying") return parser.parse_args() def main() -> None: args = _parse_args() _add_vendor_path(args.vendor_dir) mysql_url = _normalize_mysql_url(args.mysql_url.strip()) if not mysql_url: raise SystemExit("Missing --mysql-url or MYSQL_DATABASE_URL") if args.create_database: asyncio.run(_create_database_if_needed(mysql_url)) if not args.skip_schema: if args.drop_empty_target_schema: asyncio.run(_drop_empty_target_schema(mysql_url, BUSINESS_TABLES + ["alembic_version"])) asyncio.run(_create_schema(mysql_url)) _run_alembic(mysql_url) asyncio.run(_migrate_data(args, mysql_url)) print("Business table migration completed successfully.") if __name__ == "__main__": main()