feat(database): add database configuration wizard and migration tool

Add a guided database configuration wizard and a data migration tool that
allows users to:
- Build database connection strings through a step-by-step UI
- Test database connections before applying
- Preview and execute data migrations from SQLite to PostgreSQL/MySQL
- Copy to clipboard for easy .env file updates

New files:
- app/utils/db_wizard.py — connection string builder, parser, and tester
- app/utils/db_migrate.py — table-by-table data migration utility
- app/api/database.py — REST API endpoints for wizard operations
- app/views/db_wizard.py — view route for the wizard page
- frontend/templates/db_wizard.html — multi-tab wizard UI
- tests/test_db_wizard.py — unit tests for db_wizard utilities
- tests/test_db_migrate.py — unit tests for db_migrate utilities
- tests/test_db_wizard_api.py — integration tests for API and views

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-03-05 22:10:14 +00:00
parent 0f408f67b4
commit f6fcaaeccc
10 changed files with 1848 additions and 0 deletions
+248
View File
@@ -0,0 +1,248 @@
"""
Database migration utility for transferring data between databases.
Copies all table rows from a *source* SQLAlchemy database to a *target*
database. This is designed for the common scenario of migrating from the
built-in SQLite database to an external PostgreSQL / MySQL instance.
The utility:
1. Creates the schema in the target via ``Base.metadata.create_all``.
2. Copies rows table-by-table in dependency order.
3. Stamps the Alembic version in the target to ``head``.
"""
import logging
from typing import Any
from sqlalchemy import MetaData, create_engine, inspect, text
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import make_url
from sqlalchemy.orm import sessionmaker
logger = logging.getLogger(__name__)
# Tables to skip during migration (Alembic manages its own state).
_SKIP_TABLES = {"alembic_version"}
# Ordered list — parent tables first to respect foreign-key constraints.
_TABLE_ORDER = [
"documents",
"files",
"file_processing_steps",
"processing_logs",
"application_settings",
"settings_audit_log",
"saved_searches",
"webhook_configs",
]
def _make_engine(url: str) -> Engine:
"""Create a SQLAlchemy engine from *url* with sensible defaults."""
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)
def _ordered_tables(inspector: Any) -> list[str]:
"""Return table names in safe insertion order.
Tables listed in ``_TABLE_ORDER`` come first (in that order); any
remaining tables are appended alphabetically.
"""
existing = set(inspector.get_table_names())
ordered: list[str] = []
for name in _TABLE_ORDER:
if name in existing and name not in _SKIP_TABLES:
ordered.append(name)
for name in sorted(existing):
if name not in ordered and name not in _SKIP_TABLES:
ordered.append(name)
return ordered
def preview_migration(source_url: str) -> dict[str, Any]:
"""Preview what a migration would do without actually copying data.
Args:
source_url: Connection string for the source database.
Returns:
Dict with ``tables`` (list of dicts with ``name`` and ``row_count``)
and ``total_rows``.
"""
try:
src_engine = _make_engine(source_url)
src_inspector = inspect(src_engine)
tables = _ordered_tables(src_inspector)
result: list[dict[str, Any]] = []
total = 0
with src_engine.connect() as conn:
for table_name in tables:
row = conn.execute(text(f'SELECT COUNT(*) FROM "{table_name}"')).fetchone() # noqa: S608
count = row[0] if row else 0
result.append({"name": table_name, "row_count": count})
total += count
src_engine.dispose()
return {"tables": result, "total_rows": total, "success": True}
except Exception as exc:
logger.error(f"Migration preview failed: {exc}")
return {"success": False, "error": str(exc), "tables": [], "total_rows": 0}
def migrate_data(
source_url: str,
target_url: str,
*,
batch_size: int = 500,
progress_callback: Any | None = None,
) -> dict[str, Any]:
"""Copy all data from *source_url* to *target_url*.
The target schema is created automatically from the application models.
Alembic is stamped to ``head`` in the target after a successful copy.
Args:
source_url: SQLAlchemy connection string for the source DB.
target_url: SQLAlchemy connection string for the target DB.
batch_size: Number of rows to insert per batch.
progress_callback: Optional ``callable(table_name, copied, total)``
invoked after each batch.
Returns:
Dict with ``success`` (bool), ``tables_copied`` (int),
``rows_copied`` (int), and ``errors`` (list of str).
"""
errors: list[str] = []
tables_copied = 0
rows_copied = 0
try:
src_engine = _make_engine(source_url)
tgt_engine = _make_engine(target_url)
# ------------------------------------------------------------------
# 1. Create schema in target from application models
# ------------------------------------------------------------------
from app.database import Base # local import to avoid circular deps
Base.metadata.create_all(bind=tgt_engine)
logger.info("Target schema created from application models.")
# ------------------------------------------------------------------
# 2. Reflect source schema & determine copy order
# ------------------------------------------------------------------
src_meta = MetaData()
src_meta.reflect(bind=src_engine)
src_inspector = inspect(src_engine)
table_names = _ordered_tables(src_inspector)
SrcSession = sessionmaker(bind=src_engine)
TgtSession = sessionmaker(bind=tgt_engine)
# ------------------------------------------------------------------
# 3. Copy data table-by-table
# ------------------------------------------------------------------
for table_name in table_names:
try:
src_session = SrcSession()
tgt_session = TgtSession()
src_table = src_meta.tables.get(table_name)
if src_table is None:
continue
# Read all rows from source
rows = src_session.execute(src_table.select()).fetchall()
column_names = [c.name for c in src_table.columns]
if not rows:
logger.info(f"Skipping empty table: {table_name}")
tables_copied += 1
src_session.close()
tgt_session.close()
continue
# Reflect the target table to insert into
tgt_meta = MetaData()
tgt_meta.reflect(bind=tgt_engine, only=[table_name])
tgt_table = tgt_meta.tables.get(table_name)
if tgt_table is None:
errors.append(f"Target table {table_name} not found after schema creation")
src_session.close()
tgt_session.close()
continue
# Batch insert
total_for_table = len(rows)
for i in range(0, total_for_table, batch_size):
batch = rows[i : i + batch_size]
insert_data = [dict(zip(column_names, row, strict=False)) for row in batch]
tgt_session.execute(tgt_table.insert(), insert_data)
tgt_session.commit()
rows_copied += len(batch)
if progress_callback:
progress_callback(table_name, min(i + batch_size, total_for_table), total_for_table)
tables_copied += 1
logger.info(f"Copied {total_for_table} rows from {table_name}")
src_session.close()
tgt_session.close()
except Exception as exc:
msg = f"Error copying table {table_name}: {exc}"
logger.error(msg)
errors.append(msg)
# ------------------------------------------------------------------
# 4. Stamp Alembic to head in the target
# ------------------------------------------------------------------
try:
_stamp_alembic_head(tgt_engine)
logger.info("Alembic version stamped to head in target database.")
except Exception as exc:
msg = f"Failed to stamp Alembic version: {exc}"
logger.error(msg)
errors.append(msg)
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(f"Migration failed: {exc}")
return {
"success": False,
"tables_copied": tables_copied,
"rows_copied": rows_copied,
"errors": errors + [str(exc)],
}
def _stamp_alembic_head(engine: Engine) -> None:
"""Stamp the Alembic version table to ``head`` in the given engine."""
from pathlib import Path
from alembic import command
from alembic.config import Config
migrations_dir = str(Path(__file__).resolve().parent.parent.parent / "migrations")
alembic_cfg = Config()
alembic_cfg.set_main_option("script_location", migrations_dir)
alembic_cfg.set_main_option("sqlalchemy.url", "")
with engine.begin() as connection:
alembic_cfg.attributes["connection"] = connection
command.stamp(alembic_cfg, "head")
+255
View File
@@ -0,0 +1,255 @@
"""
Database configuration wizard utilities.
Provides helpers for building, validating, and testing database connection
strings. Used by both the interactive wizard UI and the REST API.
"""
import logging
from typing import Any
from sqlalchemy import create_engine, text
from sqlalchemy.engine.url import make_url
logger = logging.getLogger(__name__)
# Supported database backends with human-readable labels and defaults.
SUPPORTED_BACKENDS: list[dict[str, Any]] = [
{
"id": "sqlite",
"label": "SQLite (Development)",
"driver": "",
"default_port": None,
"description": "File-based database. Best for development and single-user setups.",
"requires_host": False,
},
{
"id": "postgresql",
"label": "PostgreSQL (Recommended for Production)",
"driver": "",
"default_port": 5432,
"description": "Robust, full-featured database. Recommended for production.",
"requires_host": True,
},
{
"id": "mysql",
"label": "MySQL / MariaDB",
"driver": "pymysql",
"default_port": 3306,
"description": "Popular open-source database. Requires pymysql driver.",
"requires_host": True,
},
]
def get_supported_backends() -> list[dict[str, Any]]:
"""Return the list of supported database backends with metadata.
Returns:
List of backend descriptor dicts.
"""
return SUPPORTED_BACKENDS
def build_connection_string(
backend: str,
host: str = "",
port: int | None = None,
database: str = "",
username: str = "",
password: str = "",
ssl_mode: str = "",
extra_options: str = "",
sqlite_path: str = "",
) -> str:
"""Build a SQLAlchemy connection string from individual components.
Args:
backend: Database backend identifier (``sqlite``, ``postgresql``, ``mysql``).
host: Database server hostname or IP.
port: Database server port (uses backend default when ``None``).
database: Database / schema name.
username: Authentication username.
password: Authentication password.
ssl_mode: SSL mode (e.g. ``require``, ``verify-full``). PostgreSQL only.
extra_options: Additional query-string options appended to the URL.
sqlite_path: File path for SQLite databases.
Returns:
A SQLAlchemy-compatible connection URL string.
Raises:
ValueError: If required fields are missing for the chosen backend.
"""
if backend == "sqlite":
path = sqlite_path.strip() if sqlite_path else "./app/database.db"
return f"sqlite:///{path}"
# Resolve driver prefix
backend_info = next((b for b in SUPPORTED_BACKENDS if b["id"] == backend), None)
if backend_info is None:
raise ValueError(f"Unsupported backend: {backend}")
if not host:
raise ValueError("Host is required for non-SQLite backends")
if not database:
raise ValueError("Database name is required for non-SQLite backends")
if not username:
raise ValueError("Username is required for non-SQLite backends")
driver_suffix = f"+{backend_info['driver']}" if backend_info["driver"] else ""
scheme = f"{backend}{driver_suffix}"
resolved_port = port if port else backend_info["default_port"]
# Build query parameters
params: list[str] = []
if ssl_mode:
params.append(f"sslmode={ssl_mode}")
if extra_options:
params.append(extra_options)
if backend == "mysql" and "charset=" not in extra_options:
params.append("charset=utf8mb4")
query_string = "&".join(params)
# Construct URL
auth = username
if password:
auth = f"{username}:{password}"
url = f"{scheme}://{auth}@{host}:{resolved_port}/{database}"
if query_string:
url = f"{url}?{query_string}"
return url
def parse_connection_string(url: str) -> dict[str, Any]:
"""Parse a SQLAlchemy connection string into its components.
Args:
url: A SQLAlchemy database URL string.
Returns:
Dict with keys: ``backend``, ``host``, ``port``, ``database``,
``username``, ``password``, ``ssl_mode``, ``is_sqlite``.
"""
try:
parsed = make_url(url)
backend_name = parsed.get_backend_name()
return {
"backend": backend_name,
"host": parsed.host or "",
"port": parsed.port,
"database": parsed.database or "",
"username": parsed.username or "",
"password": parsed.password or "",
"ssl_mode": "",
"is_sqlite": backend_name == "sqlite",
"valid": True,
}
except Exception as exc:
logger.warning(f"Failed to parse connection string: {exc}")
return {"valid": False, "error": str(exc)}
def test_connection(url: str, timeout: int = 10) -> dict[str, Any]:
"""Attempt to connect to a database and return status information.
The function creates a short-lived engine, executes a simple ``SELECT 1``
query, and disposes the engine. It does **not** modify any global state.
Args:
url: SQLAlchemy database URL to test.
timeout: Connection timeout in seconds.
Returns:
Dict with ``success`` (bool), ``message`` (str), and optional
``server_version`` (str).
"""
try:
parsed = make_url(url)
backend = parsed.get_backend_name()
connect_args: dict[str, Any] = {}
kwargs: dict[str, Any] = {"pool_pre_ping": True}
if backend == "sqlite":
connect_args["check_same_thread"] = False
else:
kwargs["pool_timeout"] = timeout
test_engine = create_engine(
url,
connect_args=connect_args,
**kwargs,
)
with test_engine.connect() as conn:
result = conn.execute(text("SELECT 1"))
result.fetchone()
# Try to fetch server version for informational display
server_version = _get_server_version(conn, backend)
test_engine.dispose()
return {
"success": True,
"message": "Connection successful",
"backend": backend,
"server_version": server_version,
}
except Exception as exc:
logger.warning(f"Connection test failed: {exc}")
return {
"success": False,
"message": str(exc),
"backend": "",
"server_version": "",
}
def _get_server_version(conn: Any, backend: str) -> str:
"""Retrieve a human-readable server version string.
Args:
conn: An active SQLAlchemy connection.
backend: Backend identifier (``sqlite``, ``postgresql``, ``mysql``).
Returns:
Server version string, or empty string on failure.
"""
try:
if backend == "postgresql":
row = conn.execute(text("SELECT version()")).fetchone()
return str(row[0]) if row else ""
elif backend == "mysql":
row = conn.execute(text("SELECT version()")).fetchone()
return str(row[0]) if row else ""
elif backend == "sqlite":
row = conn.execute(text("SELECT sqlite_version()")).fetchone()
return f"SQLite {row[0]}" if row else ""
except Exception:
logger.debug("Could not retrieve server version")
return ""
def validate_url_format(url: str) -> dict[str, Any]:
"""Validate that a connection string is syntactically correct.
Args:
url: The connection string to validate.
Returns:
Dict with ``valid`` (bool) and optional ``error`` (str).
"""
try:
parsed = make_url(url)
backend = parsed.get_backend_name()
if backend not in ("sqlite", "postgresql", "mysql"):
return {"valid": False, "error": f"Unsupported backend: {backend}"}
return {"valid": True, "backend": backend}
except Exception as exc:
return {"valid": False, "error": str(exc)}