434 lines
15 KiB
Python
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()
|