From f6fcaaecccf885203ef76add850de424233de9b4 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 5 Mar 2026 22:10:14 +0000 Subject: [PATCH] feat(database): add database configuration wizard and migration tool MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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> --- app/api/__init__.py | 2 + app/api/database.py | 169 +++++++++ app/utils/db_migrate.py | 248 ++++++++++++ app/utils/db_wizard.py | 255 +++++++++++++ app/views/__init__.py | 2 + app/views/db_wizard.py | 28 ++ frontend/templates/db_wizard.html | 606 ++++++++++++++++++++++++++++++ tests/test_db_migrate.py | 149 ++++++++ tests/test_db_wizard.py | 233 ++++++++++++ tests/test_db_wizard_api.py | 156 ++++++++ 10 files changed, 1848 insertions(+) create mode 100644 app/api/database.py create mode 100644 app/utils/db_migrate.py create mode 100644 app/utils/db_wizard.py create mode 100644 app/views/db_wizard.py create mode 100644 frontend/templates/db_wizard.html create mode 100644 tests/test_db_migrate.py create mode 100644 tests/test_db_wizard.py create mode 100644 tests/test_db_wizard_api.py diff --git a/app/api/__init__.py b/app/api/__init__.py index d340418a..f99c3823 100644 --- a/app/api/__init__.py +++ b/app/api/__init__.py @@ -7,6 +7,7 @@ import logging from fastapi import APIRouter from app.api.azure import router as azure_router +from app.api.database import router as database_router from app.api.diagnostic import router as diagnostic_router from app.api.dropbox import router as dropbox_router from app.api.duplicates import router as duplicates_router @@ -52,3 +53,4 @@ router.include_router(saved_searches_router) router.include_router(similarity_router) router.include_router(duplicates_router) router.include_router(webhooks_router) +router.include_router(database_router) diff --git a/app/api/database.py b/app/api/database.py new file mode 100644 index 00000000..501f19b2 --- /dev/null +++ b/app/api/database.py @@ -0,0 +1,169 @@ +""" +API endpoints for the database configuration wizard and migration tool. + +Provides REST endpoints for: +- Testing database connections +- Building connection strings from form components +- Previewing and executing data migrations between databases +""" + +import logging + +from fastapi import APIRouter, HTTPException, Request, status +from pydantic import BaseModel, Field + +from app.utils.db_migrate import migrate_data, preview_migration +from app.utils.db_wizard import ( + build_connection_string, + get_supported_backends, + parse_connection_string, + test_connection, + validate_url_format, +) + +logger = logging.getLogger(__name__) +router = APIRouter(prefix="/database", tags=["database"]) + + +def _require_admin(request: Request) -> dict: + """Ensure the caller is an admin. Raises 403 otherwise.""" + user = request.session.get("user") + if not user or not user.get("is_admin"): + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required") + return user + + +# --------------------------------------------------------------------------- +# Request / Response models +# --------------------------------------------------------------------------- + + +class ConnectionStringRequest(BaseModel): + """Request body for building a connection string.""" + + backend: str = Field(..., description="Database backend: sqlite, postgresql, mysql") + host: str = Field("", description="Database server hostname") + port: int | None = Field(None, description="Database server port") + database: str = Field("", description="Database name") + username: str = Field("", description="Authentication username") + password: str = Field("", description="Authentication password") + ssl_mode: str = Field("", description="SSL mode (e.g. require, verify-full)") + extra_options: str = Field("", description="Additional query-string options") + sqlite_path: str = Field("", description="File path for SQLite databases") + + +class TestConnectionRequest(BaseModel): + """Request body for testing a database connection.""" + + url: str = Field(..., description="Full SQLAlchemy connection URL to test") + + +class MigrateRequest(BaseModel): + """Request body for data migration.""" + + source_url: str = Field(..., description="Source database connection URL") + target_url: str = Field(..., description="Target database connection URL") + + +# --------------------------------------------------------------------------- +# Endpoints +# --------------------------------------------------------------------------- + + +@router.get("/backends") +async def list_backends() -> list[dict]: + """List all supported database backends with metadata.""" + return get_supported_backends() + + +@router.post("/build-url") +async def build_url(body: ConnectionStringRequest, request: Request) -> dict: + """Build a SQLAlchemy connection string from individual components. + + Returns the assembled URL string. + """ + _require_admin(request) + try: + url = build_connection_string( + backend=body.backend, + host=body.host, + port=body.port, + database=body.database, + username=body.username, + password=body.password, + ssl_mode=body.ssl_mode, + extra_options=body.extra_options, + sqlite_path=body.sqlite_path, + ) + return {"url": url} + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + + +@router.post("/parse-url") +async def parse_url(body: TestConnectionRequest, request: Request) -> dict: + """Parse a connection string into its components.""" + _require_admin(request) + return parse_connection_string(body.url) + + +@router.post("/validate-url") +async def validate_url(body: TestConnectionRequest, request: Request) -> dict: + """Validate a connection string format without connecting.""" + _require_admin(request) + return validate_url_format(body.url) + + +@router.post("/test-connection") +async def test_db_connection(body: TestConnectionRequest, request: Request) -> dict: + """Test connectivity to a database and return status info. + + This creates a temporary engine, executes ``SELECT 1``, and disposes + of the engine. It does **not** modify any global application state. + """ + _require_admin(request) + return test_connection(body.url) + + +@router.post("/preview-migration") +async def preview_db_migration(body: TestConnectionRequest, request: Request) -> dict: + """Preview what a migration from the given source would include. + + Returns a table-by-table row count without actually copying data. + """ + _require_admin(request) + return preview_migration(body.url) + + +@router.post("/migrate") +async def execute_migration(body: MigrateRequest, request: Request) -> dict: + """Execute a full data migration from source to target database. + + **Warning:** This copies all data from the source database into the + target. The target schema is created from the current application + models. Existing data in the target is **not** deleted first — use + on an empty target database. + """ + _require_admin(request) + + # Validate both URLs first + src_check = validate_url_format(body.source_url) + if not src_check.get("valid"): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Invalid source URL: {src_check.get('error', 'unknown')}", + ) + tgt_check = validate_url_format(body.target_url) + if not tgt_check.get("valid"): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Invalid target URL: {tgt_check.get('error', 'unknown')}", + ) + + result = migrate_data(body.source_url, body.target_url) + if not result["success"]: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail={"message": "Migration completed with errors", **result}, + ) + return result diff --git a/app/utils/db_migrate.py b/app/utils/db_migrate.py new file mode 100644 index 00000000..815b7cef --- /dev/null +++ b/app/utils/db_migrate.py @@ -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") diff --git a/app/utils/db_wizard.py b/app/utils/db_wizard.py new file mode 100644 index 00000000..c6e83e03 --- /dev/null +++ b/app/utils/db_wizard.py @@ -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)} diff --git a/app/views/__init__.py b/app/views/__init__.py index ed55ff29..30e4d9ed 100644 --- a/app/views/__init__.py +++ b/app/views/__init__.py @@ -4,6 +4,7 @@ Aggregated view routers for the application. from fastapi import APIRouter +from app.views.db_wizard import router as db_wizard_router from app.views.dropbox import router as dropbox_router from app.views.filemanager import router as filemanager_router @@ -21,6 +22,7 @@ from app.views.wizard import router as wizard_router # Create a main router that includes all the view routers router = APIRouter() router.include_router(wizard_router) # Wizard first (for /setup) +router.include_router(db_wizard_router) # Database wizard router.include_router(general_router) router.include_router(status_router) router.include_router(onedrive_router) diff --git a/app/views/db_wizard.py b/app/views/db_wizard.py new file mode 100644 index 00000000..9c1b715a --- /dev/null +++ b/app/views/db_wizard.py @@ -0,0 +1,28 @@ +""" +Database configuration wizard view. + +Serves the guided UI for configuring a database connection string +and migrating data from one database to another. +""" + +import logging + +from fastapi import Request + +from app.config import settings +from app.views.base import APIRouter, templates + +logger = logging.getLogger(__name__) +router = APIRouter() + + +@router.get("/database-wizard") +async def database_wizard(request: Request) -> templates.TemplateResponse: + """Render the database configuration wizard page.""" + return templates.TemplateResponse( + "db_wizard.html", + { + "request": request, + "current_database_url": settings.database_url, + }, + ) diff --git a/frontend/templates/db_wizard.html b/frontend/templates/db_wizard.html new file mode 100644 index 00000000..ba13ec07 --- /dev/null +++ b/frontend/templates/db_wizard.html @@ -0,0 +1,606 @@ +{% extends "base.html" %} +{% block title %}Database Configuration Wizard - DocuElevate{% endblock %} + +{% block head_extra %} + +{% endblock %} + +{% block content %} +
+ +
+ + {# ── Header ── #} +
+

+ + Database Configuration Wizard +

+

+ Configure a new database connection or migrate your data to an external database. +

+
+ + {# ── Tab Navigation ── #} +
+ + +
+ + {# ═══════════════════════════════════════════════════════════════════ #} + {# TAB 1 — Configure #} + {# ═══════════════════════════════════════════════════════════════════ #} +
+
+ + {# Step indicator #} +
+

+ Step 1: Choose Database Type + Step 2: Connection Details + Step 3: Test & Apply +

+
+ +
+
+ +
+ + {# ── Step 1: Choose backend ── #} +
+

Select the database engine you want to use.

+
+ + + + + + + +
+ +
+ +
+
+ + {# ── Step 2: Connection details ── #} +
+ + {# SQLite path #} + + + {# Host-based databases #} + + + {# Live preview of the URL #} +
+ + +
+ +
+ + +
+
+ + {# ── Step 3: Test & Apply ── #} +
+ +
+ + +
+ + {# Test button #} +
+ +
+
+ + + +
+ +
+
+ + {# Apply as DATABASE_URL #} +
+

+ + To use this database, set the DATABASE_URL environment variable + (in your .env file or Docker Compose config) to the connection string above, + then restart DocuElevate. +

+
+ +
+ +
+ + +
+

+ Copied! +

+
+ +
+ + + Go to Settings + +
+
+ +
+
+
+ + {# ═══════════════════════════════════════════════════════════════════ #} + {# TAB 2 — Migrate #} + {# ═══════════════════════════════════════════════════════════════════ #} +
+
+ +
+

+ + Migrate Data Between Databases +

+

+ Copy all your data from one database to another (e.g. SQLite → PostgreSQL). +

+
+ +
+ + {# Source URL #} +
+ + +

+ This is your current database. Pre-filled with the running configuration. +

+ +
+ + {# Target URL #} +
+ + +

+ The new database to copy data into. Must be empty (schema will be created automatically). +

+ +
+ + {# Actions #} +
+ + + +
+ + {# Connection test results #} +
+ + + Source: + +
+
+ + + Target: + +
+ + {# Preview table #} +
+ + + + + + + + + + + + + + +
TableRows
Total
+
+ + {# Migrate button #} +
+
+

+ + Warning: This will copy all data to the target database. + The target must be empty. This operation cannot be undone. +

+
+ + + +
+ +
+
+ + {# Migration progress / result #} +
+
+
+
+

Migration in progress — please do not close this page…

+
+ +
+
+

Migration Successful

+

+ Copied rows + across tables. +

+

+ Update your DATABASE_URL environment variable to the target URL and restart DocuElevate. +

+
+
+

Migration Failed

+ +

+
+
+ +
+
+
+ + {# ── Help text ── #} +
+

+ + See the Database Configuration Guide for more details. +

+
+ +
+
+ + +{% endblock %} diff --git a/tests/test_db_migrate.py b/tests/test_db_migrate.py new file mode 100644 index 00000000..dd49988b --- /dev/null +++ b/tests/test_db_migrate.py @@ -0,0 +1,149 @@ +"""Tests for app/utils/db_migrate.py module.""" + +from unittest.mock import MagicMock, patch + +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import StaticPool + +from app.database import Base +from app.utils.db_migrate import migrate_data, preview_migration + + +@pytest.mark.unit +class TestPreviewMigration: + """Tests for preview_migration function.""" + + def test_preview_in_memory_sqlite(self): + """Test previewing an in-memory SQLite database.""" + # Create a temporary source DB with some data + src_engine = create_engine( + "sqlite:///:memory:", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + Base.metadata.create_all(bind=src_engine) + + # Insert a test row + Session = sessionmaker(bind=src_engine) + session = Session() + session.execute(text("INSERT INTO documents (filename) VALUES ('test.pdf')")) + session.commit() + session.close() + + # Preview using the engine's URL won't work for :memory:, + # but we can test the error path + result = preview_migration("sqlite:///:memory:") + # For :memory: this creates a new empty DB, so tables are empty + assert result["success"] is True + assert isinstance(result["tables"], list) + + def test_preview_invalid_url(self): + """Test preview with invalid URL returns error.""" + result = preview_migration("invalid://not-a-db") + assert result["success"] is False + assert "error" in result + + +@pytest.mark.unit +class TestMigrateData: + """Tests for migrate_data function.""" + + def test_migrate_empty_sqlite_to_sqlite(self): + """Test migrating an empty SQLite DB to another SQLite DB.""" + # Both are file-based temp databases for this test + src_url = "sqlite:///:memory:" + tgt_url = "sqlite://" # Another in-memory DB + + # Create source schema + src_engine = create_engine(src_url, connect_args={"check_same_thread": False}, poolclass=StaticPool) + Base.metadata.create_all(bind=src_engine) + src_engine.dispose() + + # Run migration from empty source + with patch("app.utils.db_migrate._make_engine") as mock_make: + # Create real engines for both + real_src = create_engine( + "sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool + ) + Base.metadata.create_all(bind=real_src) + real_tgt = create_engine( + "sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool + ) + mock_make.side_effect = [real_src, real_tgt] + + with patch("app.utils.db_migrate._stamp_alembic_head"): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + assert result["success"] is True + assert result["rows_copied"] == 0 + + def test_migrate_with_data(self): + """Test migrating a SQLite DB with actual data.""" + real_src = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + Base.metadata.create_all(bind=real_src) + + # Insert test data + Session = sessionmaker(bind=real_src) + session = Session() + session.execute(text("INSERT INTO documents (filename) VALUES ('invoice.pdf')")) + session.execute(text("INSERT INTO documents (filename) VALUES ('receipt.pdf')")) + session.commit() + session.close() + + real_tgt = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + + with patch("app.utils.db_migrate._make_engine") as mock_make: + mock_make.side_effect = [real_src, real_tgt] + with patch("app.utils.db_migrate._stamp_alembic_head"): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + assert result["success"] is True + assert result["rows_copied"] >= 2 # At least the 2 documents rows + + def test_migrate_with_progress_callback(self): + """Test that progress callback is invoked during migration.""" + real_src = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + Base.metadata.create_all(bind=real_src) + + Session = sessionmaker(bind=real_src) + session = Session() + session.execute(text("INSERT INTO documents (filename) VALUES ('test.pdf')")) + session.commit() + session.close() + + real_tgt = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + callback = MagicMock() + + with patch("app.utils.db_migrate._make_engine") as mock_make: + mock_make.side_effect = [real_src, real_tgt] + with patch("app.utils.db_migrate._stamp_alembic_head"): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:", progress_callback=callback) + + assert result["success"] is True + # Callback should have been called at least once for the non-empty table + if result["rows_copied"] > 0: + assert callback.call_count > 0 + + def test_migrate_global_exception(self): + """Test that a global exception is caught gracefully.""" + with patch("app.utils.db_migrate._make_engine", side_effect=Exception("boom")): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + assert result["success"] is False + assert len(result["errors"]) > 0 + + def test_migrate_stamp_failure_is_recorded(self): + """Test that Alembic stamp failure is recorded as an error.""" + real_src = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + Base.metadata.create_all(bind=real_src) + real_tgt = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + + with patch("app.utils.db_migrate._make_engine") as mock_make: + mock_make.side_effect = [real_src, real_tgt] + with patch("app.utils.db_migrate._stamp_alembic_head", side_effect=Exception("stamp failed")): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + # Data copy succeeds but stamp fails — errors list non-empty + assert len(result["errors"]) > 0 + assert any("stamp" in e.lower() for e in result["errors"]) diff --git a/tests/test_db_wizard.py b/tests/test_db_wizard.py new file mode 100644 index 00000000..ba7b6dca --- /dev/null +++ b/tests/test_db_wizard.py @@ -0,0 +1,233 @@ +"""Tests for app/utils/db_wizard.py module.""" + +import pytest + +from app.utils.db_wizard import ( + build_connection_string, + get_supported_backends, + parse_connection_string, + test_connection as db_test_connection, + validate_url_format, +) + + +@pytest.mark.unit +class TestGetSupportedBackends: + """Tests for get_supported_backends function.""" + + def test_returns_list(self): + """Test that it returns a non-empty list.""" + result = get_supported_backends() + assert isinstance(result, list) + assert len(result) >= 3 + + def test_each_backend_has_required_keys(self): + """Test that each backend has expected keys.""" + required_keys = {"id", "label", "description", "requires_host"} + for backend in get_supported_backends(): + assert required_keys.issubset(set(backend.keys())), f"Missing keys in {backend.get('id')}" + + def test_includes_sqlite(self): + """Test that SQLite is included.""" + ids = [b["id"] for b in get_supported_backends()] + assert "sqlite" in ids + + def test_includes_postgresql(self): + """Test that PostgreSQL is included.""" + ids = [b["id"] for b in get_supported_backends()] + assert "postgresql" in ids + + def test_includes_mysql(self): + """Test that MySQL is included.""" + ids = [b["id"] for b in get_supported_backends()] + assert "mysql" in ids + + +@pytest.mark.unit +class TestBuildConnectionString: + """Tests for build_connection_string function.""" + + def test_sqlite_default_path(self): + """Test building a SQLite URL with default path.""" + url = build_connection_string(backend="sqlite") + assert url == "sqlite:///./app/database.db" + + def test_sqlite_custom_path(self): + """Test building a SQLite URL with custom path.""" + url = build_connection_string(backend="sqlite", sqlite_path="/data/mydb.db") + assert url == "sqlite:////data/mydb.db" + + def test_postgresql_basic(self): + """Test building a basic PostgreSQL URL.""" + url = build_connection_string( + backend="postgresql", + host="localhost", + database="docuelevate", + username="user", + password="pass", + ) + assert url == "postgresql://user:pass@localhost:5432/docuelevate" + + def test_postgresql_with_ssl(self): + """Test building a PostgreSQL URL with SSL.""" + url = build_connection_string( + backend="postgresql", + host="rds.amazonaws.com", + database="docuelevate", + username="admin", + password="secret", + ssl_mode="require", + ) + assert "sslmode=require" in url + assert "postgresql://admin:secret@rds.amazonaws.com:5432/docuelevate" in url + + def test_postgresql_custom_port(self): + """Test building a PostgreSQL URL with custom port.""" + url = build_connection_string( + backend="postgresql", + host="localhost", + port=5433, + database="testdb", + username="user", + password="pass", + ) + assert ":5433/" in url + + def test_mysql_basic(self): + """Test building a MySQL URL.""" + url = build_connection_string( + backend="mysql", + host="localhost", + database="docuelevate", + username="root", + password="password", + ) + assert url.startswith("mysql+pymysql://") + assert "charset=utf8mb4" in url + + def test_mysql_no_duplicate_charset(self): + """Test that charset is not duplicated when passed in extra_options.""" + url = build_connection_string( + backend="mysql", + host="localhost", + database="docuelevate", + username="root", + password="pass", + extra_options="charset=utf8mb4", + ) + assert url.count("charset=utf8mb4") == 1 + + def test_unsupported_backend_raises(self): + """Test that unsupported backend raises ValueError.""" + with pytest.raises(ValueError, match="Unsupported backend"): + build_connection_string(backend="oracle") + + def test_missing_host_raises(self): + """Test that missing host for non-SQLite raises ValueError.""" + with pytest.raises(ValueError, match="Host is required"): + build_connection_string(backend="postgresql", database="db", username="u") + + def test_missing_database_raises(self): + """Test that missing database name raises ValueError.""" + with pytest.raises(ValueError, match="Database name is required"): + build_connection_string(backend="postgresql", host="localhost", username="u") + + def test_missing_username_raises(self): + """Test that missing username raises ValueError.""" + with pytest.raises(ValueError, match="Username is required"): + build_connection_string(backend="postgresql", host="localhost", database="db") + + def test_no_password(self): + """Test building URL without password.""" + url = build_connection_string( + backend="postgresql", + host="localhost", + database="db", + username="user", + ) + assert "user@localhost" in url + assert ":@" not in url + + +@pytest.mark.unit +class TestParseConnectionString: + """Tests for parse_connection_string function.""" + + def test_parse_sqlite(self): + """Test parsing a SQLite URL.""" + result = parse_connection_string("sqlite:///./app/database.db") + assert result["valid"] is True + assert result["backend"] == "sqlite" + assert result["is_sqlite"] is True + + def test_parse_postgresql(self): + """Test parsing a PostgreSQL URL.""" + result = parse_connection_string("postgresql://user:pass@host:5432/mydb") + assert result["valid"] is True + assert result["backend"] == "postgresql" + assert result["host"] == "host" + assert result["port"] == 5432 + assert result["database"] == "mydb" + assert result["username"] == "user" + assert result["is_sqlite"] is False + + def test_parse_mysql(self): + """Test parsing a MySQL URL.""" + result = parse_connection_string("mysql+pymysql://root:pass@localhost:3306/db") + assert result["valid"] is True + assert result["backend"] == "mysql" + + def test_parse_invalid_url(self): + """Test parsing an invalid URL returns error.""" + result = parse_connection_string("not-a-valid-url://") + # Should still return a dict (make_url may or may not raise) + assert isinstance(result, dict) + + +@pytest.mark.unit +class TestValidateUrlFormat: + """Tests for validate_url_format function.""" + + def test_valid_sqlite(self): + """Test valid SQLite URL.""" + result = validate_url_format("sqlite:///./db.sqlite") + assert result["valid"] is True + assert result["backend"] == "sqlite" + + def test_valid_postgresql(self): + """Test valid PostgreSQL URL.""" + result = validate_url_format("postgresql://u:p@host/db") + assert result["valid"] is True + + def test_valid_mysql(self): + """Test valid MySQL URL.""" + result = validate_url_format("mysql+pymysql://u:p@host/db") + assert result["valid"] is True + + def test_unsupported_backend(self): + """Test that unsupported backends are flagged.""" + result = validate_url_format("mssql://u:p@host/db") + assert result["valid"] is False + assert "Unsupported" in result.get("error", "") + + def test_invalid_format(self): + """Test that garbage input is invalid.""" + result = validate_url_format("") + assert result["valid"] is False + + +@pytest.mark.unit +class TestTestConnection: + """Tests for test_connection function.""" + + def test_sqlite_memory_succeeds(self): + """Test connecting to an in-memory SQLite database.""" + result = db_test_connection("sqlite:///:memory:") + assert result["success"] is True + assert "SQLite" in result.get("server_version", "") + + def test_unreachable_host_fails(self): + """Test that an unreachable host returns failure.""" + result = db_test_connection("postgresql://u:p@192.0.2.1:5432/db", timeout=2) + assert result["success"] is False + assert result["message"] # Should contain an error message diff --git a/tests/test_db_wizard_api.py b/tests/test_db_wizard_api.py new file mode 100644 index 00000000..37437a11 --- /dev/null +++ b/tests/test_db_wizard_api.py @@ -0,0 +1,156 @@ +"""Tests for app/api/database.py and app/views/db_wizard.py modules.""" + +from unittest.mock import patch + +import pytest + + +@pytest.mark.integration +class TestDatabaseApiEndpoints: + """Tests for the database API endpoints.""" + + def test_list_backends(self, client): + """Test GET /api/database/backends returns supported backends.""" + response = client.get("/api/database/backends") + assert response.status_code == 200 + data = response.json() + assert isinstance(data, list) + assert len(data) >= 3 + ids = [b["id"] for b in data] + assert "sqlite" in ids + assert "postgresql" in ids + + def test_build_url_requires_admin(self, client): + """Test POST /api/database/build-url requires admin.""" + response = client.post( + "/api/database/build-url", + json={"backend": "sqlite"}, + ) + assert response.status_code == 403 + + def test_build_url_sqlite(self, client): + """Test building a SQLite URL as admin.""" + # Simulate admin session + with client.session_transaction() if hasattr(client, "session_transaction") else _noop(): + pass + # Use the session cookie approach + client.cookies.set("session", "test") + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/build-url", + json={"backend": "sqlite", "sqlite_path": "/data/test.db"}, + ) + assert response.status_code == 200 + assert "sqlite:////data/test.db" in response.json().get("url", "") + + def test_build_url_missing_host(self, client): + """Test building URL with missing host returns 400.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/build-url", + json={"backend": "postgresql", "database": "db", "username": "u"}, + ) + assert response.status_code == 400 + + def test_test_connection_sqlite(self, client): + """Test connection to in-memory SQLite.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/test-connection", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["success"] is True + + def test_validate_url_valid(self, client): + """Test validate-url with valid SQLite URL.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/validate-url", + json={"url": "sqlite:///test.db"}, + ) + assert response.status_code == 200 + assert response.json()["valid"] is True + + def test_validate_url_invalid(self, client): + """Test validate-url with unsupported backend.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/validate-url", + json={"url": "mssql://u:p@h/d"}, + ) + assert response.status_code == 200 + assert response.json()["valid"] is False + + def test_parse_url(self, client): + """Test parse-url endpoint.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/parse-url", + json={"url": "postgresql://user:pass@host:5432/db"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["backend"] == "postgresql" + assert data["host"] == "host" + + def test_preview_migration(self, client): + """Test preview-migration endpoint.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/preview-migration", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 200 + data = response.json() + assert "tables" in data + + def test_migrate_invalid_source(self, client): + """Test migrate endpoint with invalid source URL.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/migrate", + json={"source_url": "mssql://bad", "target_url": "sqlite:///:memory:"}, + ) + assert response.status_code == 400 + + def test_migrate_invalid_target(self, client): + """Test migrate endpoint with invalid target URL.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/migrate", + json={"source_url": "sqlite:///:memory:", "target_url": "mssql://bad"}, + ) + assert response.status_code == 400 + + +@pytest.mark.integration +class TestDatabaseWizardView: + """Tests for the database wizard view.""" + + def test_database_wizard_page_loads(self, client): + """Test GET /database-wizard returns 200.""" + response = client.get("/database-wizard") + assert response.status_code == 200 + + def test_database_wizard_contains_title(self, client): + """Test that the wizard page contains expected content.""" + response = client.get("/database-wizard") + assert response.status_code == 200 + assert "Database Configuration Wizard" in response.text + + def test_database_wizard_contains_tabs(self, client): + """Test that the wizard page contains configure and migrate tabs.""" + response = client.get("/database-wizard") + assert "Configure Database" in response.text + assert "Migrate Data" in response.text + + +# Context manager helper for tests that don't need session_transaction +class _noop: + def __enter__(self): + return None + + def __exit__(self, *args): + pass