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..b543aa04 --- /dev/null +++ b/app/api/database.py @@ -0,0 +1,170 @@ +""" +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"]: + error_summary = "; ".join(result.get("errors", ["Unknown error"])) + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"Migration completed with errors: {error_summary}", + ) + return result diff --git a/app/utils/db_migrate.py b/app/utils/db_migrate.py new file mode 100644 index 00000000..f95d97f0 --- /dev/null +++ b/app/utils/db_migrate.py @@ -0,0 +1,251 @@ +""" +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: + # table_name is safe — sourced from inspect().get_table_names(), not user input + 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] + # strict=False: column count should always match, but tolerate + # minor schema drift (e.g. extra columns) to avoid crashing mid-migration. + 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..e8d53a89 --- /dev/null +++ b/app/utils/db_wizard.py @@ -0,0 +1,257 @@ +""" +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 "" + if not path: + path = "./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/utils/settings_service.py b/app/utils/settings_service.py index c324c84d..9c79e2d8 100644 --- a/app/utils/settings_service.py +++ b/app/utils/settings_service.py @@ -23,11 +23,13 @@ SETTING_METADATA = { # Core Settings "database_url": { "category": "Core", - "description": "Database connection URL (e.g., sqlite:///path/to/db.sqlite)", + "description": "Database connection URL (e.g., sqlite:///path/to/db.sqlite). Use the Database Wizard for guided setup.", "type": "string", "sensitive": False, "required": True, "restart_required": True, + "help_link": "/database-wizard", + "help_link_label": "Open Database Wizard", }, "redis_url": { "category": "Core", 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/docs/API.md b/docs/API.md index 2afd5ac2..1372dbdb 100644 --- a/docs/API.md +++ b/docs/API.md @@ -1238,6 +1238,125 @@ Send a test notification through all configured notification channels. The API implements rate limiting to ensure system stability. If you exceed the limits, you'll receive a `429 Too Many Requests` response. +## Database Configuration Wizard + +Endpoints for building and testing database connection strings and migrating data between databases. All write endpoints require admin authentication. + +### GET /api/database/backends + +List supported database backends with metadata. + +**Response (200):** +```json +[ + { + "id": "sqlite", + "label": "SQLite (Development)", + "default_port": null, + "description": "File-based database. Best for development and single-user setups.", + "requires_host": false + }, + { + "id": "postgresql", + "label": "PostgreSQL (Recommended for Production)", + "default_port": 5432, + "description": "Robust, full-featured database. Recommended for production.", + "requires_host": true + } +] +``` + +### POST /api/database/build-url + +Build a SQLAlchemy connection string from individual components. + +**Request:** +```json +{ + "backend": "postgresql", + "host": "my-db.rds.amazonaws.com", + "port": 5432, + "database": "docuelevate", + "username": "admin", + "password": "secret", + "ssl_mode": "require" +} +``` + +**Response (200):** +```json +{ + "url": "postgresql://admin:secret@my-db.rds.amazonaws.com:5432/docuelevate?sslmode=require" +} +``` + +### POST /api/database/test-connection + +Test connectivity to a database. + +**Request:** +```json +{ + "url": "postgresql://admin:secret@my-db.rds.amazonaws.com:5432/docuelevate?sslmode=require" +} +``` + +**Response (200):** +```json +{ + "success": true, + "message": "Connection successful", + "backend": "postgresql", + "server_version": "PostgreSQL 16.2 on x86_64-pc-linux-gnu" +} +``` + +### POST /api/database/preview-migration + +Preview a data migration (table-by-table row counts) without copying data. + +**Request:** +```json +{ + "url": "sqlite:///./app/database.db" +} +``` + +**Response (200):** +```json +{ + "success": true, + "tables": [ + {"name": "documents", "row_count": 42}, + {"name": "files", "row_count": 150} + ], + "total_rows": 192 +} +``` + +### POST /api/database/migrate + +Execute a full data migration from source to target database. + +**Request:** +```json +{ + "source_url": "sqlite:///./app/database.db", + "target_url": "postgresql://admin:secret@host:5432/docuelevate" +} +``` + +**Response (200):** +```json +{ + "success": true, + "tables_copied": 8, + "rows_copied": 192, + "errors": [] +} +``` + + ## Further Assistance For additional help with the API, please contact our support team or refer to the [Development Guide](../CONTRIBUTING.md). diff --git a/docs/ConfigurationGuide.md b/docs/ConfigurationGuide.md index 9650c7d3..5c4dc3a0 100644 --- a/docs/ConfigurationGuide.md +++ b/docs/ConfigurationGuide.md @@ -10,7 +10,7 @@ Configuration is primarily done through environment variables specified in a `.e | **Variable** | **Description** | **Example** | |------------------------|----------------------------------------------------------|--------------------------------| -| `DATABASE_URL` | Path/URL to the SQLite database (or other SQL backend). | `sqlite:///./app/database.db` | +| `DATABASE_URL` | Path/URL to the SQLite database (or other SQL backend). Use the [Database Wizard](/database-wizard) for guided setup. See [Database Configuration](DatabaseConfiguration.md). | `sqlite:///./app/database.db` | | `REDIS_URL` | URL for Redis, used by Celery for broker & result store. | `redis://redis:6379/0` | | `WORKDIR` | Working directory for the application. | `/workdir` | | `GOTENBERG_URL` | Gotenberg PDF processing URL. | `http://gotenberg:3000` | diff --git a/docs/DatabaseConfiguration.md b/docs/DatabaseConfiguration.md index 5d08e64a..0e5811aa 100644 --- a/docs/DatabaseConfiguration.md +++ b/docs/DatabaseConfiguration.md @@ -4,6 +4,8 @@ DocuElevate uses [SQLAlchemy](https://www.sqlalchemy.org/) as its ORM and [Alemb ## Table of Contents +- [Database Configuration Wizard](#database-configuration-wizard) +- [Database Migration Tool](#database-migration-tool) - [Supported Databases](#supported-databases) - [Configuration](#configuration) - [SQLite (Development)](#sqlite-development) @@ -17,6 +19,74 @@ DocuElevate uses [SQLAlchemy](https://www.sqlalchemy.org/) as its ORM and [Alemb --- +## Database Configuration Wizard + +DocuElevate includes a guided **Database Configuration Wizard** accessible at `/database-wizard`. The wizard walks you through building a connection string step by step — no need to remember the exact URL format. + +### How to Access + +Navigate to **`/database-wizard`** in your browser, or find the link under **Admin → Settings**. On the Settings page, click the **DB Wizard** button in the toolbar, or look for the **Open Database Wizard** link next to the `database_url` setting. + +### Wizard Steps + +1. **Choose Database Type** — select SQLite, PostgreSQL, or MySQL/MariaDB. +2. **Connection Details** — enter host, port, database name, credentials, and SSL mode (auto-populated with sensible defaults). +3. **Test & Apply** — test the connection before committing, then copy the generated `DATABASE_URL` into your `.env` file. + +The wizard generates the full SQLAlchemy connection string and lets you test connectivity directly from the UI. After testing, copy the `DATABASE_URL=…` line into your `.env` file (or Docker Compose environment) and restart DocuElevate. + +### REST API + +The wizard is backed by a REST API under `/api/database/`: + +| Endpoint | Method | Description | +|----------|--------|-------------| +| `/api/database/backends` | GET | List supported database backends | +| `/api/database/build-url` | POST | Build a connection string from components | +| `/api/database/parse-url` | POST | Parse a connection string into components | +| `/api/database/validate-url` | POST | Validate URL format without connecting | +| `/api/database/test-connection` | POST | Test connectivity to a database | + +> **Note:** All write endpoints require admin authentication. + +--- + +## Database Migration Tool + +The **Migrate Data** tab (on the same `/database-wizard` page) lets you copy all your data from one database to another — for example, migrating from the built-in SQLite database to an external PostgreSQL or MySQL instance. + +### When to Use + +- Moving from a development SQLite database to a production PostgreSQL instance. +- Migrating to a managed cloud database (AWS RDS, Google Cloud SQL, Azure Database for PostgreSQL, Supabase, etc.). +- Consolidating data from one database engine to another. + +### How It Works + +1. Enter (or auto-fill) the **Source Database URL** — this is your current database. +2. Enter the **Target Database URL** — the new, empty database to copy data into. +3. Click **Test Source** and **Test Target** to verify both connections. +4. Click **Preview Migration** to see a table-by-table row count. +5. Confirm and click **Start Migration** to copy all data. + +The migration tool: +- Creates the full schema in the target database from the application models. +- Copies all rows table by table in dependency order (parent tables first). +- Stamps the Alembic migration version to `head` in the target. + +After migration, update your `DATABASE_URL` environment variable to point at the new database and restart DocuElevate. + +### REST API + +| Endpoint | Method | Description | +|----------|--------|-------------| +| `/api/database/preview-migration` | POST | Preview tables and row counts | +| `/api/database/migrate` | POST | Execute the full data migration | + +> **Warning:** Always run the migration against an **empty** target database. The tool does not delete existing data in the target before copying. + +--- + ## Supported Databases | Database | Recommended Use | Notes | diff --git a/frontend/templates/db_wizard.html b/frontend/templates/db_wizard.html new file mode 100644 index 00000000..447de046 --- /dev/null +++ b/frontend/templates/db_wizard.html @@ -0,0 +1,618 @@ +{% 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/frontend/templates/settings.html b/frontend/templates/settings.html index 523e4983..40311ea1 100644 --- a/frontend/templates/settings.html +++ b/frontend/templates/settings.html @@ -60,6 +60,10 @@ class="inline-flex items-center px-3 py-2 border border-gray-300 text-sm font-medium rounded-md text-gray-700 bg-white hover:bg-gray-50"> Wizard + + DB Wizard +

{{ setting.metadata.description }}

+ {% if setting.metadata.get('help_link') %} +

+ + {{ setting.metadata.get('help_link_label', 'More info') }} + +

+ {% endif %} {% if setting.metadata.type == 'boolean' %} diff --git a/tests/test_db_migrate.py b/tests/test_db_migrate.py new file mode 100644 index 00000000..a368a538 --- /dev/null +++ b/tests/test_db_migrate.py @@ -0,0 +1,359 @@ +"""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 ( + _make_engine, + _ordered_tables, + _stamp_alembic_head, + migrate_data, + preview_migration, +) + + +@pytest.mark.unit +class TestMakeEngine: + """Tests for _make_engine helper function.""" + + def test_sqlite_engine_has_check_same_thread(self): + """Test that SQLite engine has check_same_thread set.""" + engine = _make_engine("sqlite:///:memory:") + assert engine is not None + engine.dispose() + + def test_non_sqlite_engine_created(self): + """Test that non-SQLite engine can be created (even if driver is missing).""" + # _make_engine only creates the engine object; it doesn't connect. + # If the driver isn't installed, create_engine raises at creation time. + try: + engine = _make_engine("postgresql://u:p@localhost:5432/test") + assert engine is not None + engine.dispose() + except Exception: + # Driver not installed in test environment — acceptable + pass + + +@pytest.mark.unit +class TestOrderedTables: + """Tests for _ordered_tables helper function.""" + + def test_known_tables_come_first(self): + """Test that known tables from _TABLE_ORDER come first.""" + mock_inspector = MagicMock() + mock_inspector.get_table_names.return_value = [ + "webhook_configs", + "documents", + "files", + "custom_table", + "alembic_version", + ] + result = _ordered_tables(mock_inspector) + # alembic_version should be skipped + assert "alembic_version" not in result + # Known tables should come first in their predefined order + assert result.index("documents") < result.index("files") + assert result.index("files") < result.index("webhook_configs") + # custom_table is not in _TABLE_ORDER so comes after known tables + assert "custom_table" in result + + def test_skips_alembic_version(self): + """Test that alembic_version table is always skipped.""" + mock_inspector = MagicMock() + mock_inspector.get_table_names.return_value = ["alembic_version", "documents"] + result = _ordered_tables(mock_inspector) + assert "alembic_version" not in result + assert "documents" in result + + def test_unknown_tables_appended_alphabetically(self): + """Test that tables not in _TABLE_ORDER are appended alphabetically.""" + mock_inspector = MagicMock() + mock_inspector.get_table_names.return_value = ["zebra", "apple", "documents"] + result = _ordered_tables(mock_inspector) + assert result[0] == "documents" + # apple and zebra should be after documents, in alpha order + remaining = result[1:] + assert remaining == sorted(remaining) + + def test_empty_database(self): + """Test with an empty database returns empty list.""" + mock_inspector = MagicMock() + mock_inspector.get_table_names.return_value = [] + result = _ordered_tables(mock_inspector) + assert result == [] + + +@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 + + def test_preview_with_patched_source_shows_tables(self): + """Test preview with source that has tables and 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 ('test.pdf')")) + session.commit() + session.close() + + with patch("app.utils.db_migrate._make_engine", return_value=real_src): + result = preview_migration("sqlite:///:memory:") + + assert result["success"] is True + assert result["total_rows"] >= 1 + # At least the documents table should be in results + table_names = [t["name"] for t in result["tables"]] + assert "documents" in table_names + doc_table = next(t for t in result["tables"] if t["name"] == "documents") + assert doc_table["row_count"] >= 1 + + +@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"]) + + def test_migrate_table_copy_exception(self): + """Test that per-table copy exception is recorded but migration continues.""" + real_src = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool) + Base.metadata.create_all(bind=real_src) + + # Insert data so the table isn't empty + 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) + Base.metadata.create_all(bind=real_tgt) + + # Make the target reflect fail for one table to trigger the error path + original_reflect = MagicMock(side_effect=Exception("reflect error")) + + 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"): + # Patch MetaData so that reflecting target raises for the first table + with patch("app.utils.db_migrate.MetaData") as mock_meta_cls: + # First MetaData() is for source reflect (should work) + src_meta = MagicMock() + src_table = MagicMock() + src_table.columns = [] + src_table.select.return_value = text("SELECT 1") + src_meta.tables = {"documents": src_table} + src_meta.reflect = MagicMock() + + # Second MetaData() is for target reflect (should fail) + tgt_meta = MagicMock() + tgt_meta.reflect.side_effect = Exception("target reflect error") + + mock_meta_cls.side_effect = [src_meta, tgt_meta] + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + assert any("Error copying table" in e for e in result["errors"]) + + def test_migrate_target_table_not_found(self): + """Test that missing target table after reflect is recorded.""" + 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) + # Create schema in target so reflect works but returns empty + Base.metadata.create_all(bind=real_tgt) + + 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:") + + # With schema in target, migration should succeed normally + assert result["success"] is True + + def test_migrate_returns_tables_and_rows_counts(self): + """Test that successful migration returns expected count fields.""" + 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"): + result = migrate_data("sqlite:///:memory:", "sqlite:///:memory:") + + assert "tables_copied" in result + assert "rows_copied" in result + assert "errors" in result + assert isinstance(result["errors"], list) + + +@pytest.mark.unit +class TestStampAlembicHead: + """Tests for _stamp_alembic_head helper function.""" + + def test_stamp_calls_alembic_command(self): + """Test that stamping calls alembic command.stamp with 'head'.""" + mock_engine = MagicMock() + mock_connection = MagicMock() + mock_engine.begin.return_value.__enter__ = MagicMock(return_value=mock_connection) + mock_engine.begin.return_value.__exit__ = MagicMock(return_value=False) + + with patch("alembic.command.stamp") as mock_stamp: + _stamp_alembic_head(mock_engine) + mock_stamp.assert_called_once() + # Verify it stamps to "head" + args = mock_stamp.call_args + assert args[0][1] == "head" + + def test_stamp_raises_on_error(self): + """Test that stamp propagates exceptions.""" + mock_engine = MagicMock() + mock_connection = MagicMock() + mock_engine.begin.return_value.__enter__ = MagicMock(return_value=mock_connection) + mock_engine.begin.return_value.__exit__ = MagicMock(return_value=False) + + with patch("alembic.command.stamp", side_effect=Exception("stamp error")): + with pytest.raises(Exception, match="stamp error"): + _stamp_alembic_head(mock_engine) diff --git a/tests/test_db_wizard.py b/tests/test_db_wizard.py new file mode 100644 index 00000000..d94251b9 --- /dev/null +++ b/tests/test_db_wizard.py @@ -0,0 +1,378 @@ +"""Tests for app/utils/db_wizard.py module.""" + +from unittest.mock import MagicMock + +import pytest + +from app.utils.db_wizard import ( + _get_server_version, + build_connection_string, + get_supported_backends, + parse_connection_string, + validate_url_format, +) +from app.utils.db_wizard import ( + test_connection as db_test_connection, +) + + +@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 + + def test_sqlite_does_not_require_host(self): + """Test that SQLite backend does not require host.""" + sqlite = next(b for b in get_supported_backends() if b["id"] == "sqlite") + assert sqlite["requires_host"] is False + + def test_postgresql_requires_host(self): + """Test that PostgreSQL backend requires host.""" + pg = next(b for b in get_supported_backends() if b["id"] == "postgresql") + assert pg["requires_host"] is True + + def test_mysql_default_port(self): + """Test that MySQL has default port 3306.""" + mysql = next(b for b in get_supported_backends() if b["id"] == "mysql") + assert mysql["default_port"] == 3306 + + +@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_sqlite_whitespace_path(self): + """Test building a SQLite URL with whitespace-only path uses default.""" + url = build_connection_string(backend="sqlite", sqlite_path=" ") + assert url == "sqlite:///./app/database.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_mysql_extra_options(self): + """Test MySQL URL with extra options appended.""" + url = build_connection_string( + backend="mysql", + host="localhost", + database="docuelevate", + username="root", + password="pass", + extra_options="connect_timeout=10", + ) + assert "connect_timeout=10" in url + assert "charset=utf8mb4" in url + + 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 + + def test_postgresql_with_extra_options(self): + """Test PostgreSQL URL with extra query options.""" + url = build_connection_string( + backend="postgresql", + host="localhost", + database="db", + username="user", + extra_options="application_name=docuelevate", + ) + assert "application_name=docuelevate" in url + + def test_postgresql_ssl_and_extra_options(self): + """Test PostgreSQL URL with both SSL and extra options combined.""" + url = build_connection_string( + backend="postgresql", + host="localhost", + database="db", + username="user", + ssl_mode="require", + extra_options="application_name=docuelevate", + ) + assert "sslmode=require" in url + assert "application_name=docuelevate" 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) + + def test_parse_postgresql_no_password(self): + """Test parsing a PostgreSQL URL without password.""" + result = parse_connection_string("postgresql://user@host:5432/mydb") + assert result["valid"] is True + assert result["password"] == "" + + def test_parse_sqlite_memory(self): + """Test parsing a SQLite in-memory URL.""" + result = parse_connection_string("sqlite:///:memory:") + assert result["valid"] is True + assert result["is_sqlite"] is True + + +@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 + + def test_returns_backend_field(self): + """Test that the backend field is populated on success.""" + result = db_test_connection("sqlite:///:memory:") + assert result["backend"] == "sqlite" + + def test_failure_returns_empty_backend(self): + """Test that failure returns empty backend.""" + result = db_test_connection("postgresql://u:p@192.0.2.1:5432/db", timeout=1) + assert result["backend"] == "" + assert result["server_version"] == "" + + +@pytest.mark.unit +class TestGetServerVersion: + """Tests for _get_server_version internal function.""" + + def test_postgresql_version(self): + """Test PostgreSQL version retrieval.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = ("PostgreSQL 16.2 on x86_64",) + result = _get_server_version(mock_conn, "postgresql") + assert result == "PostgreSQL 16.2 on x86_64" + + def test_mysql_version(self): + """Test MySQL version retrieval.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = ("8.0.36",) + result = _get_server_version(mock_conn, "mysql") + assert result == "8.0.36" + + def test_sqlite_version(self): + """Test SQLite version retrieval.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = ("3.45.1",) + result = _get_server_version(mock_conn, "sqlite") + assert result == "SQLite 3.45.1" + + def test_postgresql_empty_row(self): + """Test PostgreSQL version with empty row returns empty string.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = None + result = _get_server_version(mock_conn, "postgresql") + assert result == "" + + def test_mysql_empty_row(self): + """Test MySQL version with empty row returns empty string.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = None + result = _get_server_version(mock_conn, "mysql") + assert result == "" + + def test_sqlite_empty_row(self): + """Test SQLite version with empty row returns empty string.""" + mock_conn = MagicMock() + mock_conn.execute.return_value.fetchone.return_value = None + result = _get_server_version(mock_conn, "sqlite") + assert result == "" + + def test_unknown_backend_returns_empty(self): + """Test that an unknown backend returns empty string.""" + mock_conn = MagicMock() + result = _get_server_version(mock_conn, "oracle") + assert result == "" + + def test_exception_returns_empty(self): + """Test that an exception returns empty string.""" + mock_conn = MagicMock() + mock_conn.execute.side_effect = Exception("Connection lost") + result = _get_server_version(mock_conn, "postgresql") + assert result == "" diff --git a/tests/test_db_wizard_api.py b/tests/test_db_wizard_api.py new file mode 100644 index 00000000..2bcf3247 --- /dev/null +++ b/tests/test_db_wizard_api.py @@ -0,0 +1,377 @@ +"""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 _NoOpContextManager(): + 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 + + def test_migrate_success(self, client): + """Test migrate endpoint with successful migration.""" + mock_result = {"success": True, "tables_copied": 5, "rows_copied": 100, "errors": []} + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + with patch("app.api.database.migrate_data", return_value=mock_result): + response = client.post( + "/api/database/migrate", + json={ + "source_url": "sqlite:///:memory:", + "target_url": "sqlite:///:memory:", + }, + ) + assert response.status_code == 200 + data = response.json() + assert data["success"] is True + assert data["rows_copied"] == 100 + + def test_migrate_failure_returns_500(self, client): + """Test migrate endpoint returns 500 on migration failure.""" + mock_result = { + "success": False, + "tables_copied": 2, + "rows_copied": 50, + "errors": ["Table X failed", "Stamp failed"], + } + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + with patch("app.api.database.migrate_data", return_value=mock_result): + response = client.post( + "/api/database/migrate", + json={ + "source_url": "sqlite:///:memory:", + "target_url": "sqlite:///:memory:", + }, + ) + assert response.status_code == 500 + assert "Table X failed" in response.json()["detail"] + + def test_test_connection_requires_admin(self, client): + """Test POST /api/database/test-connection requires admin.""" + response = client.post( + "/api/database/test-connection", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 403 + + def test_parse_url_requires_admin(self, client): + """Test POST /api/database/parse-url requires admin.""" + response = client.post( + "/api/database/parse-url", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 403 + + def test_validate_url_requires_admin(self, client): + """Test POST /api/database/validate-url requires admin.""" + response = client.post( + "/api/database/validate-url", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 403 + + def test_preview_migration_requires_admin(self, client): + """Test POST /api/database/preview-migration requires admin.""" + response = client.post( + "/api/database/preview-migration", + json={"url": "sqlite:///:memory:"}, + ) + assert response.status_code == 403 + + def test_migrate_requires_admin(self, client): + """Test POST /api/database/migrate requires admin.""" + response = client.post( + "/api/database/migrate", + json={"source_url": "sqlite:///:memory:", "target_url": "sqlite:///:memory:"}, + ) + assert response.status_code == 403 + + def test_build_url_postgresql(self, client): + """Test building a PostgreSQL URL.""" + with patch("app.api.database._require_admin", return_value={"is_admin": True}): + response = client.post( + "/api/database/build-url", + json={ + "backend": "postgresql", + "host": "localhost", + "port": 5432, + "database": "mydb", + "username": "admin", + "password": "secret", + }, + ) + assert response.status_code == 200 + url = response.json()["url"] + assert "postgresql://admin:secret@localhost:5432/mydb" in url + + +@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 + + def test_database_wizard_has_skip_link(self, client): + """Test that the wizard page includes a skip-to-content link.""" + response = client.get("/database-wizard") + assert "Skip to main content" in response.text + + def test_database_wizard_has_main_landmark(self, client): + """Test that the wizard page has a main landmark.""" + response = client.get("/database-wizard") + assert 'id="main-content"' in response.text + + def test_database_wizard_has_tablist_role(self, client): + """Test that the tab navigation has proper ARIA tablist role.""" + response = client.get("/database-wizard") + assert 'role="tablist"' in response.text + assert 'role="tab"' in response.text + assert 'role="tabpanel"' in response.text + + def test_database_wizard_has_aria_labels_on_backend_buttons(self, client): + """Test that backend selection buttons have aria-label attributes.""" + response = client.get("/database-wizard") + assert 'aria-label="Select SQLite"' in response.text + assert 'aria-label="Select PostgreSQL"' in response.text + assert 'aria-label="Select MySQL / MariaDB"' in response.text + + def test_database_wizard_has_form_labels(self, client): + """Test that form inputs have associated labels.""" + response = client.get("/database-wizard") + assert 'for="sqlite_path"' in response.text + assert 'for="db_host"' in response.text + assert 'for="db_port"' in response.text + assert 'for="db_name"' in response.text + assert 'for="db_user"' in response.text + assert 'for="db_pass"' in response.text + assert 'for="ssl_mode"' in response.text + + def test_database_wizard_has_aria_describedby(self, client): + """Test that inputs have aria-describedby pointing to help text.""" + response = client.get("/database-wizard") + assert 'aria-describedby="sqlite_path_help"' in response.text + assert 'id="sqlite_path_help"' in response.text + assert 'aria-describedby="ssl_mode_help"' in response.text + assert 'id="ssl_mode_help"' in response.text + assert 'aria-describedby="mig_source_help"' in response.text + assert 'aria-describedby="mig_target_help"' in response.text + + def test_database_wizard_has_status_roles(self, client): + """Test that dynamic feedback areas have role=status or role=alert.""" + response = client.get("/database-wizard") + assert 'role="status"' in response.text + assert 'role="alert"' in response.text + + def test_database_wizard_has_aria_live(self, client): + """Test that dynamic areas have aria-live for screen reader announcements.""" + response = client.get("/database-wizard") + assert 'aria-live="polite"' in response.text + + def test_database_wizard_has_progressbar(self, client): + """Test that the migration progress indicator has role=progressbar.""" + response = client.get("/database-wizard") + assert 'role="progressbar"' in response.text + + def test_database_wizard_has_focus_ring_styles(self, client): + """Test that interactive elements have focus ring styling.""" + response = client.get("/database-wizard") + assert "focus:ring-2" in response.text + assert "focus:outline-none" in response.text + + def test_database_wizard_has_table_scope_headers(self, client): + """Test that migration preview table has proper scope attributes.""" + response = client.get("/database-wizard") + assert 'scope="col"' in response.text + + def test_database_wizard_has_dark_mode_styles(self, client): + """Test that the wizard includes dark mode CSS overrides.""" + response = client.get("/database-wizard") + assert "html.dark" in response.text + + def test_database_wizard_copy_button_has_aria_label(self, client): + """Test that the copy-to-clipboard button has an aria-label.""" + response = client.get("/database-wizard") + assert 'aria-label="Copy to clipboard"' in response.text + + def test_database_wizard_decorative_icons_hidden(self, client): + """Test that decorative icons have aria-hidden=true.""" + response = client.get("/database-wizard") + assert 'aria-hidden="true"' in response.text + + def test_database_wizard_reduced_motion(self, client): + """Test that wizard respects prefers-reduced-motion media query.""" + response = client.get("/database-wizard") + assert "prefers-reduced-motion" in response.text + + +@pytest.mark.unit +class TestSettingsPageWizardLink: + """Tests for the database wizard link on the settings page.""" + + def test_settings_template_has_db_wizard_link(self): + """Test that the settings template contains a link to the database wizard.""" + from pathlib import Path + + template_path = Path(__file__).resolve().parent.parent / "frontend" / "templates" / "settings.html" + content = template_path.read_text() + assert "/database-wizard" in content + assert "DB Wizard" in content + + def test_settings_template_has_help_link_rendering(self): + """Test that the settings template renders help_link metadata.""" + from pathlib import Path + + template_path = Path(__file__).resolve().parent.parent / "frontend" / "templates" / "settings.html" + content = template_path.read_text() + assert "setting.metadata.get('help_link')" in content + assert "help_link_label" in content + + def test_database_url_metadata_has_help_link(self): + """Test that database_url SETTING_METADATA includes help_link to wizard.""" + from app.utils.settings_service import SETTING_METADATA + + meta = SETTING_METADATA["database_url"] + assert "help_link" in meta + assert meta["help_link"] == "/database-wizard" + assert "help_link_label" in meta + + +# Context manager helper for tests that don't need session_transaction +class _NoOpContextManager: + """Dummy context manager for tests that don't need session_transaction.""" + + def __enter__(self): + return None + + def __exit__(self, *args): + pass