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:
@@ -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")
|
||||
@@ -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)}
|
||||
Reference in New Issue
Block a user