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

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

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

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-03-05 22:10:14 +00:00
parent 0f408f67b4
commit f6fcaaeccc
10 changed files with 1848 additions and 0 deletions
+2
View File
@@ -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)
+169
View File
@@ -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
+248
View File
@@ -0,0 +1,248 @@
"""
Database migration utility for transferring data between databases.
Copies all table rows from a *source* SQLAlchemy database to a *target*
database. This is designed for the common scenario of migrating from the
built-in SQLite database to an external PostgreSQL / MySQL instance.
The utility:
1. Creates the schema in the target via ``Base.metadata.create_all``.
2. Copies rows table-by-table in dependency order.
3. Stamps the Alembic version in the target to ``head``.
"""
import logging
from typing import Any
from sqlalchemy import MetaData, create_engine, inspect, text
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import make_url
from sqlalchemy.orm import sessionmaker
logger = logging.getLogger(__name__)
# Tables to skip during migration (Alembic manages its own state).
_SKIP_TABLES = {"alembic_version"}
# Ordered list — parent tables first to respect foreign-key constraints.
_TABLE_ORDER = [
"documents",
"files",
"file_processing_steps",
"processing_logs",
"application_settings",
"settings_audit_log",
"saved_searches",
"webhook_configs",
]
def _make_engine(url: str) -> Engine:
"""Create a SQLAlchemy engine from *url* with sensible defaults."""
parsed = make_url(url)
connect_args: dict[str, Any] = {}
if parsed.get_backend_name() == "sqlite":
connect_args["check_same_thread"] = False
return create_engine(url, connect_args=connect_args)
def _ordered_tables(inspector: Any) -> list[str]:
"""Return table names in safe insertion order.
Tables listed in ``_TABLE_ORDER`` come first (in that order); any
remaining tables are appended alphabetically.
"""
existing = set(inspector.get_table_names())
ordered: list[str] = []
for name in _TABLE_ORDER:
if name in existing and name not in _SKIP_TABLES:
ordered.append(name)
for name in sorted(existing):
if name not in ordered and name not in _SKIP_TABLES:
ordered.append(name)
return ordered
def preview_migration(source_url: str) -> dict[str, Any]:
"""Preview what a migration would do without actually copying data.
Args:
source_url: Connection string for the source database.
Returns:
Dict with ``tables`` (list of dicts with ``name`` and ``row_count``)
and ``total_rows``.
"""
try:
src_engine = _make_engine(source_url)
src_inspector = inspect(src_engine)
tables = _ordered_tables(src_inspector)
result: list[dict[str, Any]] = []
total = 0
with src_engine.connect() as conn:
for table_name in tables:
row = conn.execute(text(f'SELECT COUNT(*) FROM "{table_name}"')).fetchone() # noqa: S608
count = row[0] if row else 0
result.append({"name": table_name, "row_count": count})
total += count
src_engine.dispose()
return {"tables": result, "total_rows": total, "success": True}
except Exception as exc:
logger.error(f"Migration preview failed: {exc}")
return {"success": False, "error": str(exc), "tables": [], "total_rows": 0}
def migrate_data(
source_url: str,
target_url: str,
*,
batch_size: int = 500,
progress_callback: Any | None = None,
) -> dict[str, Any]:
"""Copy all data from *source_url* to *target_url*.
The target schema is created automatically from the application models.
Alembic is stamped to ``head`` in the target after a successful copy.
Args:
source_url: SQLAlchemy connection string for the source DB.
target_url: SQLAlchemy connection string for the target DB.
batch_size: Number of rows to insert per batch.
progress_callback: Optional ``callable(table_name, copied, total)``
invoked after each batch.
Returns:
Dict with ``success`` (bool), ``tables_copied`` (int),
``rows_copied`` (int), and ``errors`` (list of str).
"""
errors: list[str] = []
tables_copied = 0
rows_copied = 0
try:
src_engine = _make_engine(source_url)
tgt_engine = _make_engine(target_url)
# ------------------------------------------------------------------
# 1. Create schema in target from application models
# ------------------------------------------------------------------
from app.database import Base # local import to avoid circular deps
Base.metadata.create_all(bind=tgt_engine)
logger.info("Target schema created from application models.")
# ------------------------------------------------------------------
# 2. Reflect source schema & determine copy order
# ------------------------------------------------------------------
src_meta = MetaData()
src_meta.reflect(bind=src_engine)
src_inspector = inspect(src_engine)
table_names = _ordered_tables(src_inspector)
SrcSession = sessionmaker(bind=src_engine)
TgtSession = sessionmaker(bind=tgt_engine)
# ------------------------------------------------------------------
# 3. Copy data table-by-table
# ------------------------------------------------------------------
for table_name in table_names:
try:
src_session = SrcSession()
tgt_session = TgtSession()
src_table = src_meta.tables.get(table_name)
if src_table is None:
continue
# Read all rows from source
rows = src_session.execute(src_table.select()).fetchall()
column_names = [c.name for c in src_table.columns]
if not rows:
logger.info(f"Skipping empty table: {table_name}")
tables_copied += 1
src_session.close()
tgt_session.close()
continue
# Reflect the target table to insert into
tgt_meta = MetaData()
tgt_meta.reflect(bind=tgt_engine, only=[table_name])
tgt_table = tgt_meta.tables.get(table_name)
if tgt_table is None:
errors.append(f"Target table {table_name} not found after schema creation")
src_session.close()
tgt_session.close()
continue
# Batch insert
total_for_table = len(rows)
for i in range(0, total_for_table, batch_size):
batch = rows[i : i + batch_size]
insert_data = [dict(zip(column_names, row, strict=False)) for row in batch]
tgt_session.execute(tgt_table.insert(), insert_data)
tgt_session.commit()
rows_copied += len(batch)
if progress_callback:
progress_callback(table_name, min(i + batch_size, total_for_table), total_for_table)
tables_copied += 1
logger.info(f"Copied {total_for_table} rows from {table_name}")
src_session.close()
tgt_session.close()
except Exception as exc:
msg = f"Error copying table {table_name}: {exc}"
logger.error(msg)
errors.append(msg)
# ------------------------------------------------------------------
# 4. Stamp Alembic to head in the target
# ------------------------------------------------------------------
try:
_stamp_alembic_head(tgt_engine)
logger.info("Alembic version stamped to head in target database.")
except Exception as exc:
msg = f"Failed to stamp Alembic version: {exc}"
logger.error(msg)
errors.append(msg)
src_engine.dispose()
tgt_engine.dispose()
return {
"success": len(errors) == 0,
"tables_copied": tables_copied,
"rows_copied": rows_copied,
"errors": errors,
}
except Exception as exc:
logger.error(f"Migration failed: {exc}")
return {
"success": False,
"tables_copied": tables_copied,
"rows_copied": rows_copied,
"errors": errors + [str(exc)],
}
def _stamp_alembic_head(engine: Engine) -> None:
"""Stamp the Alembic version table to ``head`` in the given engine."""
from pathlib import Path
from alembic import command
from alembic.config import Config
migrations_dir = str(Path(__file__).resolve().parent.parent.parent / "migrations")
alembic_cfg = Config()
alembic_cfg.set_main_option("script_location", migrations_dir)
alembic_cfg.set_main_option("sqlalchemy.url", "")
with engine.begin() as connection:
alembic_cfg.attributes["connection"] = connection
command.stamp(alembic_cfg, "head")
+255
View File
@@ -0,0 +1,255 @@
"""
Database configuration wizard utilities.
Provides helpers for building, validating, and testing database connection
strings. Used by both the interactive wizard UI and the REST API.
"""
import logging
from typing import Any
from sqlalchemy import create_engine, text
from sqlalchemy.engine.url import make_url
logger = logging.getLogger(__name__)
# Supported database backends with human-readable labels and defaults.
SUPPORTED_BACKENDS: list[dict[str, Any]] = [
{
"id": "sqlite",
"label": "SQLite (Development)",
"driver": "",
"default_port": None,
"description": "File-based database. Best for development and single-user setups.",
"requires_host": False,
},
{
"id": "postgresql",
"label": "PostgreSQL (Recommended for Production)",
"driver": "",
"default_port": 5432,
"description": "Robust, full-featured database. Recommended for production.",
"requires_host": True,
},
{
"id": "mysql",
"label": "MySQL / MariaDB",
"driver": "pymysql",
"default_port": 3306,
"description": "Popular open-source database. Requires pymysql driver.",
"requires_host": True,
},
]
def get_supported_backends() -> list[dict[str, Any]]:
"""Return the list of supported database backends with metadata.
Returns:
List of backend descriptor dicts.
"""
return SUPPORTED_BACKENDS
def build_connection_string(
backend: str,
host: str = "",
port: int | None = None,
database: str = "",
username: str = "",
password: str = "",
ssl_mode: str = "",
extra_options: str = "",
sqlite_path: str = "",
) -> str:
"""Build a SQLAlchemy connection string from individual components.
Args:
backend: Database backend identifier (``sqlite``, ``postgresql``, ``mysql``).
host: Database server hostname or IP.
port: Database server port (uses backend default when ``None``).
database: Database / schema name.
username: Authentication username.
password: Authentication password.
ssl_mode: SSL mode (e.g. ``require``, ``verify-full``). PostgreSQL only.
extra_options: Additional query-string options appended to the URL.
sqlite_path: File path for SQLite databases.
Returns:
A SQLAlchemy-compatible connection URL string.
Raises:
ValueError: If required fields are missing for the chosen backend.
"""
if backend == "sqlite":
path = sqlite_path.strip() if sqlite_path else "./app/database.db"
return f"sqlite:///{path}"
# Resolve driver prefix
backend_info = next((b for b in SUPPORTED_BACKENDS if b["id"] == backend), None)
if backend_info is None:
raise ValueError(f"Unsupported backend: {backend}")
if not host:
raise ValueError("Host is required for non-SQLite backends")
if not database:
raise ValueError("Database name is required for non-SQLite backends")
if not username:
raise ValueError("Username is required for non-SQLite backends")
driver_suffix = f"+{backend_info['driver']}" if backend_info["driver"] else ""
scheme = f"{backend}{driver_suffix}"
resolved_port = port if port else backend_info["default_port"]
# Build query parameters
params: list[str] = []
if ssl_mode:
params.append(f"sslmode={ssl_mode}")
if extra_options:
params.append(extra_options)
if backend == "mysql" and "charset=" not in extra_options:
params.append("charset=utf8mb4")
query_string = "&".join(params)
# Construct URL
auth = username
if password:
auth = f"{username}:{password}"
url = f"{scheme}://{auth}@{host}:{resolved_port}/{database}"
if query_string:
url = f"{url}?{query_string}"
return url
def parse_connection_string(url: str) -> dict[str, Any]:
"""Parse a SQLAlchemy connection string into its components.
Args:
url: A SQLAlchemy database URL string.
Returns:
Dict with keys: ``backend``, ``host``, ``port``, ``database``,
``username``, ``password``, ``ssl_mode``, ``is_sqlite``.
"""
try:
parsed = make_url(url)
backend_name = parsed.get_backend_name()
return {
"backend": backend_name,
"host": parsed.host or "",
"port": parsed.port,
"database": parsed.database or "",
"username": parsed.username or "",
"password": parsed.password or "",
"ssl_mode": "",
"is_sqlite": backend_name == "sqlite",
"valid": True,
}
except Exception as exc:
logger.warning(f"Failed to parse connection string: {exc}")
return {"valid": False, "error": str(exc)}
def test_connection(url: str, timeout: int = 10) -> dict[str, Any]:
"""Attempt to connect to a database and return status information.
The function creates a short-lived engine, executes a simple ``SELECT 1``
query, and disposes the engine. It does **not** modify any global state.
Args:
url: SQLAlchemy database URL to test.
timeout: Connection timeout in seconds.
Returns:
Dict with ``success`` (bool), ``message`` (str), and optional
``server_version`` (str).
"""
try:
parsed = make_url(url)
backend = parsed.get_backend_name()
connect_args: dict[str, Any] = {}
kwargs: dict[str, Any] = {"pool_pre_ping": True}
if backend == "sqlite":
connect_args["check_same_thread"] = False
else:
kwargs["pool_timeout"] = timeout
test_engine = create_engine(
url,
connect_args=connect_args,
**kwargs,
)
with test_engine.connect() as conn:
result = conn.execute(text("SELECT 1"))
result.fetchone()
# Try to fetch server version for informational display
server_version = _get_server_version(conn, backend)
test_engine.dispose()
return {
"success": True,
"message": "Connection successful",
"backend": backend,
"server_version": server_version,
}
except Exception as exc:
logger.warning(f"Connection test failed: {exc}")
return {
"success": False,
"message": str(exc),
"backend": "",
"server_version": "",
}
def _get_server_version(conn: Any, backend: str) -> str:
"""Retrieve a human-readable server version string.
Args:
conn: An active SQLAlchemy connection.
backend: Backend identifier (``sqlite``, ``postgresql``, ``mysql``).
Returns:
Server version string, or empty string on failure.
"""
try:
if backend == "postgresql":
row = conn.execute(text("SELECT version()")).fetchone()
return str(row[0]) if row else ""
elif backend == "mysql":
row = conn.execute(text("SELECT version()")).fetchone()
return str(row[0]) if row else ""
elif backend == "sqlite":
row = conn.execute(text("SELECT sqlite_version()")).fetchone()
return f"SQLite {row[0]}" if row else ""
except Exception:
logger.debug("Could not retrieve server version")
return ""
def validate_url_format(url: str) -> dict[str, Any]:
"""Validate that a connection string is syntactically correct.
Args:
url: The connection string to validate.
Returns:
Dict with ``valid`` (bool) and optional ``error`` (str).
"""
try:
parsed = make_url(url)
backend = parsed.get_backend_name()
if backend not in ("sqlite", "postgresql", "mysql"):
return {"valid": False, "error": f"Unsupported backend: {backend}"}
return {"valid": True, "backend": backend}
except Exception as exc:
return {"valid": False, "error": str(exc)}
+2
View File
@@ -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)
+28
View File
@@ -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,
},
)
+606
View File
@@ -0,0 +1,606 @@
{% extends "base.html" %}
{% block title %}Database Configuration Wizard - DocuElevate{% endblock %}
{% block head_extra %}
<style>
.wizard-step { display: none; }
.wizard-step.active { display: block; }
.db-card { transition: all 0.2s ease; }
.db-card.selected { border-color: #4f46e5; box-shadow: 0 0 0 3px rgba(79,70,229,0.3); }
.db-card:hover { border-color: #6366f1; }
.mono-input { font-family: 'Courier New', monospace; }
.migration-table { border-collapse: collapse; }
.migration-table th, .migration-table td { padding: 0.5rem 1rem; text-align: left; }
@keyframes pulse-bar { 0%,100%{opacity:1} 50%{opacity:.5} }
.animate-pulse-bar { animation: pulse-bar 1.5s ease-in-out infinite; }
</style>
{% endblock %}
{% block content %}
<div class="min-h-screen bg-gradient-to-br from-blue-50 to-indigo-100 py-12 px-4 sm:px-6 lg:px-8"
x-data="dbWizard()" x-init="init()">
<div class="max-w-4xl mx-auto">
{# ── Header ── #}
<div class="text-center mb-8">
<h1 class="text-4xl font-bold text-gray-900 mb-2">
<i class="fas fa-database text-indigo-600" aria-hidden="true"></i>
Database Configuration Wizard
</h1>
<p class="text-lg text-gray-600">
Configure a new database connection or migrate your data to an external database.
</p>
</div>
{# ── Tab Navigation ── #}
<div class="flex justify-center mb-8 space-x-4" role="tablist" aria-label="Wizard tabs">
<button @click="activeTab='configure'"
:class="activeTab==='configure' ? 'bg-indigo-600 text-white' : 'bg-white text-gray-700 hover:bg-gray-100'"
class="px-6 py-3 rounded-lg font-semibold shadow transition"
role="tab" :aria-selected="activeTab==='configure'" id="tab-configure" aria-controls="panel-configure">
<i class="fas fa-cog mr-2" aria-hidden="true"></i> Configure Database
</button>
<button @click="activeTab='migrate'"
:class="activeTab==='migrate' ? 'bg-indigo-600 text-white' : 'bg-white text-gray-700 hover:bg-gray-100'"
class="px-6 py-3 rounded-lg font-semibold shadow transition"
role="tab" :aria-selected="activeTab==='migrate'" id="tab-migrate" aria-controls="panel-migrate">
<i class="fas fa-exchange-alt mr-2" aria-hidden="true"></i> Migrate Data
</button>
</div>
{# ═══════════════════════════════════════════════════════════════════ #}
{# TAB 1 — Configure #}
{# ═══════════════════════════════════════════════════════════════════ #}
<div x-show="activeTab==='configure'" role="tabpanel" id="panel-configure" aria-labelledby="tab-configure">
<div class="bg-white rounded-lg shadow-xl overflow-hidden">
{# Step indicator #}
<div class="bg-indigo-600 px-6 py-4">
<h2 class="text-2xl font-bold text-white">
<span x-show="cfgStep===1"><i class="fas fa-server mr-2" aria-hidden="true"></i> Step 1: Choose Database Type</span>
<span x-show="cfgStep===2"><i class="fas fa-plug mr-2" aria-hidden="true"></i> Step 2: Connection Details</span>
<span x-show="cfgStep===3"><i class="fas fa-check-circle mr-2" aria-hidden="true"></i> Step 3: Test &amp; Apply</span>
</h2>
<div class="flex mt-3 space-x-2">
<template x-for="s in [1,2,3]" :key="s">
<div class="h-2 flex-1 rounded-full"
:class="s <= cfgStep ? 'bg-white' : 'bg-indigo-400'"></div>
</template>
</div>
</div>
<div class="px-6 py-8">
{# ── Step 1: Choose backend ── #}
<div x-show="cfgStep===1">
<p class="text-gray-600 mb-6">Select the database engine you want to use.</p>
<div class="grid grid-cols-1 md:grid-cols-3 gap-4">
<button type="button" @click="selectBackend('sqlite')"
:class="form.backend==='sqlite' ? 'selected' : ''"
class="db-card border-2 rounded-lg p-6 text-left focus:outline-none focus:ring-2 focus:ring-indigo-500"
aria-label="Select SQLite">
<div class="text-3xl mb-2 text-yellow-600"><i class="fas fa-file-alt" aria-hidden="true"></i></div>
<h3 class="font-bold text-lg">SQLite</h3>
<p class="text-sm text-gray-500 mt-1">File-based. Best for development &amp; single-user setups.</p>
</button>
<button type="button" @click="selectBackend('postgresql')"
:class="form.backend==='postgresql' ? 'selected' : ''"
class="db-card border-2 rounded-lg p-6 text-left focus:outline-none focus:ring-2 focus:ring-indigo-500"
aria-label="Select PostgreSQL">
<div class="text-3xl mb-2 text-blue-600"><i class="fas fa-elephant" aria-hidden="true"></i></div>
<h3 class="font-bold text-lg">PostgreSQL</h3>
<p class="text-sm text-gray-500 mt-1">Recommended for production. Full feature support.</p>
<span class="inline-block mt-2 px-2 py-0.5 bg-green-100 text-green-800 text-xs rounded font-semibold">Recommended</span>
</button>
<button type="button" @click="selectBackend('mysql')"
:class="form.backend==='mysql' ? 'selected' : ''"
class="db-card border-2 rounded-lg p-6 text-left focus:outline-none focus:ring-2 focus:ring-indigo-500"
aria-label="Select MySQL / MariaDB">
<div class="text-3xl mb-2 text-orange-600"><i class="fas fa-database" aria-hidden="true"></i></div>
<h3 class="font-bold text-lg">MySQL / MariaDB</h3>
<p class="text-sm text-gray-500 mt-1">Popular alternative. Requires pymysql driver.</p>
</button>
</div>
<div class="flex justify-end mt-8">
<button @click="cfgStep=2" :disabled="!form.backend"
class="px-6 py-3 bg-indigo-600 text-white rounded-md hover:bg-indigo-700 disabled:opacity-50 disabled:cursor-not-allowed shadow-lg focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500">
Next <i class="fas fa-arrow-right ml-2" aria-hidden="true"></i>
</button>
</div>
</div>
{# ── Step 2: Connection details ── #}
<div x-show="cfgStep===2">
{# SQLite path #}
<template x-if="form.backend==='sqlite'">
<div class="space-y-4">
<label for="sqlite_path" class="block text-sm font-medium text-gray-700">Database File Path</label>
<input id="sqlite_path" x-model="form.sqlite_path" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md shadow-sm focus:ring-2 focus:ring-indigo-500 focus:border-transparent"
placeholder="./app/database.db" />
<p class="text-xs text-gray-500">Leave blank to use the default path <code>./app/database.db</code>.</p>
</div>
</template>
{# Host-based databases #}
<template x-if="form.backend!=='sqlite'">
<div class="space-y-4">
<div class="grid grid-cols-1 md:grid-cols-2 gap-4">
<div>
<label for="db_host" class="block text-sm font-medium text-gray-700">Host <span class="text-red-600">*</span></label>
<input id="db_host" x-model="form.host" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
placeholder="localhost or my-db.rds.amazonaws.com" required />
</div>
<div>
<label for="db_port" class="block text-sm font-medium text-gray-700">Port</label>
<input id="db_port" x-model.number="form.port" type="number"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
:placeholder="form.backend==='postgresql' ? '5432' : '3306'" />
</div>
</div>
<div>
<label for="db_name" class="block text-sm font-medium text-gray-700">Database Name <span class="text-red-600">*</span></label>
<input id="db_name" x-model="form.database" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
placeholder="docuelevate" required />
</div>
<div class="grid grid-cols-1 md:grid-cols-2 gap-4">
<div>
<label for="db_user" class="block text-sm font-medium text-gray-700">Username <span class="text-red-600">*</span></label>
<input id="db_user" x-model="form.username" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
placeholder="docuelevate" required />
</div>
<div>
<label for="db_pass" class="block text-sm font-medium text-gray-700">Password</label>
<input id="db_pass" x-model="form.password" type="password"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
placeholder="••••••••" />
</div>
</div>
<div x-show="form.backend==='postgresql'">
<label for="ssl_mode" class="block text-sm font-medium text-gray-700">SSL Mode</label>
<select id="ssl_mode" x-model="form.ssl_mode"
class="w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500">
<option value="">None (disable)</option>
<option value="require">require</option>
<option value="verify-ca">verify-ca</option>
<option value="verify-full">verify-full</option>
</select>
<p class="text-xs text-gray-500 mt-1">Use <code>require</code> for managed cloud databases (AWS RDS, Supabase, etc.).</p>
</div>
</div>
</template>
{# Live preview of the URL #}
<div class="mt-6 p-4 bg-gray-50 rounded-lg border">
<label class="block text-sm font-medium text-gray-700 mb-1">Generated Connection String</label>
<code class="block text-sm break-all text-indigo-700" x-text="builtUrl || '(fill in details above)'"></code>
</div>
<div class="flex justify-between mt-8">
<button @click="cfgStep=1"
class="px-6 py-3 border border-gray-300 text-gray-700 rounded-md hover:bg-gray-50 focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500">
<i class="fas fa-arrow-left mr-2" aria-hidden="true"></i> Back
</button>
<button @click="cfgStep=3; buildUrl()"
class="px-6 py-3 bg-indigo-600 text-white rounded-md hover:bg-indigo-700 shadow-lg focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500">
Next <i class="fas fa-arrow-right ml-2" aria-hidden="true"></i>
</button>
</div>
</div>
{# ── Step 3: Test & Apply ── #}
<div x-show="cfgStep===3">
<div class="mb-6 p-4 bg-gray-50 rounded-lg border">
<label class="block text-sm font-medium text-gray-700 mb-1">Connection String</label>
<code class="block text-sm break-all text-indigo-700" x-text="builtUrl"></code>
</div>
{# Test button #}
<div class="flex items-center space-x-4 mb-6">
<button @click="testConnection()"
:disabled="testLoading"
class="px-6 py-3 bg-green-600 text-white rounded-md hover:bg-green-700 disabled:opacity-50 shadow-lg focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-green-500">
<i class="fas fa-plug mr-2" aria-hidden="true"></i>
<span x-text="testLoading ? 'Testing…' : 'Test Connection'"></span>
</button>
<div x-show="testResult" class="flex-1">
<div x-show="testResult?.success" class="bg-green-100 border-l-4 border-green-500 text-green-800 p-3 rounded" role="status">
<i class="fas fa-check-circle mr-1" aria-hidden="true"></i>
<span x-text="testResult?.message"></span>
<span x-show="testResult?.server_version" class="block text-xs mt-1" x-text="'Server: ' + testResult?.server_version"></span>
</div>
<div x-show="!testResult?.success" class="bg-red-100 border-l-4 border-red-500 text-red-800 p-3 rounded" role="alert">
<i class="fas fa-times-circle mr-1" aria-hidden="true"></i>
<span x-text="testResult?.message"></span>
</div>
</div>
</div>
{# Apply as DATABASE_URL #}
<div class="bg-amber-50 border-l-4 border-amber-400 p-4 mb-6">
<p class="text-sm text-amber-800">
<i class="fas fa-exclamation-triangle mr-1" aria-hidden="true"></i>
<strong>To use this database</strong>, set the <code>DATABASE_URL</code> environment variable
(in your <code>.env</code> file or Docker Compose config) to the connection string above,
then restart DocuElevate.
</p>
</div>
<div class="p-4 border rounded-lg bg-gray-50">
<label for="env_snippet" class="block text-sm font-medium text-gray-700 mb-2">
Copy this into your <code>.env</code> file:
</label>
<div class="flex items-center">
<input id="env_snippet" type="text" readonly
:value="'DATABASE_URL=' + builtUrl"
class="mono-input flex-1 px-4 py-3 border border-gray-300 rounded-l-md bg-white text-sm" />
<button @click="copyToClipboard('DATABASE_URL=' + builtUrl)"
class="px-4 py-3 bg-indigo-600 text-white rounded-r-md hover:bg-indigo-700 focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500"
aria-label="Copy to clipboard">
<i class="fas fa-copy" aria-hidden="true"></i>
</button>
</div>
<p x-show="copied" x-transition class="text-xs text-green-600 mt-1" role="status">
<i class="fas fa-check" aria-hidden="true"></i> Copied!
</p>
</div>
<div class="flex justify-between mt-8">
<button @click="cfgStep=2"
class="px-6 py-3 border border-gray-300 text-gray-700 rounded-md hover:bg-gray-50 focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500">
<i class="fas fa-arrow-left mr-2" aria-hidden="true"></i> Back
</button>
<a href="/settings"
class="px-6 py-3 bg-indigo-600 text-white rounded-md hover:bg-indigo-700 shadow-lg focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500 inline-flex items-center">
Go to Settings <i class="fas fa-arrow-right ml-2" aria-hidden="true"></i>
</a>
</div>
</div>
</div>
</div>
</div>
{# ═══════════════════════════════════════════════════════════════════ #}
{# TAB 2 — Migrate #}
{# ═══════════════════════════════════════════════════════════════════ #}
<div x-show="activeTab==='migrate'" role="tabpanel" id="panel-migrate" aria-labelledby="tab-migrate">
<div class="bg-white rounded-lg shadow-xl overflow-hidden">
<div class="bg-indigo-600 px-6 py-4">
<h2 class="text-2xl font-bold text-white">
<i class="fas fa-exchange-alt mr-2" aria-hidden="true"></i>
Migrate Data Between Databases
</h2>
<p class="text-indigo-100 mt-1">
Copy all your data from one database to another (e.g. SQLite → PostgreSQL).
</p>
</div>
<div class="px-6 py-8 space-y-6">
{# Source URL #}
<div>
<label for="mig_source" class="block text-sm font-medium text-gray-700 mb-1">
Source Database URL
</label>
<input id="mig_source" x-model="migrate.source" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
:placeholder="currentDbUrl || 'sqlite:///./app/database.db'" />
<p class="text-xs text-gray-500 mt-1">
This is your current database. Pre-filled with the running configuration.
</p>
<button @click="migrate.source = currentDbUrl"
class="mt-1 text-xs text-indigo-600 hover:underline focus:outline-none"
type="button">
<i class="fas fa-sync-alt mr-1" aria-hidden="true"></i> Use current database
</button>
</div>
{# Target URL #}
<div>
<label for="mig_target" class="block text-sm font-medium text-gray-700 mb-1">
Target Database URL
</label>
<input id="mig_target" x-model="migrate.target" type="text"
class="mono-input w-full px-4 py-3 border border-gray-300 rounded-md focus:ring-2 focus:ring-indigo-500"
placeholder="postgresql://user:pass@host:5432/docuelevate" />
<p class="text-xs text-gray-500 mt-1">
The new database to copy data into. Must be empty (schema will be created automatically).
</p>
<button @click="if(builtUrl) migrate.target = builtUrl"
x-show="builtUrl"
class="mt-1 text-xs text-indigo-600 hover:underline focus:outline-none"
type="button">
<i class="fas fa-magic mr-1" aria-hidden="true"></i> Use URL from Configure tab
</button>
</div>
{# Actions #}
<div class="flex space-x-4">
<button @click="testMigrateConnection('source')"
:disabled="!migrate.source || migrateLoading"
class="px-4 py-2 bg-gray-600 text-white rounded-md hover:bg-gray-700 disabled:opacity-50 text-sm focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-gray-500">
<i class="fas fa-plug mr-1" aria-hidden="true"></i> Test Source
</button>
<button @click="testMigrateConnection('target')"
:disabled="!migrate.target || migrateLoading"
class="px-4 py-2 bg-gray-600 text-white rounded-md hover:bg-gray-700 disabled:opacity-50 text-sm focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-gray-500">
<i class="fas fa-plug mr-1" aria-hidden="true"></i> Test Target
</button>
<button @click="previewMigration()"
:disabled="!migrate.source || migrateLoading"
class="px-4 py-2 bg-blue-600 text-white rounded-md hover:bg-blue-700 disabled:opacity-50 text-sm focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-blue-500">
<i class="fas fa-search mr-1" aria-hidden="true"></i> Preview Migration
</button>
</div>
{# Connection test results #}
<div x-show="migrate.sourceTest" class="text-sm" role="status">
<span :class="migrate.sourceTest?.success ? 'text-green-700' : 'text-red-700'">
<i :class="migrate.sourceTest?.success ? 'fas fa-check-circle' : 'fas fa-times-circle'" aria-hidden="true"></i>
Source: <span x-text="migrate.sourceTest?.message"></span>
</span>
</div>
<div x-show="migrate.targetTest" class="text-sm" role="status">
<span :class="migrate.targetTest?.success ? 'text-green-700' : 'text-red-700'">
<i :class="migrate.targetTest?.success ? 'fas fa-check-circle' : 'fas fa-times-circle'" aria-hidden="true"></i>
Target: <span x-text="migrate.targetTest?.message"></span>
</span>
</div>
{# Preview table #}
<div x-show="migrate.preview" class="border rounded-lg overflow-hidden">
<table class="migration-table w-full text-sm">
<thead class="bg-gray-100">
<tr>
<th scope="col">Table</th>
<th scope="col" class="text-right">Rows</th>
</tr>
</thead>
<tbody>
<template x-for="t in migrate.preview?.tables || []" :key="t.name">
<tr class="border-t">
<td x-text="t.name"></td>
<td class="text-right" x-text="t.row_count.toLocaleString()"></td>
</tr>
</template>
<tr class="border-t font-bold bg-gray-50">
<td>Total</td>
<td class="text-right" x-text="(migrate.preview?.total_rows || 0).toLocaleString()"></td>
</tr>
</tbody>
</table>
</div>
{# Migrate button #}
<div x-show="migrate.preview && migrate.preview.success" class="mt-4">
<div class="bg-amber-50 border-l-4 border-amber-400 p-4 mb-4">
<p class="text-sm text-amber-800">
<i class="fas fa-exclamation-triangle mr-1" aria-hidden="true"></i>
<strong>Warning:</strong> This will copy all data to the target database.
The target must be empty. This operation cannot be undone.
</p>
</div>
<label class="inline-flex items-center mb-4">
<input type="checkbox" x-model="migrate.confirmed" class="form-checkbox text-indigo-600 rounded" />
<span class="ml-2 text-sm">I understand and want to proceed with the migration</span>
</label>
<div>
<button @click="executeMigration()"
:disabled="!migrate.confirmed || migrateLoading"
class="px-6 py-3 bg-red-600 text-white rounded-md hover:bg-red-700 disabled:opacity-50 shadow-lg focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-red-500">
<i class="fas fa-play mr-2" aria-hidden="true"></i>
<span x-text="migrateLoading ? 'Migrating…' : 'Start Migration'"></span>
</button>
</div>
</div>
{# Migration progress / result #}
<div x-show="migrateLoading" class="mt-4" role="status" aria-live="polite">
<div class="w-full bg-gray-200 rounded-full h-3">
<div class="bg-indigo-600 h-3 rounded-full animate-pulse-bar" style="width:100%"></div>
</div>
<p class="text-sm text-gray-600 mt-2">Migration in progress — please do not close this page…</p>
</div>
<div x-show="migrate.result" class="mt-4" role="status" aria-live="polite">
<div x-show="migrate.result?.success" class="bg-green-100 border-l-4 border-green-500 text-green-800 p-4 rounded">
<p class="font-bold"><i class="fas fa-check-circle mr-1" aria-hidden="true"></i> Migration Successful</p>
<p class="text-sm mt-1">
Copied <strong x-text="migrate.result?.rows_copied?.toLocaleString()"></strong> rows
across <strong x-text="migrate.result?.tables_copied"></strong> tables.
</p>
<p class="text-sm mt-2">
Update your <code>DATABASE_URL</code> environment variable to the target URL and restart DocuElevate.
</p>
</div>
<div x-show="!migrate.result?.success" class="bg-red-100 border-l-4 border-red-500 text-red-800 p-4 rounded">
<p class="font-bold"><i class="fas fa-times-circle mr-1" aria-hidden="true"></i> Migration Failed</p>
<template x-for="err in migrate.result?.errors || migrate.result?.detail?.errors || []" :key="err">
<p class="text-sm mt-1" x-text="err"></p>
</template>
<p x-show="migrate.result?.detail && typeof migrate.result.detail === 'string'" class="text-sm mt-1" x-text="migrate.result?.detail"></p>
</div>
</div>
</div>
</div>
</div>
{# ── Help text ── #}
<div class="mt-6 text-center">
<p class="text-sm text-gray-600">
<i class="fas fa-book mr-1" aria-hidden="true"></i>
See the <a href="/docs/DatabaseConfiguration" class="text-indigo-600 hover:underline">Database Configuration Guide</a> for more details.
</p>
</div>
</div>
</div>
<script>
function dbWizard() {
const csrfToken = document.querySelector('meta[name="csrf-token"]')?.content || '';
return {
activeTab: 'configure',
cfgStep: 1,
currentDbUrl: '{{ current_database_url | default("", true) }}',
form: {
backend: '',
host: '',
port: null,
database: '',
username: '',
password: '',
ssl_mode: '',
sqlite_path: '',
},
builtUrl: '',
testLoading: false,
testResult: null,
copied: false,
migrate: {
source: '{{ current_database_url | default("", true) }}',
target: '',
sourceTest: null,
targetTest: null,
preview: null,
confirmed: false,
result: null,
},
migrateLoading: false,
init() {
this.$watch('form', () => this.buildUrlLocal(), {deep: true});
},
selectBackend(b) {
this.form.backend = b;
// Set sensible defaults
if (b === 'postgresql') { this.form.port = 5432; }
else if (b === 'mysql') { this.form.port = 3306; }
else { this.form.port = null; }
},
buildUrlLocal() {
const f = this.form;
if (f.backend === 'sqlite') {
this.builtUrl = 'sqlite:///' + (f.sqlite_path || './app/database.db');
return;
}
if (!f.backend || !f.host || !f.database || !f.username) {
this.builtUrl = '';
return;
}
let driver = f.backend === 'mysql' ? 'mysql+pymysql' : f.backend;
let auth = f.username;
if (f.password) auth += ':' + f.password;
let port = f.port || (f.backend === 'postgresql' ? 5432 : 3306);
let url = driver + '://' + auth + '@' + f.host + ':' + port + '/' + f.database;
let params = [];
if (f.ssl_mode) params.push('sslmode=' + f.ssl_mode);
if (f.backend === 'mysql') params.push('charset=utf8mb4');
if (params.length) url += '?' + params.join('&');
this.builtUrl = url;
},
async buildUrl() {
this.buildUrlLocal();
},
async testConnection() {
this.testLoading = true;
this.testResult = null;
try {
const res = await fetch('/api/database/test-connection', {
method: 'POST',
headers: {'Content-Type': 'application/json', 'X-CSRF-Token': csrfToken},
body: JSON.stringify({url: this.builtUrl})
});
this.testResult = await res.json();
} catch (e) {
this.testResult = {success: false, message: 'Request failed: ' + e.message};
}
this.testLoading = false;
},
async testMigrateConnection(which) {
this.migrateLoading = true;
const url = which === 'source' ? this.migrate.source : this.migrate.target;
try {
const res = await fetch('/api/database/test-connection', {
method: 'POST',
headers: {'Content-Type': 'application/json', 'X-CSRF-Token': csrfToken},
body: JSON.stringify({url})
});
const data = await res.json();
if (which === 'source') this.migrate.sourceTest = data;
else this.migrate.targetTest = data;
} catch (e) {
const err = {success: false, message: 'Request failed: ' + e.message};
if (which === 'source') this.migrate.sourceTest = err;
else this.migrate.targetTest = err;
}
this.migrateLoading = false;
},
async previewMigration() {
this.migrateLoading = true;
this.migrate.preview = null;
try {
const res = await fetch('/api/database/preview-migration', {
method: 'POST',
headers: {'Content-Type': 'application/json', 'X-CSRF-Token': csrfToken},
body: JSON.stringify({url: this.migrate.source})
});
this.migrate.preview = await res.json();
} catch (e) {
this.migrate.preview = {success: false, error: e.message, tables: [], total_rows: 0};
}
this.migrateLoading = false;
},
async executeMigration() {
this.migrateLoading = true;
this.migrate.result = null;
try {
const res = await fetch('/api/database/migrate', {
method: 'POST',
headers: {'Content-Type': 'application/json', 'X-CSRF-Token': csrfToken},
body: JSON.stringify({source_url: this.migrate.source, target_url: this.migrate.target})
});
this.migrate.result = await res.json();
} catch (e) {
this.migrate.result = {success: false, errors: ['Request failed: ' + e.message]};
}
this.migrateLoading = false;
},
copyToClipboard(text) {
navigator.clipboard.writeText(text).then(() => {
this.copied = true;
setTimeout(() => this.copied = false, 2000);
});
}
};
}
</script>
{% endblock %}
+149
View File
@@ -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"])
+233
View File
@@ -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
+156
View File
@@ -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