From ebb00fcacee754a8d74e39707b931d96114c8584 Mon Sep 17 00:00:00 2001 From: Christian Krakau-Louis Date: Sun, 31 May 2026 01:16:23 +0200 Subject: [PATCH] fix(database): add SQL backend migration support to 0.30.2 --- app/database.py | 283 +++++++++++++++++++++------------------ app/utils/db_migrate.py | 212 +++++++++++++++++++++++++++++ requirements.txt | 73 +++++----- tests/test_db_migrate.py | 102 ++++++++++++++ 4 files changed, 503 insertions(+), 167 deletions(-) create mode 100644 app/utils/db_migrate.py create mode 100644 tests/test_db_migrate.py diff --git a/app/database.py b/app/database.py index 8a73f71b..e762222a 100644 --- a/app/database.py +++ b/app/database.py @@ -1,131 +1,152 @@ -# app/database.py - -import logging -import os -from collections.abc import Generator -from typing import Any - -from sqlalchemy import create_engine, exc -from sqlalchemy.engine.url import make_url -from sqlalchemy.orm import Session, declarative_base, sessionmaker - -from app.config import settings - -logger = logging.getLogger(__name__) - -Base = declarative_base() - -# Parse the DATABASE_URL -DB_URL = settings.database_url -engine = create_engine(DB_URL, connect_args={"check_same_thread": False}) -SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) - - -def init_db() -> None: - """ - Ensures the SQLite database file and its parent directory exist (if using sqlite). - Then runs Base.metadata.create_all(bind=engine) to initialize tables. - Logs a message if a new SQLite DB file is created. - """ - # 1. Parse the DB URL to see if it's sqlite - url = make_url(DB_URL) - if url.get_backend_name() == "sqlite": - # 2. Extract the database path from the URL - database_path = url.database # e.g. "/workdir/db/database.db" or ":memory:" - - if database_path != ":memory:": - # 3. Ensure directory exists - db_dir = os.path.dirname(database_path) - if db_dir and not os.path.exists(db_dir): - logger.info(f"Creating directory for SQLite DB: {db_dir}") - os.makedirs(db_dir, exist_ok=True) - - # 4. If the file does not exist, create an empty one - if not os.path.exists(database_path): - logger.info(f"Creating new SQLite database file at {database_path}") - open(database_path, "a").close() - - # 5. Now create tables if they don't exist yet - try: - Base.metadata.create_all(bind=engine) - logger.info("Database initialization complete (tables created if not exist).") - - # 6. Run lightweight schema migrations for existing databases - _run_schema_migrations(engine) - except exc.SQLAlchemyError as e: - logger.error(f"Error initializing database: {e}") - raise - - -def _run_schema_migrations(engine: Any) -> None: - """ - Apply lightweight schema migrations for columns added after the initial release. - Each migration is idempotent and safe to run multiple times. - """ - from sqlalchemy import inspect, text - - inspector = inspect(engine) - - # Migration: Add 'detail' column to processing_logs (added for verbose worker log output) - if "processing_logs" in inspector.get_table_names(): - columns = [col["name"] for col in inspector.get_columns("processing_logs")] - if "detail" not in columns: - logger.info("Migrating processing_logs: adding 'detail' column") - with engine.begin() as conn: - conn.execute(text("ALTER TABLE processing_logs ADD COLUMN detail TEXT")) - logger.info("Migration complete: 'detail' column added to processing_logs") - - # Migration: Add file path columns to files table - if "files" in inspector.get_table_names(): - columns = [col["name"] for col in inspector.get_columns("files")] - if "original_file_path" not in columns: - logger.info("Migrating files: adding 'original_file_path' column") - with engine.begin() as conn: - conn.execute(text("ALTER TABLE files ADD COLUMN original_file_path VARCHAR")) - logger.info("Migration complete: 'original_file_path' column added to files") - - if "processed_file_path" not in columns: - logger.info("Migrating files: adding 'processed_file_path' column") - with engine.begin() as conn: - conn.execute(text("ALTER TABLE files ADD COLUMN processed_file_path VARCHAR")) - logger.info("Migration complete: 'processed_file_path' column added to files") - - # Migration: Add deduplication columns to files table - if "is_duplicate" not in columns: - logger.info("Migrating files: adding 'is_duplicate' column") - with engine.begin() as conn: - conn.execute(text("ALTER TABLE files ADD COLUMN is_duplicate BOOLEAN DEFAULT FALSE NOT NULL")) - logger.info("Migration complete: 'is_duplicate' column added to files") - - if "duplicate_of_id" not in columns: - logger.info("Migrating files: adding 'duplicate_of_id' column") - with engine.begin() as conn: - conn.execute(text("ALTER TABLE files ADD COLUMN duplicate_of_id INTEGER")) - logger.info("Migration complete: 'duplicate_of_id' column added to files") - - # Migration: Drop unique index on filehash to allow duplicate records - try: - indexes = inspector.get_indexes("files") - unique_filehash_indexes = [ - index for index in indexes if index.get("unique") and "filehash" in index.get("column_names", []) - ] - if unique_filehash_indexes: - logger.info("Migrating files: dropping unique index on 'filehash'") - with engine.begin() as conn: - for index in unique_filehash_indexes: - conn.execute(text(f"DROP INDEX IF EXISTS {index['name']}")) - logger.info("Migration complete: unique index on 'filehash' removed") - except Exception as exc: - logger.warning(f"Skipping filehash unique index drop: {exc}") - - -def get_db() -> Generator[Session, None, None]: - """ - Dependency for FastAPI routes or general DB usage. - Yields a SQLAlchemy session, and closes it upon exit. - """ - db = SessionLocal() - try: - yield db - finally: - db.close() +# app/database.py + +import logging +import os +from collections.abc import Generator +from typing import Any + +from sqlalchemy import create_engine, exc +from sqlalchemy.engine.url import make_url +from sqlalchemy.orm import Session, declarative_base, sessionmaker +from sqlalchemy.pool import NullPool, QueuePool + +from app.config import settings + +logger = logging.getLogger(__name__) + +Base = declarative_base() + +DB_URL = settings.database_url +_parsed_url = make_url(DB_URL) + +_connect_args: dict[str, Any] = {} +_engine_kwargs: dict[str, Any] = {"pool_pre_ping": True} + +if _parsed_url.get_backend_name() == "sqlite": + _connect_args["check_same_thread"] = False + _engine_kwargs["poolclass"] = NullPool +else: + _engine_kwargs.update( + { + "poolclass": QueuePool, + "pool_size": int(os.getenv("DB_POOL_SIZE", "10")), + "max_overflow": int(os.getenv("DB_MAX_OVERFLOW", "20")), + "pool_timeout": int(os.getenv("DB_POOL_TIMEOUT", "30")), + "pool_recycle": int(os.getenv("DB_POOL_RECYCLE", "1800")), + } + ) + +engine = create_engine(DB_URL, connect_args=_connect_args, **_engine_kwargs) +SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) + + +def init_db() -> None: + """ + Ensures the SQLite database file and its parent directory exist (if using sqlite). + Then runs Base.metadata.create_all(bind=engine) to initialize tables. + Logs a message if a new SQLite DB file is created. + """ + # 1. Parse the DB URL to see if it's sqlite + url = make_url(DB_URL) + if url.get_backend_name() == "sqlite": + # 2. Extract the database path from the URL + database_path = url.database # e.g. "/workdir/db/database.db" or ":memory:" + + if database_path != ":memory:": + # 3. Ensure directory exists + db_dir = os.path.dirname(database_path) + if db_dir and not os.path.exists(db_dir): + logger.info(f"Creating directory for SQLite DB: {db_dir}") + os.makedirs(db_dir, exist_ok=True) + + # 4. If the file does not exist, create an empty one + if not os.path.exists(database_path): + logger.info(f"Creating new SQLite database file at {database_path}") + open(database_path, "a").close() + + # 5. Now create tables if they don't exist yet + try: + import app.models # noqa: F401 + + Base.metadata.create_all(bind=engine) + logger.info("Database initialization complete (tables created if not exist).") + + # 6. Run lightweight schema migrations for existing databases + _run_schema_migrations(engine) + except exc.SQLAlchemyError as e: + logger.error(f"Error initializing database: {e}") + raise + + +def _run_schema_migrations(engine: Any) -> None: + """ + Apply lightweight schema migrations for columns added after the initial release. + Each migration is idempotent and safe to run multiple times. + """ + from sqlalchemy import inspect, text + + inspector = inspect(engine) + + # Migration: Add 'detail' column to processing_logs (added for verbose worker log output) + if "processing_logs" in inspector.get_table_names(): + columns = [col["name"] for col in inspector.get_columns("processing_logs")] + if "detail" not in columns: + logger.info("Migrating processing_logs: adding 'detail' column") + with engine.begin() as conn: + conn.execute(text("ALTER TABLE processing_logs ADD COLUMN detail TEXT")) + logger.info("Migration complete: 'detail' column added to processing_logs") + + # Migration: Add file path columns to files table + if "files" in inspector.get_table_names(): + columns = [col["name"] for col in inspector.get_columns("files")] + if "original_file_path" not in columns: + logger.info("Migrating files: adding 'original_file_path' column") + with engine.begin() as conn: + conn.execute(text("ALTER TABLE files ADD COLUMN original_file_path VARCHAR")) + logger.info("Migration complete: 'original_file_path' column added to files") + + if "processed_file_path" not in columns: + logger.info("Migrating files: adding 'processed_file_path' column") + with engine.begin() as conn: + conn.execute(text("ALTER TABLE files ADD COLUMN processed_file_path VARCHAR")) + logger.info("Migration complete: 'processed_file_path' column added to files") + + # Migration: Add deduplication columns to files table + if "is_duplicate" not in columns: + logger.info("Migrating files: adding 'is_duplicate' column") + with engine.begin() as conn: + conn.execute(text("ALTER TABLE files ADD COLUMN is_duplicate BOOLEAN DEFAULT FALSE NOT NULL")) + logger.info("Migration complete: 'is_duplicate' column added to files") + + if "duplicate_of_id" not in columns: + logger.info("Migrating files: adding 'duplicate_of_id' column") + with engine.begin() as conn: + conn.execute(text("ALTER TABLE files ADD COLUMN duplicate_of_id INTEGER")) + logger.info("Migration complete: 'duplicate_of_id' column added to files") + + # Migration: Drop unique index on filehash to allow duplicate records + try: + indexes = inspector.get_indexes("files") + unique_filehash_indexes = [ + index for index in indexes if index.get("unique") and "filehash" in index.get("column_names", []) + ] + if unique_filehash_indexes: + logger.info("Migrating files: dropping unique index on 'filehash'") + with engine.begin() as conn: + for index in unique_filehash_indexes: + conn.execute(text(f"DROP INDEX IF EXISTS {index['name']}")) + logger.info("Migration complete: unique index on 'filehash' removed") + except Exception as exc: + logger.warning(f"Skipping filehash unique index drop: {exc}") + + +def get_db() -> Generator[Session, None, None]: + """ + Dependency for FastAPI routes or general DB usage. + Yields a SQLAlchemy session, and closes it upon exit. + """ + db = SessionLocal() + try: + yield db + finally: + db.close() diff --git a/app/utils/db_migrate.py b/app/utils/db_migrate.py new file mode 100644 index 00000000..88329cf1 --- /dev/null +++ b/app/utils/db_migrate.py @@ -0,0 +1,212 @@ +"""Database migration utility for copying SQLite state to another SQL backend. + +The production 0.30.2 deployment stores state in SQLite. This module provides +a conservative data-copy path for moving that state into PostgreSQL while +preserving tables that may have been created by newer application builds. +""" + +from __future__ import annotations + +import argparse +import json +import logging +import re +from typing import Any + +from sqlalchemy import MetaData, column, create_engine, func, inspect, select, table, text +from sqlalchemy.engine import Engine +from sqlalchemy.engine.url import make_url + +logger = logging.getLogger(__name__) + +_SKIP_TABLES = {"alembic_version", "sqlite_sequence"} +_TABLE_ORDER = [ + "documents", + "files", + "file_processing_steps", + "processing_logs", + "application_settings", +] +_SAFE_IDENTIFIER = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$") + + +def _make_engine(url: str) -> Engine: + """Create a SQLAlchemy engine with SQLite-only connection arguments.""" + parsed = make_url(url) + connect_args: dict[str, Any] = {} + if parsed.get_backend_name() == "sqlite": + connect_args["check_same_thread"] = False + return create_engine(url, connect_args=connect_args, pool_pre_ping=True) + + +def _safe_table_names(inspector: Any) -> list[str]: + """Return source table names, excluding internal and unsafe identifiers.""" + result = [] + for name in inspector.get_table_names(): + if name in _SKIP_TABLES: + continue + if not _SAFE_IDENTIFIER.match(name): + logger.warning("Skipping table with unsafe name: %s", name) + continue + result.append(name) + return result + + +def _ordered_tables(inspector: Any) -> list[str]: + """Return table names in stable parent-first order.""" + existing = set(_safe_table_names(inspector)) + ordered = [name for name in _TABLE_ORDER if name in existing] + ordered.extend(sorted(existing - set(ordered))) + return ordered + + +def preview_migration(source_url: str) -> dict[str, Any]: + """Preview source tables and row counts without copying data.""" + try: + engine = _make_engine(source_url) + inspector = inspect(engine) + tables = [] + total_rows = 0 + + with engine.connect() as conn: + for table_name in _ordered_tables(inspector): + row = conn.execute(select(func.count()).select_from(table(table_name))).fetchone() + count = row[0] if row else 0 + tables.append({"name": table_name, "row_count": count}) + total_rows += count + + engine.dispose() + return {"success": True, "tables": tables, "total_rows": total_rows} + except Exception as exc: + logger.error("Migration preview failed: %s", exc) + return {"success": False, "error": str(exc), "tables": [], "total_rows": 0} + + +def _create_target_schema_from_source(src_engine: Engine, tgt_engine: Engine) -> MetaData: + """Reflect the source schema and create equivalent target tables.""" + source_metadata = MetaData() + source_metadata.reflect(bind=src_engine, views=False) + for table_name in list(source_metadata.tables): + if table_name in _SKIP_TABLES or not _SAFE_IDENTIFIER.match(table_name): + source_metadata.remove(source_metadata.tables[table_name]) + source_metadata.create_all(bind=tgt_engine) + return source_metadata + + +def _reset_postgres_sequences(engine: Engine, table_names: list[str]) -> None: + """Move PostgreSQL serial sequences past copied explicit primary keys.""" + if engine.dialect.name != "postgresql": + return + + with engine.begin() as conn: + for table_name in table_names: + if not _SAFE_IDENTIFIER.match(table_name): + continue + sequence = conn.execute( + text("SELECT pg_get_serial_sequence(:table_name, 'id')"), {"table_name": table_name} + ).scalar() + if not sequence: + continue + reflected_table = table(table_name, column("id")) + max_id = conn.execute(select(func.max(reflected_table.c.id))).scalar() + conn.execute( + text("SELECT setval(:sequence_name, :value, :is_called)"), + {"sequence_name": sequence, "value": max_id or 1, "is_called": max_id is not None}, + ) + + +def migrate_data( + source_url: str, + target_url: str, + *, + batch_size: int = 500, + progress_callback: Any | None = None, +) -> dict[str, Any]: + """Copy all supported tables from *source_url* to *target_url*.""" + errors: list[str] = [] + tables_copied = 0 + rows_copied = 0 + + try: + src_engine = _make_engine(source_url) + tgt_engine = _make_engine(target_url) + source_metadata = _create_target_schema_from_source(src_engine, tgt_engine) + source_inspector = inspect(src_engine) + table_names = _ordered_tables(source_inspector) + + for table_name in table_names: + try: + src_table = source_metadata.tables.get(table_name) + if src_table is None: + continue + + rows = [] + with src_engine.connect() as conn: + for row in conn.execute(src_table.select()): + rows.append(dict(row._mapping)) + + if not rows: + tables_copied += 1 + continue + + target_metadata = MetaData() + target_metadata.reflect(bind=tgt_engine, only=[table_name]) + target_table = target_metadata.tables[table_name] + total_for_table = len(rows) + + with tgt_engine.begin() as conn: + for offset in range(0, total_for_table, batch_size): + batch = rows[offset : offset + batch_size] + conn.execute(target_table.insert(), batch) + rows_copied += len(batch) + if progress_callback: + progress_callback(table_name, min(offset + batch_size, total_for_table), total_for_table) + + tables_copied += 1 + logger.info("Copied %s rows from %s", total_for_table, table_name) + except Exception as exc: + message = f"Error copying table {table_name}: {exc}" + logger.error(message) + errors.append(message) + + try: + _reset_postgres_sequences(tgt_engine, table_names) + except Exception as exc: + message = f"Failed to reset PostgreSQL sequences: {exc}" + logger.error(message) + errors.append(message) + + src_engine.dispose() + tgt_engine.dispose() + return { + "success": len(errors) == 0, + "tables_copied": tables_copied, + "rows_copied": rows_copied, + "errors": errors, + } + except Exception as exc: + logger.error("Migration failed: %s", exc) + return { + "success": False, + "tables_copied": tables_copied, + "rows_copied": rows_copied, + "errors": errors + [str(exc)], + } + + +def _main() -> None: + parser = argparse.ArgumentParser(description="Copy DocuElevate database rows between SQLAlchemy URLs.") + parser.add_argument("source_url") + parser.add_argument("target_url") + parser.add_argument("--batch-size", type=int, default=500) + args = parser.parse_args() + + logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") + result = migrate_data(args.source_url, args.target_url, batch_size=args.batch_size) + print(json.dumps(result, indent=2, sort_keys=True)) + if not result["success"]: + raise SystemExit(1) + + +if __name__ == "__main__": + _main() diff --git a/requirements.txt b/requirements.txt index b1389fd6..d9f70f55 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,37 +1,38 @@ -fastapi[all] # Web framework with all extras -uvicorn # ASGI server -celery # Task queue -redis # Message broker for Celery -sqlalchemy # Database ORM -pydantic # Data validation -cryptography>=41.0.0 # Encryption for sensitive settings in database -openai # GPT integration for metadata extraction -pypdf>=3.9.0 # PDF processing for text extraction, metadata editing and rotation (upgraded from PyPDF2 to fix CVE-2023-36464) -requests # HTTP client -puremagic>=1.25,<2.0 # File type detection (pure Python) -filetype>=1.2.0,<2.0 # File type detection fallback (pure Python) -dropbox>=11.36.0 # Dropbox integration -azure-ai-documentintelligence # Azure OCR service -authlib>=1.6.5 # Authentication - fixed security vulnerabilities (GHSA-xxx) -python-dotenv # Environment variables -starlette>=0.49.1 # ASGI toolkit (used by FastAPI) - fixed DoS vulnerability -alembic # Database migrations -slowapi>=0.1.9 # Rate limiting middleware for FastAPI - -# Google Drive API -google-api-python-client>=2.79.0 -google-auth>=2.22.0 -google-auth-oauthlib>=1.0.0 - -# OneDrive/Microsoft Graph API -msgraph-core>=1.0.0 -msal>=1.20.0 - -# AWS S3 -boto3>=1.28.0 - -# SFTP -paramiko>=3.4.0 # SSH/SFTP implementation for Python (LGPL license) - -# Notification service +fastapi[all] # Web framework with all extras +uvicorn # ASGI server +celery # Task queue +redis # Message broker for Celery +sqlalchemy # Database ORM +psycopg[binary]>=3.2,<4.0 # PostgreSQL driver for external HA database deployments +pydantic # Data validation +cryptography>=41.0.0 # Encryption for sensitive settings in database +openai # GPT integration for metadata extraction +pypdf>=3.9.0 # PDF processing for text extraction, metadata editing and rotation (upgraded from PyPDF2 to fix CVE-2023-36464) +requests # HTTP client +puremagic>=1.25,<2.0 # File type detection (pure Python) +filetype>=1.2.0,<2.0 # File type detection fallback (pure Python) +dropbox>=11.36.0 # Dropbox integration +azure-ai-documentintelligence # Azure OCR service +authlib>=1.6.5 # Authentication - fixed security vulnerabilities (GHSA-xxx) +python-dotenv # Environment variables +starlette>=0.49.1 # ASGI toolkit (used by FastAPI) - fixed DoS vulnerability +alembic # Database migrations +slowapi>=0.1.9 # Rate limiting middleware for FastAPI + +# Google Drive API +google-api-python-client>=2.79.0 +google-auth>=2.22.0 +google-auth-oauthlib>=1.0.0 + +# OneDrive/Microsoft Graph API +msgraph-core>=1.0.0 +msal>=1.20.0 + +# AWS S3 +boto3>=1.28.0 + +# SFTP +paramiko>=3.4.0 # SSH/SFTP implementation for Python (LGPL license) + +# Notification service apprise>=1.4.0 diff --git a/tests/test_db_migrate.py b/tests/test_db_migrate.py new file mode 100644 index 00000000..1f889138 --- /dev/null +++ b/tests/test_db_migrate.py @@ -0,0 +1,102 @@ +"""Tests for the SQLite-to-SQL database migration helper.""" + +from unittest.mock import MagicMock, patch + +import pytest +from sqlalchemy import create_engine, inspect, text +from sqlalchemy.pool import StaticPool + +from app.utils.db_migrate import _make_engine, _ordered_tables, migrate_data, preview_migration + + +@pytest.mark.unit +class TestMakeEngine: + def test_sqlite_engine_is_created(self): + engine = _make_engine("sqlite:///:memory:") + assert engine is not None + engine.dispose() + + def test_non_sqlite_engine_does_not_use_sqlite_connect_args(self): + engine = _make_engine("postgresql+psycopg://user:pass@localhost:5432/docuelevate") + assert engine is not None + engine.dispose() + + +@pytest.mark.unit +class TestOrderedTables: + def test_known_tables_are_ordered_before_unknown_tables(self): + inspector = MagicMock() + inspector.get_table_names.return_value = ["z_table", "files", "documents", "sqlite_sequence"] + + assert _ordered_tables(inspector) == ["documents", "files", "z_table"] + + def test_unsafe_table_names_are_skipped(self): + inspector = MagicMock() + inspector.get_table_names.return_value = ["files", "bad-table"] + + assert _ordered_tables(inspector) == ["files"] + + +@pytest.mark.unit +class TestPreviewMigration: + def test_preview_returns_row_counts(self, tmp_path): + db_path = tmp_path / "source.db" + engine = create_engine(f"sqlite:///{db_path}") + with engine.begin() as conn: + conn.execute(text("CREATE TABLE files (id INTEGER PRIMARY KEY, local_filename VARCHAR NOT NULL)")) + conn.execute(text("INSERT INTO files (local_filename) VALUES ('a.pdf'), ('b.pdf')")) + + result = preview_migration(f"sqlite:///{db_path}") + + assert result["success"] is True + assert result["total_rows"] == 2 + assert result["tables"] == [{"name": "files", "row_count": 2}] + + +@pytest.mark.unit +class TestMigrateData: + def test_migrates_reflected_schema_and_rows(self, tmp_path): + source_path = tmp_path / "source.db" + target_path = tmp_path / "target.db" + source = create_engine(f"sqlite:///{source_path}") + + with source.begin() as conn: + conn.execute(text("CREATE TABLE files (id INTEGER PRIMARY KEY, local_filename VARCHAR NOT NULL)")) + conn.execute(text("CREATE TABLE future_table (id INTEGER PRIMARY KEY, value VARCHAR)")) + conn.execute(text("INSERT INTO files (id, local_filename) VALUES (7, 'stable.pdf')")) + conn.execute(text("INSERT INTO future_table (id, value) VALUES (1, 'kept')")) + + result = migrate_data(f"sqlite:///{source_path}", f"sqlite:///{target_path}") + + assert result["success"] is True + assert result["rows_copied"] == 2 + + target = create_engine(f"sqlite:///{target_path}") + inspector = inspect(target) + assert "files" in inspector.get_table_names() + assert "future_table" in inspector.get_table_names() + with target.connect() as conn: + assert conn.execute(text("SELECT local_filename FROM files WHERE id = 7")).scalar_one() == "stable.pdf" + assert conn.execute(text("SELECT value FROM future_table WHERE id = 1")).scalar_one() == "kept" + + def test_migration_reports_global_errors(self): + with patch("app.utils.db_migrate._make_engine", side_effect=RuntimeError("boom")): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + assert result["success"] is False + assert "boom" in result["errors"][0] + + def test_migration_progress_callback_is_called(self): + source = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + target = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + with source.begin() as conn: + conn.execute(text("CREATE TABLE files (id INTEGER PRIMARY KEY, local_filename VARCHAR NOT NULL)")) + conn.execute(text("INSERT INTO files (id, local_filename) VALUES (1, 'a.pdf'), (2, 'b.pdf')")) + + callback = MagicMock() + with patch("app.utils.db_migrate._make_engine") as make_engine: + make_engine.side_effect = [source, target] + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:", batch_size=1, progress_callback=callback) + + assert result["success"] is True + assert callback.call_count == 2