deerflow-code/offline-backend-20260512/backend/scripts/migrate_business_sqlite_to_mysql.py
2026-09-07 18:24:55 +08:00

434 lines
15 KiB
Python

"""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()