fix(database): add SQL backend migration support to 0.30.2

This commit is contained in:
Christian Krakau-Louis
2026-05-31 01:16:23 +02:00
parent 47aa6fc3d4
commit ebb00fcace
4 changed files with 503 additions and 167 deletions
+23 -2
View File
@@ -8,6 +8,7 @@ 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
@@ -15,9 +16,27 @@ 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})
_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)
@@ -47,6 +66,8 @@ def init_db() -> None:
# 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).")
+212
View File
@@ -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()
+1
View File
@@ -3,6 +3,7 @@ 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
+102
View File
@@ -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