{% if show_oauth %}
From 4857203d08966e5c470517255f8e2d74854cbbc1 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Sun, 8 Mar 2026 09:49:29 +0000
Subject: [PATCH 05/25] Initial plan
From 3aa5364e0ca3eaceb37616bb9b3a9a55fc08b223 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Sun, 8 Mar 2026 09:51:23 +0000
Subject: [PATCH 06/25] fix(auth): return 401 for API paths in require_login to
prevent wrong post-login redirect
The common.js fetch('/api/auth/whoami') probe on every page load was
overwriting the redirect_after_login session key with the API endpoint URL.
After login, users were sent to the JSON endpoint instead of the original page.
Fix: require_login now returns HTTP 401 for any /api/* path, consistent
with REST conventions, and never stores API URLs as the post-login redirect.
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
---
app/auth.py | 13 +++++++++
tests/test_auth.py | 55 +++++++++++++++++++++++++++++++++++++++
tests/test_auth_module.py | 27 +++++++++++++++++++
3 files changed, 95 insertions(+)
diff --git a/app/auth.py b/app/auth.py
index d6b2ec39..f570bc2a 100644
--- a/app/auth.py
+++ b/app/auth.py
@@ -3,9 +3,11 @@ import inspect
import logging
import pathlib
from functools import wraps
+from urllib.parse import urlparse
from authlib.integrations.starlette_client import OAuth
from fastapi import APIRouter, Depends, Request, status
+from fastapi.responses import JSONResponse
from fastapi.templating import Jinja2Templates
from sqlalchemy.orm import Session
from starlette.responses import RedirectResponse
@@ -81,6 +83,17 @@ def require_login(func):
@wraps(func)
async def wrapper(request: Request, *args, **kwargs):
if not request.session.get("user"):
+ # For API endpoints return 401 instead of storing the URL in the session
+ # and redirecting to /login. Without this guard, the /api/auth/whoami
+ # probe issued by common.js on every page load would overwrite
+ # redirect_after_login with the API URL, causing the post-login redirect
+ # to land on a JSON endpoint rather than the original page.
+ url_path = urlparse(str(request.url)).path
+ if url_path.startswith("/api/"):
+ return JSONResponse(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ content={"error": "Not authenticated"},
+ )
request.session["redirect_after_login"] = str(request.url)
return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)
# Check if the wrapped function is a coroutine function
diff --git a/tests/test_auth.py b/tests/test_auth.py
index c16fe318..81b6276e 100644
--- a/tests/test_auth.py
+++ b/tests/test_auth.py
@@ -175,6 +175,61 @@ class TestRequireLogin:
assert result["message"] == "sync"
assert result["param"] == "test_value"
+ @pytest.mark.asyncio
+ async def test_returns_401_for_api_paths_when_not_authenticated(self):
+ """Test that require_login returns 401 (not redirect) for /api/* paths.
+
+ This prevents the /api/auth/whoami JS probe from overwriting
+ redirect_after_login with an API URL, which would send the user to a
+ JSON endpoint after login instead of the page they actually wanted.
+ """
+ from fastapi.responses import JSONResponse
+
+ with patch("app.auth.AUTH_ENABLED", True):
+ from app.auth import require_login
+
+ @require_login
+ async def api_endpoint(request: Request):
+ return {"message": "success"}
+
+ mock_request = MagicMock(spec=Request)
+ mock_request.session = {}
+ mock_request.url = MagicMock()
+ mock_request.url.__str__ = MagicMock(return_value="http://test.com/api/auth/whoami")
+
+ result = await api_endpoint(mock_request)
+
+ assert isinstance(result, JSONResponse)
+ assert result.status_code == status.HTTP_401_UNAUTHORIZED
+ # Redirect URL must NOT be stored for API paths
+ assert "redirect_after_login" not in mock_request.session
+
+ @pytest.mark.asyncio
+ async def test_does_not_save_redirect_for_api_paths(self):
+ """Test that redirect_after_login is never set for any /api/* request."""
+ from fastapi.responses import JSONResponse
+
+ with patch("app.auth.AUTH_ENABLED", True):
+ from app.auth import require_login
+
+ @require_login
+ async def api_endpoint(request: Request):
+ return {"data": "ok"}
+
+ for api_path in ["/api/documents/upload", "/api/v1/resource", "/api/users/me"]:
+ mock_request = MagicMock(spec=Request)
+ mock_request.session = {}
+ mock_request.url = MagicMock()
+ mock_request.url.__str__ = MagicMock(return_value=f"http://test.com{api_path}")
+
+ result = await api_endpoint(mock_request)
+
+ assert isinstance(result, JSONResponse), f"Expected JSONResponse for {api_path}"
+ assert result.status_code == status.HTTP_401_UNAUTHORIZED
+ assert "redirect_after_login" not in mock_request.session, (
+ f"redirect_after_login must not be set for {api_path}"
+ )
+
@pytest.mark.integration
class TestWhoamiEndpoint:
diff --git a/tests/test_auth_module.py b/tests/test_auth_module.py
index d08830ff..233dfd5e 100644
--- a/tests/test_auth_module.py
+++ b/tests/test_auth_module.py
@@ -185,6 +185,33 @@ class TestRequireLogin:
assert isinstance(result, RedirectResponse)
assert result.status_code == status.HTTP_302_FOUND
+ @pytest.mark.asyncio
+ async def test_returns_401_for_api_path_when_not_authenticated(self):
+ """Test returns 401 for /api/* paths instead of redirect-to-login.
+
+ Prevents the common.js /api/auth/whoami probe from overwriting
+ redirect_after_login, which would send the user to a JSON endpoint
+ after login instead of the page they originally requested.
+ """
+ from fastapi.responses import JSONResponse
+
+ with patch("app.auth.AUTH_ENABLED", True):
+
+ @require_login
+ async def test_api_endpoint(request: Request):
+ return {"data": "ok"}
+
+ mock_request = MagicMock(spec=Request)
+ mock_request.session = {}
+ mock_request.url = MagicMock()
+ mock_request.url.__str__ = MagicMock(return_value="http://localhost/api/auth/whoami")
+
+ result = await test_api_endpoint(mock_request)
+
+ assert isinstance(result, JSONResponse)
+ assert result.status_code == status.HTTP_401_UNAUTHORIZED
+ assert "redirect_after_login" not in mock_request.session
+
@pytest.mark.unit
class TestOAuthConfiguration:
From 9b45ba62ba1e9f491765bd41304df191f2042104 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Sun, 8 Mar 2026 09:53:39 +0000
Subject: [PATCH 07/25] Initial plan
From a0f5ba179978d7c63bbdb8c564a32eaf3f78e5ad Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Sun, 8 Mar 2026 09:56:58 +0000
Subject: [PATCH 08/25] feat(backup): extend backup and restore to PostgreSQL
and MySQL/MariaDB
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
---
app/api/backup.py | 152 ++++++-----
app/tasks/backup_tasks.py | 346 ++++++++++++++++++++++++-
docs/ConfigurationGuide.md | 6 +-
docs/DatabaseConfiguration.md | 80 +++++-
tests/test_backup.py | 461 ++++++++++++++++++++++++++++++++--
5 files changed, 941 insertions(+), 104 deletions(-)
diff --git a/app/api/backup.py b/app/api/backup.py
index 2c83058c..8262814c 100644
--- a/app/api/backup.py
+++ b/app/api/backup.py
@@ -117,88 +117,104 @@ async def restore_backup(
"""Restore the database from an uploaded gzip-compressed SQL dump.
**Warning**: This overwrites the current database contents.
- Only SQLite databases are supported.
- The uploaded file must be a ``.db.gz`` file produced by the DocuElevate
- backup task (a gzip-compressed SQLite ``.dump()`` SQL script).
+ Supported formats (must match the currently configured database backend):
+
+ - ``*.db.gz`` – gzip-compressed SQLite ``.dump()`` SQL script (SQLite backend)
+ - ``*.pgsql.gz`` – gzip-compressed ``pg_dump --format=plain`` output (PostgreSQL backend)
+ - ``*.mysql.gz`` – gzip-compressed ``mysqldump`` output (MySQL / MariaDB backend)
"""
- from app.tasks.backup_tasks import _db_path
-
- db_path = _db_path()
- if db_path is None:
- raise HTTPException(
- status_code=status.HTTP_400_BAD_REQUEST,
- detail="Restore is only supported for SQLite databases.",
- )
-
- if not file.filename or not file.filename.endswith(".db.gz"):
- raise HTTPException(
- status_code=status.HTTP_400_BAD_REQUEST,
- detail="Uploaded file must be a .db.gz backup archive.",
- )
-
- import gzip
- import sqlite3
import tempfile
from pathlib import Path
- # Write the upload to a temp file first so we can validate it
- with tempfile.NamedTemporaryFile(suffix=".db.gz", delete=False) as tmp:
+ from sqlalchemy.engine.url import make_url
+
+ from app.config import settings as app_settings
+ from app.tasks.backup_tasks import (
+ _archive_ext_for_backend,
+ _db_path,
+ _restore_mysql,
+ _restore_postgresql,
+ _restore_sqlite,
+ )
+
+ url = make_url(app_settings.database_url)
+ backend = url.get_backend_name()
+ expected_ext = _archive_ext_for_backend(backend)
+
+ if not file.filename or not file.filename.endswith(expected_ext):
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail=(
+ f"Uploaded file must be a '{expected_ext}' backup archive for the current database backend ({backend})."
+ ),
+ )
+
+ # Write upload to a temp file
+ with tempfile.NamedTemporaryFile(suffix=expected_ext, delete=False) as tmp:
tmp_path = Path(tmp.name)
content = await file.read()
tmp.write(content)
try:
- # Decompress and read SQL statements
- with gzip.open(str(tmp_path), "rt", encoding="utf-8") as gz:
- sql_script = gz.read()
- except Exception as exc:
- tmp_path.unlink(missing_ok=True)
- raise HTTPException(
- status_code=status.HTTP_400_BAD_REQUEST,
- detail=f"Failed to decompress backup file: {exc}",
- ) from exc
+ if backend == "sqlite":
+ db_path = _db_path()
+ if db_path is None:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail="Restore is only supported for file-based SQLite databases.",
+ )
+ # Close the application DB session before replacing the file
+ db.close()
+ try:
+ _restore_sqlite(db_path, tmp_path)
+ except ValueError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail=str(exc),
+ ) from exc
+ except RuntimeError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail=str(exc),
+ ) from exc
- # Create a fresh in-memory DB from the script to validate it
- try:
- mem_conn = sqlite3.connect(":memory:")
- mem_conn.executescript(sql_script)
- mem_conn.close()
- except sqlite3.Error as exc:
- tmp_path.unlink(missing_ok=True)
- raise HTTPException(
- status_code=status.HTTP_400_BAD_REQUEST,
- detail=f"Backup file contains invalid SQL: {exc}",
- ) from exc
+ elif backend == "postgresql":
+ db.close()
+ try:
+ _restore_postgresql(app_settings.database_url, tmp_path)
+ except FileNotFoundError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail=f"psql binary not found – is PostgreSQL client installed? ({exc})",
+ ) from exc
+ except RuntimeError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail=f"PostgreSQL restore failed: {exc}",
+ ) from exc
- # Close the application DB session before replacing the file
- db.close()
+ elif backend == "mysql":
+ db.close()
+ try:
+ _restore_mysql(app_settings.database_url, tmp_path)
+ except FileNotFoundError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail=f"mysql binary not found – is MySQL client installed? ({exc})",
+ ) from exc
+ except RuntimeError as exc:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail=f"MySQL restore failed: {exc}",
+ ) from exc
- # Preserve the current DB before overwriting
- import shutil
+ else:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail=f"Database backend '{backend}' does not support restore.",
+ )
- bak = str(db_path) + ".pre_restore"
- try:
- shutil.copy2(str(db_path), bak)
- except OSError as exc:
- logger.warning(f"Could not create pre-restore backup at {bak}: {exc}")
-
- try:
- # Write the restored database
- restore_conn = sqlite3.connect(str(db_path))
- restore_conn.executescript(sql_script)
- restore_conn.close()
- except sqlite3.Error as exc:
- # Attempt rollback
- try:
- if os.path.exists(bak):
- shutil.copy2(bak, str(db_path))
- except OSError as rollback_exc:
- logger.error(f"Rollback failed; database may be corrupted: {rollback_exc}")
- raise HTTPException(
- status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
- detail=f"Restore failed: {exc}",
- ) from exc
finally:
tmp_path.unlink(missing_ok=True)
diff --git a/app/tasks/backup_tasks.py b/app/tasks/backup_tasks.py
index c94ca1eb..f40014f7 100644
--- a/app/tasks/backup_tasks.py
+++ b/app/tasks/backup_tasks.py
@@ -16,12 +16,19 @@ Three separate Celery-beat entries call ``create_backup`` with the appropriate
After each backup is created ``_apply_retention`` prunes old local backups for
that tier. Remote copies are pruned by ``_prune_remote_backups`` which mirrors
the same retention limits.
+
+Supported database backends
+----------------------------
+- **SQLite** – dumped via Python's built-in ``sqlite3.iterdump()``; archive extension ``.db.gz``
+- **PostgreSQL** – dumped via ``pg_dump --format=plain``; archive extension ``.pgsql.gz``
+- **MySQL / MariaDB** – dumped via ``mysqldump --single-transaction``; archive extension ``.mysql.gz``
"""
import gzip
import hashlib
import logging
import os
+import subprocess
from datetime import datetime, timezone
from pathlib import Path
@@ -42,6 +49,13 @@ _BACKUP_TYPE_RETAIN: dict[str, str] = {
"weekly": "backup_retain_weekly",
}
+#: Map of backend name → archive file extension.
+_BACKEND_EXTENSIONS: dict[str, str] = {
+ "sqlite": ".db.gz",
+ "postgresql": ".pgsql.gz",
+ "mysql": ".mysql.gz",
+}
+
def _backup_dir() -> Path:
"""Return (and create) the local backup directory."""
@@ -51,6 +65,14 @@ def _backup_dir() -> Path:
return path
+def _db_backend() -> str:
+ """Return the database backend name (e.g. ``'sqlite'``, ``'postgresql'``, ``'mysql'``)."""
+ from sqlalchemy.engine.url import make_url
+
+ url = make_url(settings.database_url)
+ return url.get_backend_name()
+
+
def _db_path() -> Path | None:
"""Return the SQLite database file path, or None for non-SQLite databases."""
from sqlalchemy.engine.url import make_url
@@ -64,6 +86,20 @@ def _db_path() -> Path | None:
return Path(db)
+def _archive_ext_for_backend(backend: str) -> str:
+ """Return the archive file extension for the given database backend.
+
+ Args:
+ backend: Backend name as returned by
+ ``sqlalchemy.engine.url.URL.get_backend_name()`` (e.g. ``'sqlite'``).
+
+ Returns:
+ File extension string including the leading dot, e.g. ``'.db.gz'``.
+ Falls back to ``'.sql.gz'`` for unknown backends.
+ """
+ return _BACKEND_EXTENSIONS.get(backend, ".sql.gz")
+
+
def _sha256(path: Path) -> str:
"""Return the SHA-256 hex digest of *path*."""
h = hashlib.sha256()
@@ -86,6 +122,277 @@ def _dump_sqlite(db_path: Path, dest: Path) -> None:
conn.close()
+def _dump_postgresql(db_url: str, dest: Path) -> None:
+ """Write a gzip-compressed ``pg_dump`` of the PostgreSQL database to *dest*.
+
+ Uses ``PGPASSWORD`` environment variable so the password is never exposed on
+ the process command line.
+
+ Args:
+ db_url: Full SQLAlchemy database URL (e.g. ``postgresql://user:pass@host/db``).
+ dest: Destination path for the ``.pgsql.gz`` archive.
+
+ Raises:
+ RuntimeError: If ``pg_dump`` exits with a non-zero return code.
+ FileNotFoundError: If the ``pg_dump`` binary is not found.
+ """
+ from sqlalchemy.engine.url import make_url
+
+ url = make_url(db_url)
+ env = os.environ.copy()
+ if url.password:
+ env["PGPASSWORD"] = str(url.password)
+
+ # Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL),
+ # not from user-controlled input. shell=False (the default when passing a list) is used
+ # so there is no shell interpretation of the argument values.
+ cmd: list[str] = ["pg_dump", "--format=plain", "--no-password"]
+ if url.host:
+ cmd.extend(["-h", url.host])
+ if url.port:
+ cmd.extend(["-p", str(url.port)])
+ if url.username:
+ cmd.extend(["-U", url.username])
+ if url.database:
+ cmd.append(url.database)
+
+ with gzip.open(str(dest), "wb") as gz:
+ proc = subprocess.Popen( # noqa: S603
+ cmd,
+ stdout=subprocess.PIPE,
+ stderr=subprocess.PIPE,
+ env=env,
+ )
+ stdout = proc.stdout
+ if stdout is None: # pragma: no cover – guaranteed by stdout=PIPE
+ raise RuntimeError("pg_dump produced no stdout pipe")
+ try:
+ while True:
+ chunk = stdout.read(65536)
+ if not chunk:
+ break
+ gz.write(chunk)
+ finally:
+ stdout.close()
+ stderr_bytes = proc.stderr.read() if proc.stderr else b""
+ proc.wait()
+
+ if proc.returncode != 0:
+ dest.unlink(missing_ok=True)
+ raise RuntimeError(
+ f"pg_dump exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}"
+ )
+
+
+def _dump_mysql(db_url: str, dest: Path) -> None:
+ """Write a gzip-compressed ``mysqldump`` of the MySQL database to *dest*.
+
+ Uses the ``MYSQL_PWD`` environment variable so the password is never exposed
+ on the process command line.
+
+ Args:
+ db_url: Full SQLAlchemy database URL
+ (e.g. ``mysql+pymysql://user:pass@host/db``).
+ dest: Destination path for the ``.mysql.gz`` archive.
+
+ Raises:
+ RuntimeError: If ``mysqldump`` exits with a non-zero return code.
+ FileNotFoundError: If the ``mysqldump`` binary is not found.
+ """
+ from sqlalchemy.engine.url import make_url
+
+ url = make_url(db_url)
+ env = os.environ.copy()
+ if url.password:
+ env["MYSQL_PWD"] = str(url.password)
+
+ # Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
+ # shell=False (list form) prevents shell interpretation of argument values.
+ cmd: list[str] = ["mysqldump", "--single-transaction", "--routines", "--triggers"]
+ if url.host:
+ cmd.extend(["-h", url.host])
+ if url.port:
+ cmd.extend(["-P", str(url.port)])
+ if url.username:
+ cmd.extend(["-u", url.username])
+ if url.database:
+ cmd.append(url.database)
+
+ with gzip.open(str(dest), "wb") as gz:
+ proc = subprocess.Popen( # noqa: S603
+ cmd,
+ stdout=subprocess.PIPE,
+ stderr=subprocess.PIPE,
+ env=env,
+ )
+ stdout = proc.stdout
+ if stdout is None: # pragma: no cover – guaranteed by stdout=PIPE
+ raise RuntimeError("mysqldump produced no stdout pipe")
+ try:
+ while True:
+ chunk = stdout.read(65536)
+ if not chunk:
+ break
+ gz.write(chunk)
+ finally:
+ stdout.close()
+ stderr_bytes = proc.stderr.read() if proc.stderr else b""
+ proc.wait()
+
+ if proc.returncode != 0:
+ dest.unlink(missing_ok=True)
+ raise RuntimeError(
+ f"mysqldump exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}"
+ )
+
+
+def _restore_sqlite(db_path: Path, archive_path: Path) -> None:
+ """Restore a SQLite database from a gzip-compressed SQL dump archive.
+
+ Validates the SQL by replaying it on an in-memory database before touching
+ the live file. Saves a ``.pre_restore`` rollback copy first.
+
+ Args:
+ db_path: Path to the live SQLite database file to overwrite.
+ archive_path: Path to the ``.db.gz`` gzip-compressed SQL dump.
+
+ Raises:
+ ValueError: If the archive cannot be decompressed or contains invalid SQL.
+ RuntimeError: If writing the restored database fails.
+ """
+ import shutil
+ import sqlite3
+
+ # Decompress and read SQL statements
+ try:
+ with gzip.open(str(archive_path), "rt", encoding="utf-8") as gz:
+ sql_script = gz.read()
+ except Exception as exc:
+ raise ValueError(f"Failed to decompress backup file: {exc}") from exc
+
+ # Validate by replaying on an in-memory database
+ try:
+ mem_conn = sqlite3.connect(":memory:")
+ mem_conn.executescript(sql_script)
+ mem_conn.close()
+ except sqlite3.Error as exc:
+ raise ValueError(f"Backup file contains invalid SQL: {exc}") from exc
+
+ # Preserve the current DB before overwriting
+ bak = str(db_path) + ".pre_restore"
+ try:
+ shutil.copy2(str(db_path), bak)
+ except OSError as exc:
+ logger.warning(f"Could not create pre-restore backup at {bak}: {exc}")
+
+ try:
+ restore_conn = sqlite3.connect(str(db_path))
+ restore_conn.executescript(sql_script)
+ restore_conn.close()
+ except sqlite3.Error as exc:
+ # Attempt rollback to the pre-restore copy
+ try:
+ if os.path.exists(bak):
+ shutil.copy2(bak, str(db_path))
+ except OSError as rollback_exc:
+ logger.error(f"Rollback failed; database may be corrupted: {rollback_exc}")
+ raise RuntimeError(f"SQLite restore failed: {exc}") from exc
+
+
+def _restore_postgresql(db_url: str, archive_path: Path) -> None:
+ """Restore a PostgreSQL database from a gzip-compressed SQL dump archive.
+
+ Pipes the decompressed dump to ``psql``. Uses ``PGPASSWORD`` so the
+ password is never exposed on the process command line.
+
+ Args:
+ db_url: Full SQLAlchemy database URL.
+ archive_path: Path to the ``.pgsql.gz`` gzip-compressed ``pg_dump`` archive.
+
+ Raises:
+ RuntimeError: If ``psql`` exits with a non-zero return code.
+ FileNotFoundError: If the ``psql`` binary is not found.
+ """
+ from sqlalchemy.engine.url import make_url
+
+ url = make_url(db_url)
+ env = os.environ.copy()
+ if url.password:
+ env["PGPASSWORD"] = str(url.password)
+
+ # Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
+ # shell=False (list form) prevents shell interpretation of argument values.
+ cmd: list[str] = ["psql", "--no-password"]
+ if url.host:
+ cmd.extend(["-h", url.host])
+ if url.port:
+ cmd.extend(["-p", str(url.port)])
+ if url.username:
+ cmd.extend(["-U", url.username])
+ if url.database:
+ cmd.append(url.database)
+
+ with gzip.open(str(archive_path), "rb") as gz:
+ proc = subprocess.Popen( # noqa: S603
+ cmd,
+ stdin=subprocess.PIPE,
+ stdout=subprocess.PIPE,
+ stderr=subprocess.PIPE,
+ env=env,
+ )
+ _, stderr_bytes = proc.communicate(input=gz.read())
+
+ if proc.returncode != 0:
+ raise RuntimeError(f"psql exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}")
+
+
+def _restore_mysql(db_url: str, archive_path: Path) -> None:
+ """Restore a MySQL database from a gzip-compressed SQL dump archive.
+
+ Pipes the decompressed dump to ``mysql``. Uses the ``MYSQL_PWD``
+ environment variable so the password is never exposed on the command line.
+
+ Args:
+ db_url: Full SQLAlchemy database URL.
+ archive_path: Path to the ``.mysql.gz`` gzip-compressed ``mysqldump`` archive.
+
+ Raises:
+ RuntimeError: If ``mysql`` exits with a non-zero return code.
+ FileNotFoundError: If the ``mysql`` binary is not found.
+ """
+ from sqlalchemy.engine.url import make_url
+
+ url = make_url(db_url)
+ env = os.environ.copy()
+ if url.password:
+ env["MYSQL_PWD"] = str(url.password)
+
+ # Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
+ # shell=False (list form) prevents shell interpretation of argument values.
+ cmd: list[str] = ["mysql"]
+ if url.host:
+ cmd.extend(["-h", url.host])
+ if url.port:
+ cmd.extend(["-P", str(url.port)])
+ if url.username:
+ cmd.extend(["-u", url.username])
+ if url.database:
+ cmd.append(url.database)
+
+ with gzip.open(str(archive_path), "rb") as gz:
+ proc = subprocess.Popen( # noqa: S603
+ cmd,
+ stdin=subprocess.PIPE,
+ stdout=subprocess.PIPE,
+ stderr=subprocess.PIPE,
+ env=env,
+ )
+ _, stderr_bytes = proc.communicate(input=gz.read())
+
+ if proc.returncode != 0:
+ raise RuntimeError(f"mysql exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}")
+
+
def _apply_retention(backup_type: str, db: object) -> None:
"""Delete local backups beyond the retention limit for *backup_type*.
@@ -314,6 +621,11 @@ def _email_backup(archive_path: Path, filename: str) -> None:
def create_backup(self, backup_type: str = "hourly") -> dict:
"""Create a database backup archive and apply retention.
+ Supports SQLite (``.db.gz``), PostgreSQL (``.pgsql.gz``), and
+ MySQL / MariaDB (``.mysql.gz``) databases. The native dump tool for the
+ configured backend (``sqlite3``, ``pg_dump``, or ``mysqldump``) must be
+ available on the worker's ``PATH``.
+
Args:
backup_type: ``"hourly"``, ``"daily"``, or ``"weekly"``.
@@ -327,19 +639,27 @@ def create_backup(self, backup_type: str = "hourly") -> dict:
logger.debug("Backup is disabled; skipping create_backup task.")
return {"status": "disabled"}
+ backend = _db_backend()
+ ext = _archive_ext_for_backend(backend)
+
ts = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H-%M-%S")
- filename = f"backup_{backup_type}_{ts}.db.gz"
+ filename = f"backup_{backup_type}_{ts}{ext}"
archive_path = _backup_dir() / filename
- db_path = _db_path()
- if db_path is None:
- logger.warning("Backup task skipped: non-SQLite databases are not supported for file-based backups.")
+ # SQLite: verify the database file exists before attempting to dump it
+ db_path: Path | None = None
+ if backend == "sqlite":
+ db_path = _db_path()
+ if db_path is None:
+ logger.warning("Backup task skipped: in-memory SQLite databases are not supported.")
+ return {"status": "unsupported_db"}
+ if not db_path.exists():
+ logger.error(f"Database file not found: {db_path}")
+ return {"status": "error", "detail": f"DB file missing: {db_path}"}
+ elif backend not in ("postgresql", "mysql"):
+ logger.warning(f"Backup task skipped: unsupported database backend '{backend}'.")
return {"status": "unsupported_db"}
- if not db_path.exists():
- logger.error(f"Database file not found: {db_path}")
- return {"status": "error", "detail": f"DB file missing: {db_path}"}
-
status = "ok"
checksum: str | None = None
size_bytes = 0
@@ -347,7 +667,15 @@ def create_backup(self, backup_type: str = "hourly") -> dict:
remote_path: str | None = None
try:
- _dump_sqlite(db_path, archive_path)
+ if backend == "sqlite":
+ # db_path is guaranteed non-None: we returned early if it were None
+ if db_path is None: # pragma: no cover
+ return {"status": "error", "detail": "db_path unexpectedly None"}
+ _dump_sqlite(db_path, archive_path)
+ elif backend == "postgresql":
+ _dump_postgresql(settings.database_url, archive_path)
+ elif backend == "mysql":
+ _dump_mysql(settings.database_url, archive_path)
size_bytes = archive_path.stat().st_size
checksum = _sha256(archive_path)
logger.info(f"Created {backup_type} backup: {archive_path} ({size_bytes:,} bytes)")
diff --git a/docs/ConfigurationGuide.md b/docs/ConfigurationGuide.md
index 2de94abd..3ea63e1e 100644
--- a/docs/ConfigurationGuide.md
+++ b/docs/ConfigurationGuide.md
@@ -1037,9 +1037,13 @@ Webhook URLs, secrets, and subscribed events are configured per-webhook via the
### Backup & Restore
-DocuElevate can automatically back up the SQLite database on a scheduled basis.
+DocuElevate automatically backs up the database on a scheduled basis.
Backups are managed from the **Admin → Backup & Restore** dashboard.
+Supported database backends: **SQLite** (`.db.gz`), **PostgreSQL** (`.pgsql.gz`), **MySQL / MariaDB** (`.mysql.gz`).
+For PostgreSQL and MySQL backups the respective CLI client (`pg_dump` / `psql` or `mysqldump` / `mysql`) must be installed on the Celery worker host.
+See the [Database Configuration Guide](DatabaseConfiguration.md#backup-procedures) for setup details.
+
| **Variable** | **Description** | **Default** |
|--------------------------------|-----------------------------------------------------------------------------------------------|---------------------|
| `BACKUP_ENABLED` | Enable or disable automatic scheduled backups (`True`/`False`). | `True` |
diff --git a/docs/DatabaseConfiguration.md b/docs/DatabaseConfiguration.md
index 0e5811aa..9eb1fece 100644
--- a/docs/DatabaseConfiguration.md
+++ b/docs/DatabaseConfiguration.md
@@ -341,29 +341,91 @@ Disable `prepared_statements` when using PgBouncer in transaction mode.
## Backup Procedures
-### PostgreSQL
+DocuElevate's built-in **Backup & Restore** feature (Admin → Backup & Restore) supports all three
+database backends natively, using the native dump tools of each database.
-**Manual backup:**
+| Backend | Backup tool | Archive extension | Restore tool |
+|----------------|--------------|-------------------|--------------|
+| SQLite | `sqlite3` (built-in Python) | `.db.gz` | `sqlite3` (built-in Python) |
+| PostgreSQL | `pg_dump` | `.pgsql.gz` | `psql` |
+| MySQL/MariaDB | `mysqldump` | `.mysql.gz` | `mysql` |
+
+Passwords are passed via the `PGPASSWORD` (PostgreSQL) and `MYSQL_PWD` (MySQL) environment
+variables so they are never exposed on the process command line.
+
+### Prerequisites
+
+For PostgreSQL and MySQL backups the corresponding CLI client must be installed on the
+worker host (the container / server that runs Celery workers):
+
+```bash
+# PostgreSQL clients (Debian/Ubuntu)
+apt-get install -y postgresql-client
+
+# MySQL clients (Debian/Ubuntu)
+apt-get install -y default-mysql-client
+```
+
+The binaries required are:
+
+- **PostgreSQL**: `pg_dump` (backup) and `psql` (restore)
+- **MySQL / MariaDB**: `mysqldump` (backup) and `mysql` (restore)
+
+### Using the Admin Dashboard
+
+Navigate to **Admin → Backup & Restore** to:
+
+- Trigger manual backups (hourly / daily / weekly)
+- Download backup archives
+- Upload and restore a backup archive
+- Configure retention and remote storage destinations
+
+### PostgreSQL – manual backup/restore
+
+**Manual backup using DocuElevate's archive format (for use with the UI restore):**
+
+```bash
+pg_dump --format=plain --no-password \
+ -h localhost -U docuelevate docuelevate \
+ | gzip > docuelevate_$(date +%Y%m%d_%H%M).pgsql.gz
+```
+
+**Restore via the DocuElevate UI:** upload the `.pgsql.gz` file on the Backup & Restore page.
+
+**Manual restore using native tools (custom format):**
```bash
pg_dump -h localhost -U docuelevate -F c docuelevate > docuelevate_$(date +%Y%m%d_%H%M).dump
-```
-
-**Restore:**
-
-```bash
pg_restore -h localhost -U docuelevate -d docuelevate docuelevate_20240101_1200.dump
```
**Automated daily backup (cron example):**
```cron
-0 2 * * * pg_dump -h localhost -U docuelevate -F c docuelevate | gzip > /backups/docuelevate_$(date +\%Y\%m\%d).dump.gz
+0 2 * * * pg_dump --format=plain -h localhost -U docuelevate docuelevate | gzip > /backups/docuelevate_$(date +\%Y\%m\%d).pgsql.gz
```
Use your cloud provider's automated backup feature when available (e.g., RDS automated snapshots, Cloud SQL backups).
-### SQLite
+### MySQL / MariaDB – manual backup/restore
+
+**Manual backup using DocuElevate's archive format (for use with the UI restore):**
+
+```bash
+MYSQL_PWD=yourpassword mysqldump --single-transaction --routines --triggers \
+ -h localhost -u docuelevate docuelevate \
+ | gzip > docuelevate_$(date +%Y%m%d_%H%M).mysql.gz
+```
+
+**Restore via the DocuElevate UI:** upload the `.mysql.gz` file on the Backup & Restore page.
+
+**Manual restore using native tools:**
+
+```bash
+gunzip -c docuelevate_20240101_1200.mysql.gz | mysql -h localhost -u docuelevate -p docuelevate
+```
+
+### SQLite – manual backup/restore
```bash
# Stop the application first, or use SQLite's online backup API
diff --git a/tests/test_backup.py b/tests/test_backup.py
index d4859493..34dd115b 100644
--- a/tests/test_backup.py
+++ b/tests/test_backup.py
@@ -236,6 +236,254 @@ class TestBackupTaskHelpers:
result = _db_path()
assert result is None
+ def test_db_backend_sqlite(self):
+ """_db_backend() returns 'sqlite' for SQLite URLs."""
+ from app.tasks.backup_tasks import _db_backend
+
+ with patch("app.tasks.backup_tasks.settings") as mock_settings:
+ mock_settings.database_url = "sqlite:////tmp/test.db"
+ assert _db_backend() == "sqlite"
+
+ def test_db_backend_postgresql(self):
+ """_db_backend() returns 'postgresql' for PostgreSQL URLs."""
+ from app.tasks.backup_tasks import _db_backend
+
+ with patch("app.tasks.backup_tasks.settings") as mock_settings:
+ mock_settings.database_url = "postgresql://user:pass@localhost/db"
+ assert _db_backend() == "postgresql"
+
+ def test_db_backend_mysql(self):
+ """_db_backend() returns 'mysql' for MySQL URLs."""
+ from app.tasks.backup_tasks import _db_backend
+
+ with patch("app.tasks.backup_tasks.settings") as mock_settings:
+ mock_settings.database_url = "mysql+pymysql://user:pass@localhost/db"
+ assert _db_backend() == "mysql"
+
+ def test_archive_ext_sqlite(self):
+ """_archive_ext_for_backend() returns '.db.gz' for sqlite."""
+ from app.tasks.backup_tasks import _archive_ext_for_backend
+
+ assert _archive_ext_for_backend("sqlite") == ".db.gz"
+
+ def test_archive_ext_postgresql(self):
+ """_archive_ext_for_backend() returns '.pgsql.gz' for postgresql."""
+ from app.tasks.backup_tasks import _archive_ext_for_backend
+
+ assert _archive_ext_for_backend("postgresql") == ".pgsql.gz"
+
+ def test_archive_ext_mysql(self):
+ """_archive_ext_for_backend() returns '.mysql.gz' for mysql."""
+ from app.tasks.backup_tasks import _archive_ext_for_backend
+
+ assert _archive_ext_for_backend("mysql") == ".mysql.gz"
+
+ def test_archive_ext_unknown(self):
+ """_archive_ext_for_backend() falls back to '.sql.gz' for unknown backends."""
+ from app.tasks.backup_tasks import _archive_ext_for_backend
+
+ assert _archive_ext_for_backend("mssql") == ".sql.gz"
+
+ def test_dump_postgresql_success(self, tmp_path):
+ """_dump_postgresql() streams pg_dump output into a gzip archive."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _dump_postgresql
+
+ dest = tmp_path / "dump.pgsql.gz"
+ fake_sql = b"-- PostgreSQL database dump\nSELECT 1;\n"
+
+ mock_proc = MagicMock()
+ mock_proc.stdout.read.side_effect = [fake_sql, b""]
+ mock_proc.stderr.read.return_value = b""
+ mock_proc.returncode = 0
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ _dump_postgresql("postgresql://user:pass@localhost/testdb", dest)
+
+ assert dest.exists()
+ with gzip.open(str(dest), "rb") as gz:
+ assert gz.read() == fake_sql
+
+ def test_dump_postgresql_failure(self, tmp_path):
+ """_dump_postgresql() raises RuntimeError when pg_dump fails."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _dump_postgresql
+
+ dest = tmp_path / "dump.pgsql.gz"
+
+ mock_proc = MagicMock()
+ mock_proc.stdout.read.side_effect = [b""]
+ mock_proc.stderr.read.return_value = b"FATAL: connection refused"
+ mock_proc.returncode = 1
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ with pytest.raises(RuntimeError, match="pg_dump exited with code 1"):
+ _dump_postgresql("postgresql://user:pass@localhost/testdb", dest)
+
+ def test_dump_mysql_success(self, tmp_path):
+ """_dump_mysql() streams mysqldump output into a gzip archive."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _dump_mysql
+
+ dest = tmp_path / "dump.mysql.gz"
+ fake_sql = b"-- MySQL dump\nCREATE TABLE t (id INT);\n"
+
+ mock_proc = MagicMock()
+ mock_proc.stdout.read.side_effect = [fake_sql, b""]
+ mock_proc.stderr.read.return_value = b""
+ mock_proc.returncode = 0
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ _dump_mysql("mysql+pymysql://user:pass@localhost/testdb", dest)
+
+ assert dest.exists()
+ with gzip.open(str(dest), "rb") as gz:
+ assert gz.read() == fake_sql
+
+ def test_dump_mysql_failure(self, tmp_path):
+ """_dump_mysql() raises RuntimeError when mysqldump fails."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _dump_mysql
+
+ dest = tmp_path / "dump.mysql.gz"
+
+ mock_proc = MagicMock()
+ mock_proc.stdout.read.side_effect = [b""]
+ mock_proc.stderr.read.return_value = b"ERROR: Access denied"
+ mock_proc.returncode = 1
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ with pytest.raises(RuntimeError, match="mysqldump exited with code 1"):
+ _dump_mysql("mysql+pymysql://user:pass@localhost/testdb", dest)
+
+ def test_restore_sqlite_success(self, tmp_path):
+ """_restore_sqlite() applies a valid SQL dump to a SQLite file."""
+ from app.tasks.backup_tasks import _restore_sqlite
+
+ db_file = tmp_path / "test.db"
+ conn = sqlite3.connect(str(db_file))
+ conn.execute("CREATE TABLE old (id INTEGER)")
+ conn.commit()
+ conn.close()
+
+ # Create a valid dump archive
+ sql = "BEGIN TRANSACTION;\nCREATE TABLE new_tbl (x TEXT);\nCOMMIT;\n"
+ archive = tmp_path / "dump.db.gz"
+ with gzip.open(str(archive), "wt") as gz:
+ gz.write(sql)
+
+ _restore_sqlite(db_file, archive)
+
+ conn2 = sqlite3.connect(str(db_file))
+ tables = [r[0] for r in conn2.execute("SELECT name FROM sqlite_master WHERE type='table'")]
+ conn2.close()
+ assert "new_tbl" in tables
+
+ def test_restore_sqlite_invalid_gz(self, tmp_path):
+ """_restore_sqlite() raises ValueError for corrupt gzip content."""
+ from app.tasks.backup_tasks import _restore_sqlite
+
+ db_file = tmp_path / "test.db"
+ db_file.write_bytes(b"")
+ archive = tmp_path / "bad.db.gz"
+ archive.write_bytes(b"not gzip data")
+
+ with pytest.raises(ValueError, match="Failed to decompress"):
+ _restore_sqlite(db_file, archive)
+
+ def test_restore_sqlite_invalid_sql(self, tmp_path):
+ """_restore_sqlite() raises ValueError for invalid SQL content."""
+ from app.tasks.backup_tasks import _restore_sqlite
+
+ db_file = tmp_path / "test.db"
+ db_file.write_bytes(b"")
+ archive = tmp_path / "bad.db.gz"
+ with gzip.open(str(archive), "wt") as gz:
+ gz.write("THIS IS NOT VALID SQL!!!;\n")
+
+ with pytest.raises(ValueError, match="invalid SQL"):
+ _restore_sqlite(db_file, archive)
+
+ def test_restore_postgresql_success(self, tmp_path):
+ """_restore_postgresql() pipes the archive to psql."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _restore_postgresql
+
+ fake_sql = b"-- PostgreSQL dump\nSELECT 1;\n"
+ archive = tmp_path / "dump.pgsql.gz"
+ with gzip.open(str(archive), "wb") as gz:
+ gz.write(fake_sql)
+
+ mock_proc = MagicMock()
+ mock_proc.communicate.return_value = (b"", b"")
+ mock_proc.returncode = 0
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ _restore_postgresql("postgresql://user:pass@localhost/testdb", archive)
+
+ mock_proc.communicate.assert_called_once_with(input=fake_sql)
+
+ def test_restore_postgresql_failure(self, tmp_path):
+ """_restore_postgresql() raises RuntimeError when psql fails."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _restore_postgresql
+
+ archive = tmp_path / "dump.pgsql.gz"
+ with gzip.open(str(archive), "wb") as gz:
+ gz.write(b"SELECT 1;")
+
+ mock_proc = MagicMock()
+ mock_proc.communicate.return_value = (b"", b"ERROR: invalid input")
+ mock_proc.returncode = 1
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ with pytest.raises(RuntimeError, match="psql exited with code 1"):
+ _restore_postgresql("postgresql://user:pass@localhost/testdb", archive)
+
+ def test_restore_mysql_success(self, tmp_path):
+ """_restore_mysql() pipes the archive to mysql."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _restore_mysql
+
+ fake_sql = b"-- MySQL dump\nSELECT 1;\n"
+ archive = tmp_path / "dump.mysql.gz"
+ with gzip.open(str(archive), "wb") as gz:
+ gz.write(fake_sql)
+
+ mock_proc = MagicMock()
+ mock_proc.communicate.return_value = (b"", b"")
+ mock_proc.returncode = 0
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ _restore_mysql("mysql+pymysql://user:pass@localhost/testdb", archive)
+
+ mock_proc.communicate.assert_called_once_with(input=fake_sql)
+
+ def test_restore_mysql_failure(self, tmp_path):
+ """_restore_mysql() raises RuntimeError when mysql fails."""
+ from unittest.mock import MagicMock
+
+ from app.tasks.backup_tasks import _restore_mysql
+
+ archive = tmp_path / "dump.mysql.gz"
+ with gzip.open(str(archive), "wb") as gz:
+ gz.write(b"SELECT 1;")
+
+ mock_proc = MagicMock()
+ mock_proc.communicate.return_value = (b"", b"ERROR: Access denied")
+ mock_proc.returncode = 1
+
+ with patch("app.tasks.backup_tasks.subprocess.Popen", return_value=mock_proc):
+ with pytest.raises(RuntimeError, match="mysql exited with code 1"):
+ _restore_mysql("mysql+pymysql://user:pass@localhost/testdb", archive)
+
def test_apply_retention_prunes_old(self, tmp_path, db_session):
"""_apply_retention() deletes backups beyond the retention limit."""
from app.tasks.backup_tasks import _apply_retention
@@ -352,20 +600,28 @@ class TestCreateBackupTask:
result = create_backup("hourly")
assert result["status"] == "disabled"
- def test_non_sqlite_db(self):
- """create_backup returns unsupported_db for non-SQLite databases."""
+ def test_unsupported_db_backend(self):
+ """create_backup returns unsupported_db for backends other than sqlite/postgresql/mysql."""
from app.tasks.backup_tasks import create_backup
- with (
- patch("app.tasks.backup_tasks.settings") as mock_settings,
- patch("app.tasks.backup_tasks._db_path", return_value=None),
- ):
+ with patch("app.tasks.backup_tasks.settings") as mock_settings:
mock_settings.backup_enabled = True
+ mock_settings.database_url = "mssql+pyodbc://user:pass@server/db"
+ result = create_backup("hourly")
+ assert result["status"] == "unsupported_db"
+
+ def test_in_memory_sqlite_unsupported(self):
+ """create_backup returns unsupported_db for in-memory SQLite."""
+ from app.tasks.backup_tasks import create_backup
+
+ with patch("app.tasks.backup_tasks.settings") as mock_settings:
+ mock_settings.backup_enabled = True
+ mock_settings.database_url = "sqlite:///:memory:"
result = create_backup("hourly")
assert result["status"] == "unsupported_db"
def test_missing_db_file(self, tmp_path):
- """create_backup returns error when the DB file does not exist."""
+ """create_backup returns error when the SQLite DB file does not exist."""
from app.tasks.backup_tasks import create_backup
missing = tmp_path / "does_not_exist.db"
@@ -375,6 +631,7 @@ class TestCreateBackupTask:
patch("app.tasks.backup_tasks._db_path", return_value=missing),
):
mock_settings.backup_enabled = True
+ mock_settings.database_url = f"sqlite:///{missing}"
result = create_backup("hourly")
assert result["status"] == "error"
@@ -400,6 +657,7 @@ class TestCreateBackupTask:
):
backup_dir.mkdir(parents=True, exist_ok=True)
mock_settings.backup_enabled = True
+ mock_settings.database_url = f"sqlite:///{db_file}"
mock_db = MagicMock()
mock_sl.return_value.__enter__ = MagicMock(return_value=mock_db)
mock_sl.return_value.__exit__ = MagicMock(return_value=False)
@@ -408,6 +666,69 @@ class TestCreateBackupTask:
assert result["status"] == "ok"
assert "filename" in result
assert result["filename"].startswith("backup_hourly_")
+ assert result["filename"].endswith(".db.gz")
+
+ def test_successful_backup_postgresql(self, tmp_path):
+ """create_backup creates a .pgsql.gz archive for PostgreSQL databases."""
+ from app.tasks.backup_tasks import create_backup
+
+ backup_dir = tmp_path / "backups"
+ backup_dir.mkdir()
+
+ def fake_pg_dump(db_url: str, dest: Path) -> None:
+ with gzip.open(str(dest), "wb") as gz:
+ gz.write(b"-- PostgreSQL dump\n")
+
+ with (
+ patch("app.tasks.backup_tasks.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._backup_dir", return_value=backup_dir),
+ patch("app.tasks.backup_tasks._dump_postgresql", side_effect=fake_pg_dump),
+ patch("app.tasks.backup_tasks._upload_remote", return_value=None),
+ patch("app.tasks.backup_tasks._apply_retention"),
+ patch("app.tasks.backup_tasks._prune_remote_backups"),
+ patch("app.tasks.backup_tasks.SessionLocal") as mock_sl,
+ ):
+ mock_settings.backup_enabled = True
+ mock_settings.database_url = "postgresql://user:pass@localhost/testdb"
+ mock_db = MagicMock()
+ mock_sl.return_value.__enter__ = MagicMock(return_value=mock_db)
+ mock_sl.return_value.__exit__ = MagicMock(return_value=False)
+ result = create_backup("daily")
+
+ assert result["status"] == "ok"
+ assert result["filename"].endswith(".pgsql.gz")
+ assert "daily" in result["filename"]
+
+ def test_successful_backup_mysql(self, tmp_path):
+ """create_backup creates a .mysql.gz archive for MySQL databases."""
+ from app.tasks.backup_tasks import create_backup
+
+ backup_dir = tmp_path / "backups"
+ backup_dir.mkdir()
+
+ def fake_mysql_dump(db_url: str, dest: Path) -> None:
+ with gzip.open(str(dest), "wb") as gz:
+ gz.write(b"-- MySQL dump\n")
+
+ with (
+ patch("app.tasks.backup_tasks.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._backup_dir", return_value=backup_dir),
+ patch("app.tasks.backup_tasks._dump_mysql", side_effect=fake_mysql_dump),
+ patch("app.tasks.backup_tasks._upload_remote", return_value=None),
+ patch("app.tasks.backup_tasks._apply_retention"),
+ patch("app.tasks.backup_tasks._prune_remote_backups"),
+ patch("app.tasks.backup_tasks.SessionLocal") as mock_sl,
+ ):
+ mock_settings.backup_enabled = True
+ mock_settings.database_url = "mysql+pymysql://user:pass@localhost/testdb"
+ mock_db = MagicMock()
+ mock_sl.return_value.__enter__ = MagicMock(return_value=mock_db)
+ mock_sl.return_value.__exit__ = MagicMock(return_value=False)
+ result = create_backup("weekly")
+
+ assert result["status"] == "ok"
+ assert result["filename"].endswith(".mysql.gz")
+ assert "weekly" in result["filename"]
def test_invalid_backup_type_defaults_to_hourly(self, tmp_path):
"""create_backup normalises unknown backup_type to 'hourly'."""
@@ -431,6 +752,7 @@ class TestCreateBackupTask:
patch("app.tasks.backup_tasks.SessionLocal") as mock_sl,
):
mock_settings.backup_enabled = True
+ mock_settings.database_url = f"sqlite:///{db_file}"
mock_db = MagicMock()
mock_sl.return_value.__enter__ = MagicMock(return_value=mock_db)
mock_sl.return_value.__exit__ = MagicMock(return_value=False)
@@ -549,19 +871,24 @@ class TestBackupAPIEndpoints:
assert resp.status_code == 403
def test_restore_wrong_extension(self, admin_client):
- """POST /api/admin/backup/restore rejects non-.db.gz files."""
+ """POST /api/admin/backup/restore rejects files with wrong extension for current backend."""
+ # Default test env uses sqlite:///:memory: → expects .db.gz
resp = admin_client.post(
"/api/admin/backup/restore",
files={"file": ("backup.zip", b"data", "application/zip")},
)
assert resp.status_code == 400
- def test_restore_invalid_gz_content(self, admin_client):
+ def test_restore_invalid_gz_content(self, admin_client, tmp_path):
"""POST /api/admin/backup/restore rejects corrupt gzip data."""
- resp = admin_client.post(
- "/api/admin/backup/restore",
- files={"file": ("backup.db.gz", b"not gzip data at all", "application/gzip")},
- )
+ db_file = tmp_path / "test.db"
+ db_file.write_bytes(b"")
+
+ with patch("app.tasks.backup_tasks._db_path", return_value=db_file):
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.db.gz", b"not gzip data at all", "application/gzip")},
+ )
assert resp.status_code == 400
def test_restore_valid_archive(self, admin_client, tmp_path):
@@ -581,11 +908,11 @@ class TestBackupAPIEndpoints:
assert resp.status_code == 200
assert resp.json()["status"] == "restored"
- def test_restore_non_sqlite_db(self, admin_client):
- """POST /api/admin/backup/restore returns 400 for non-SQLite database."""
- sql = "BEGIN TRANSACTION;\nCOMMIT;\n"
- gz_data = gzip.compress(sql.encode())
+ def test_restore_in_memory_sqlite(self, admin_client):
+ """POST /api/admin/backup/restore returns 400 for in-memory SQLite (no file to restore to)."""
+ gz_data = gzip.compress(b"BEGIN TRANSACTION;\nCOMMIT;\n")
+ # _db_path() returns None for :memory: URLs → 400
with patch("app.tasks.backup_tasks._db_path", return_value=None):
resp = admin_client.post(
"/api/admin/backup/restore",
@@ -593,6 +920,106 @@ class TestBackupAPIEndpoints:
)
assert resp.status_code == 400
+ def test_restore_wrong_extension_for_postgresql(self, admin_client):
+ """POST /api/admin/backup/restore returns 400 when uploading .db.gz for PostgreSQL backend."""
+ gz_data = gzip.compress(b"-- PostgreSQL dump")
+
+ with patch("app.config.settings") as mock_settings:
+ mock_settings.database_url = "postgresql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.db.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 400
+
+ def test_restore_wrong_extension_for_mysql(self, admin_client):
+ """POST /api/admin/backup/restore returns 400 when uploading .db.gz for MySQL backend."""
+ gz_data = gzip.compress(b"-- MySQL dump")
+
+ with patch("app.config.settings") as mock_settings:
+ mock_settings.database_url = "mysql+pymysql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.db.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 400
+
+ def test_restore_postgresql_success(self, admin_client):
+ """POST /api/admin/backup/restore succeeds for PostgreSQL database."""
+ gz_data = gzip.compress(b"-- PostgreSQL dump\n")
+
+ with (
+ patch("app.config.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._restore_postgresql") as mock_restore,
+ ):
+ mock_settings.database_url = "postgresql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.pgsql.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 200
+ assert resp.json()["status"] == "restored"
+ mock_restore.assert_called_once()
+
+ def test_restore_mysql_success(self, admin_client):
+ """POST /api/admin/backup/restore succeeds for MySQL database."""
+ gz_data = gzip.compress(b"-- MySQL dump\n")
+
+ with (
+ patch("app.config.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._restore_mysql") as mock_restore,
+ ):
+ mock_settings.database_url = "mysql+pymysql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.mysql.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 200
+ assert resp.json()["status"] == "restored"
+ mock_restore.assert_called_once()
+
+ def test_restore_postgresql_runtime_error(self, admin_client):
+ """POST /api/admin/backup/restore returns 500 when psql command fails."""
+ gz_data = gzip.compress(b"-- PostgreSQL dump\n")
+
+ with (
+ patch("app.config.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._restore_postgresql", side_effect=RuntimeError("psql failed")),
+ ):
+ mock_settings.database_url = "postgresql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.pgsql.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 500
+
+ def test_restore_postgresql_missing_binary(self, admin_client):
+ """POST /api/admin/backup/restore returns 500 when psql binary is missing."""
+ gz_data = gzip.compress(b"-- PostgreSQL dump\n")
+
+ with (
+ patch("app.config.settings") as mock_settings,
+ patch("app.tasks.backup_tasks._restore_postgresql", side_effect=FileNotFoundError("psql not found")),
+ ):
+ mock_settings.database_url = "postgresql://user:pass@localhost/testdb"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.pgsql.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 500
+
+ def test_restore_unsupported_backend(self, admin_client):
+ """POST /api/admin/backup/restore returns 400 for an unsupported database backend."""
+ gz_data = gzip.compress(b"-- some dump\n")
+
+ with patch("app.config.settings") as mock_settings:
+ mock_settings.database_url = "mssql+pyodbc://user:pass@server/db"
+ resp = admin_client.post(
+ "/api/admin/backup/restore",
+ files={"file": ("backup.sql.gz", gz_data, "application/gzip")},
+ )
+ assert resp.status_code == 400
+
# ---------------------------------------------------------------------------
# View tests
From d36ba88de765b688888c6e661256f8508da86d89 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Sun, 8 Mar 2026 10:06:08 +0000
Subject: [PATCH 09/25] feat(auth): password reset, forgot username, and admin
user management for local accounts
- Add /forgot-password and /forgot-username page routes and templates
- Update login page label to "Username or Email" (both already accepted by backend)
- Add "Forgot password?" and "Forgot username?" links to login page
- Add POST /api/auth/forgot-username endpoint + send_forgot_username_email() utility
- Add admin endpoints: PATCH /local/{id}, POST /local/{id}/send-password-reset, POST /local/{id}/set-password
- Update admin_users.html with Edit, Password, and Reset action buttons + modals
- Add 23 tests; fix code review issues (import style, display_name clearing behaviour)
- Update docs/API.md and docs/UserGuide.md
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
---
app/api/admin_users.py | 18 +-
app/api/local_auth.py | 36 ++++
app/utils/local_auth.py | 35 ++++
docs/API.md | 125 +++++++++++
docs/UserGuide.md | 25 ++-
frontend/templates/admin_users.html | 2 +-
frontend/templates/forgot_username.html | 125 +++++++++++
frontend/templates/login.html | 11 +-
tests/test_admin_users.py | 266 +++++++++++++++++++++++-
tests/test_local_auth.py | 116 +++++++++++
10 files changed, 742 insertions(+), 17 deletions(-)
create mode 100644 frontend/templates/forgot_username.html
diff --git a/app/api/admin_users.py b/app/api/admin_users.py
index 1e5fc221..d73a1467 100644
--- a/app/api/admin_users.py
+++ b/app/api/admin_users.py
@@ -15,6 +15,7 @@ from pydantic import BaseModel, Field
from sqlalchemy import func
from sqlalchemy.orm import Session
+from app.config import settings
from app.database import get_db
from app.models import FileRecord, LocalUser, UserProfile
from app.utils.local_auth import generate_token, hash_password, send_password_reset_email
@@ -393,7 +394,8 @@ def update_local_user(local_user_id: int, body: LocalUserUpdate, db: DbSession,
user.email = body.email
if body.display_name is not None:
- user.display_name = body.display_name
+ # Normalise empty string to None so that clearing the field removes the display name
+ user.display_name = body.display_name or None
if body.is_admin is not None:
user.is_admin = body.is_admin
@@ -431,9 +433,7 @@ def update_local_user(local_user_id: int, body: LocalUserUpdate, db: DbSession,
status_code=status.HTTP_200_OK,
summary="Send a password reset email to a local user",
)
-def admin_send_password_reset(
- local_user_id: int, request: Request, db: DbSession, _admin: AdminUser
-) -> dict[str, Any]:
+def admin_send_password_reset(local_user_id: int, request: Request, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
"""Generate a password reset token and email the reset link to the local user.
This is a last-resort tool for admins to help users who are locked out.
@@ -443,13 +443,11 @@ def admin_send_password_reset(
Raises:
404: Local user not found.
"""
- from app.config import settings as _settings
-
user = db.query(LocalUser).filter(LocalUser.id == local_user_id).first()
if not user:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Local user not found.")
- if not _settings.email_host:
+ if not settings.email_host:
logger.warning("Admin requested password reset for %s but SMTP is not configured", user.email)
return {"sent": False, "reason": "SMTP is not configured on this server."}
@@ -500,10 +498,10 @@ def admin_set_password(
db.rollback()
raise
- logger.info(
- "[SECURITY] ADMIN_SET_PASSWORD user=%s admin=%s", user.email, _admin.get("email", "unknown")
- )
+ logger.info("[SECURITY] ADMIN_SET_PASSWORD user=%s admin=%s", user.email, _admin.get("email", "unknown"))
return {"updated": True, "email": user.email}
+
+
def get_user(user_id: str, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
"""Return profile and document statistics for a specific user."""
doc_count = db.query(func.count(FileRecord.id)).filter(FileRecord.owner_id == user_id).scalar() or 0
diff --git a/app/api/local_auth.py b/app/api/local_auth.py
index f42cbe2b..be6df664 100644
--- a/app/api/local_auth.py
+++ b/app/api/local_auth.py
@@ -31,6 +31,7 @@ from app.utils.local_auth import (
generate_token,
hash_password,
is_token_expired,
+ send_forgot_username_email,
send_password_reset_email,
send_verification_email,
)
@@ -79,6 +80,12 @@ class PasswordResetBody(BaseModel):
new_password_confirm: str
+class ForgotUsernameBody(BaseModel):
+ """Body for the forgot-username endpoint."""
+
+ email: str
+
+
# ---------------------------------------------------------------------------
# Page routes (return HTML)
# ---------------------------------------------------------------------------
@@ -107,6 +114,19 @@ async def verify_email_sent_page(request: Request) -> Any:
return templates.TemplateResponse("verify_email_sent.html", {"request": request})
+@router.get("/forgot-username", include_in_schema=False)
+async def forgot_username_page(request: Request) -> Any:
+ """Render the forgot-username page where users can request a username reminder email."""
+ return templates.TemplateResponse(
+ "forgot_username.html",
+ {
+ "request": request,
+ "csrf_token": getattr(request.state, "csrf_token", ""),
+ "app_version": settings.version,
+ },
+ )
+
+
@router.get("/forgot-password", include_in_schema=False)
async def forgot_password_page(request: Request) -> Any:
"""Render the forgot-password page where users can request a reset email."""
@@ -350,3 +370,19 @@ async def reset_password(body: PasswordResetBody, db: DbSession) -> dict[str, st
logger.info("[SECURITY] PASSWORD_RESET_SUCCESS user=%s", user.email)
return {"message": "Password updated successfully."}
+
+
+@router.post("/api/auth/forgot-username")
+async def forgot_username(body: ForgotUsernameBody, db: DbSession) -> dict[str, str]:
+ """Send a username reminder email.
+
+ Always returns 200 to avoid leaking whether an email is registered.
+ """
+ user = db.query(LocalUser).filter(LocalUser.email == body.email).first()
+ if user:
+ try:
+ send_forgot_username_email(user.email, user.username)
+ except Exception as exc:
+ logger.warning("Failed to send forgot-username email to %s: %s", user.email, exc)
+
+ return {"message": "Username reminder sent if account exists."}
diff --git a/app/utils/local_auth.py b/app/utils/local_auth.py
index 669d6819..20a54719 100644
--- a/app/utils/local_auth.py
+++ b/app/utils/local_auth.py
@@ -164,6 +164,41 @@ def send_password_reset_email(email: str, username: str, token: str, base_url: s
_smtp_send(subject, html_body, plain_body, email)
+def send_forgot_username_email(email: str, username: str) -> None:
+ """Send an email reminding the user of their username.
+
+ Args:
+ email: Recipient email address.
+ username: The user's username to include in the message.
+ """
+ subject = "Your DocuElevate username"
+ html_body = f"""
+
+
+
+
+
Your Username
+
You requested a reminder of your DocuElevate username.
+
+
Your username is:
+
{username}
+
+
You can sign in using your username or your email address.
+
If you did not request this reminder, you can safely ignore this email.
+
+
DocuElevate · Intelligent Document Processing
+
+
+"""
+ plain_body = (
+ f"You requested a reminder of your DocuElevate username.\n\n"
+ f"Your username is: {username}\n\n"
+ "You can sign in using your username or your email address.\n\n"
+ "If you did not request this, please ignore this email."
+ )
+ _smtp_send(subject, html_body, plain_body, email)
+
+
def build_session_user(user: object) -> dict:
"""Build the session user dict for a LocalUser, matching the OAuth session format.
diff --git a/docs/API.md b/docs/API.md
index e3bca54a..809bc0a6 100644
--- a/docs/API.md
+++ b/docs/API.md
@@ -844,6 +844,131 @@ problem.
---
+**GET** `/api/admin/users/local`
+
+List all local (email/password) user accounts with basic metadata.
+
+---
+
+**POST** `/api/admin/users/local`
+
+Create a new local user account (admin-only, immediately active — no email verification required).
+
+**Request body**:
+```json
+{
+ "email": "user@example.com",
+ "username": "alice",
+ "display_name": "Alice Smith",
+ "password": "securepassword",
+ "is_admin": false
+}
+```
+
+---
+
+**PATCH** `/api/admin/users/local/{local_user_id}`
+
+Update an existing local user account. Only the provided (non-null) fields are modified.
+If the email is changed, the associated `UserProfile.user_id` is also updated automatically.
+
+**Request body** (all fields optional):
+```json
+{
+ "email": "newemail@example.com",
+ "display_name": "Alice Wonderland",
+ "is_admin": true,
+ "is_active": false
+}
+```
+
+**Error Responses**:
+- `404`: Local user not found
+- `409`: New email already taken by another account
+
+---
+
+**POST** `/api/admin/users/local/{local_user_id}/send-password-reset`
+
+Send a password reset email to a local user on their behalf. Useful when a user is locked out.
+Returns `{"sent": true}` on success or `{"sent": false, "reason": "..."}` when SMTP is not
+configured or sending fails (never returns an error status so the admin always gets feedback).
+
+**Error Responses**:
+- `404`: Local user not found
+
+---
+
+**POST** `/api/admin/users/local/{local_user_id}/set-password`
+
+Directly set a new password for a local user without requiring an email token (last resort when
+email delivery is unavailable). The user should be advised to change their password after logging in.
+
+**Request body**:
+```json
+{
+ "password": "temporarypassword"
+}
+```
+
+**Error Responses**:
+- `404`: Local user not found
+- `422`: Password shorter than 8 characters
+
+---
+
+**DELETE** `/api/admin/users/local/{local_user_id}`
+
+Delete a local user account by numeric ID. The associated `UserProfile` is also removed. Documents
+owned by this user are **not** deleted. Returns `204 No Content` on success.
+
+---
+
+### Local Authentication (self-service)
+
+These endpoints are for local (email/password) users and do not require authentication.
+
+**POST** `/api/auth/request-password-reset`
+
+Send a password reset email. Always returns 200 to avoid leaking whether an email is registered.
+
+**Request body**:
+```json
+{ "email": "user@example.com" }
+```
+
+---
+
+**POST** `/api/auth/reset-password`
+
+Set a new password using a valid reset token (received via email).
+
+**Request body**:
+```json
+{
+ "token": "the-token-from-email",
+ "new_password": "newpassword",
+ "new_password_confirm": "newpassword"
+}
+```
+
+**Error Responses**:
+- `400`: Token is invalid or expired
+- `422`: Passwords do not match
+
+---
+
+**POST** `/api/auth/forgot-username`
+
+Send a username reminder email. Always returns 200 to avoid leaking whether an email is registered.
+
+**Request body**:
+```json
+{ "email": "user@example.com" }
+```
+
+---
+
### Settings Suggestions (Autocomplete)
**GET** `/api/settings/{key}/suggestions`
diff --git a/docs/UserGuide.md b/docs/UserGuide.md
index 59ffdc40..a6332bc9 100644
--- a/docs/UserGuide.md
+++ b/docs/UserGuide.md
@@ -28,8 +28,29 @@ If OpenID Connect authentication is configured:
3. Log in with your existing credentials on that platform
4. You'll be redirected back to DocuElevate after successful authentication
-#### User Sessions
-- Once authenticated, your session will remain active until you log out or it expires
+#### Local User Accounts
+If your administrator has created a local (email/password) account for you:
+
+1. You'll see a "Sign in with username" form on the login page
+2. Enter your **username or email address** — both are accepted
+3. Enter your password and click **Sign in**
+
+##### Forgot your password?
+If you can't remember your password:
+1. Click **Forgot password?** below the sign-in form
+2. Enter your email address and click **Send reset link**
+3. Check your inbox for a password reset email (valid for 24 hours)
+4. Click the link in the email and enter your new password
+
+##### Forgot your username?
+If you can't remember your username:
+1. Click **Forgot username?** below the sign-in form
+2. Enter your email address and click **Send username reminder**
+3. You'll receive an email with your username
+
+> **Tip:** You can always sign in with your email address directly — you don't need to look up your username.
+
+
- Click the "Logout" button in the top navigation bar to end your session
- For security, sessions automatically expire after a period of inactivity
diff --git a/frontend/templates/admin_users.html b/frontend/templates/admin_users.html
index 874c1789..360dcb2c 100644
--- a/frontend/templates/admin_users.html
+++ b/frontend/templates/admin_users.html
@@ -1122,7 +1122,7 @@ function adminUsersApp() {
},
body: JSON.stringify({
email: this.editLocalUserModal.form.email || null,
- display_name: this.editLocalUserModal.form.display_name || null,
+ display_name: this.editLocalUserModal.form.display_name,
is_admin: this.editLocalUserModal.form.is_admin,
is_active: this.editLocalUserModal.form.is_active,
}),
diff --git a/frontend/templates/forgot_username.html b/frontend/templates/forgot_username.html
new file mode 100644
index 00000000..c8c9b8d2
--- /dev/null
+++ b/frontend/templates/forgot_username.html
@@ -0,0 +1,125 @@
+
+
+
+
+
+ DocuElevate - Forgot Username
+
+
+
+
+
+
+
+
+
+
+
Forgot your username?
+
+ Enter the email address associated with your account and we'll send you your username.
+ You can also sign in directly with your email address.
+
+
+
+
+
+
+
+
+
+
Check your inbox
+
+ If an account exists for that email address, your username has been sent.
+ Remember: you can also sign in using your email address directly.
+