fix(merge): resolve conflicts with main v0.92.0 keeping path-param regression tests
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
+149
-2
@@ -7,7 +7,7 @@ user accounts directly, without requiring email verification.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
@@ -15,9 +15,10 @@ 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 hash_password
|
||||
from app.utils.local_auth import generate_token, hash_password, send_password_reset_email
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/users", tags=["admin-users"])
|
||||
@@ -123,6 +124,21 @@ class LocalUserCreate(BaseModel):
|
||||
is_admin: bool = Field(default=False, description="Grant admin privileges")
|
||||
|
||||
|
||||
class LocalUserUpdate(BaseModel):
|
||||
"""Body for admin-updating a local (email/password) user account."""
|
||||
|
||||
email: str | None = Field(default=None, max_length=255, description="New email address")
|
||||
display_name: str | None = Field(default=None, max_length=255, description="New display name")
|
||||
is_admin: bool | None = Field(default=None, description="Grant or revoke admin privileges")
|
||||
is_active: bool | None = Field(default=None, description="Activate or deactivate the account")
|
||||
|
||||
|
||||
class LocalUserSetPassword(BaseModel):
|
||||
"""Body for admin setting a temporary password for a local user."""
|
||||
|
||||
password: str = Field(..., min_length=8, max_length=128, description="New temporary password")
|
||||
|
||||
|
||||
class LocalUserResponse(BaseModel):
|
||||
"""Summary of a local user account."""
|
||||
|
||||
@@ -355,6 +371,137 @@ def delete_local_user(local_user_id: int, db: DbSession, _admin: AdminUser) -> N
|
||||
logger.info("Admin deleted local user account: %s", user.email)
|
||||
|
||||
|
||||
@router.patch("/local/{local_user_id}", summary="Update a local user account")
|
||||
def update_local_user(local_user_id: int, body: LocalUserUpdate, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update the email address, display name, admin flag, or active status of a local user account.
|
||||
|
||||
Only fields explicitly provided (non-None) are modified. If the email is changed
|
||||
the associated UserProfile row is also updated to keep ``user_id`` in sync.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
409: The new email is already taken by another account.
|
||||
"""
|
||||
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.")
|
||||
|
||||
old_email = user.email
|
||||
|
||||
if body.email is not None and body.email != user.email:
|
||||
if db.query(LocalUser).filter(LocalUser.email == body.email, LocalUser.id != local_user_id).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Email already registered.")
|
||||
user.email = body.email
|
||||
|
||||
if body.display_name is not None:
|
||||
# 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
|
||||
|
||||
if body.is_active is not None:
|
||||
user.is_active = body.is_active
|
||||
|
||||
try:
|
||||
db.flush()
|
||||
# Keep UserProfile.user_id in sync when email changes
|
||||
if body.email is not None and body.email != old_email:
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == old_email).first()
|
||||
if profile:
|
||||
profile.user_id = body.email
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin updated local user %s (id=%d)", user.email, user.id)
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"username": user.username,
|
||||
"display_name": user.display_name,
|
||||
"is_active": user.is_active,
|
||||
"is_admin": user.is_admin,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/local/{local_user_id}/send-password-reset",
|
||||
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]:
|
||||
"""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.
|
||||
Returns ``{"sent": true}`` on success and ``{"sent": false, "reason": "..."}`` when
|
||||
SMTP is not configured or sending fails.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
"""
|
||||
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:
|
||||
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."}
|
||||
|
||||
token = generate_token()
|
||||
user.password_reset_token = token
|
||||
user.password_reset_sent_at = datetime.now(tz=timezone.utc)
|
||||
db.commit()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
try:
|
||||
send_password_reset_email(user.email, user.username, token, base_url)
|
||||
except Exception as exc:
|
||||
logger.warning("Admin-triggered password reset email failed for %s: %s", user.email, exc)
|
||||
return {"sent": False, "reason": str(exc)}
|
||||
|
||||
logger.info("[SECURITY] ADMIN_PASSWORD_RESET_EMAIL user=%s admin=%s", user.email, _admin.get("email", "unknown"))
|
||||
return {"sent": True, "email": user.email}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/local/{local_user_id}/set-password",
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Set a temporary password for a local user account",
|
||||
)
|
||||
def admin_set_password(
|
||||
local_user_id: int, body: LocalUserSetPassword, db: DbSession, _admin: AdminUser
|
||||
) -> dict[str, Any]:
|
||||
"""Directly set a new password for a local user without requiring an email token.
|
||||
|
||||
Use this as a last resort when email delivery is unavailable. The user
|
||||
should be advised to change their password after logging in.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
"""
|
||||
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.")
|
||||
|
||||
user.hashed_password = hash_password(body.password)
|
||||
# Clear any outstanding reset tokens
|
||||
user.password_reset_token = None
|
||||
user.password_reset_sent_at = None
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("[SECURITY] ADMIN_SET_PASSWORD user=%s admin=%s", user.email, _admin.get("email", "unknown"))
|
||||
return {"updated": True, "email": user.email}
|
||||
|
||||
|
||||
@router.get("/{user_id:path}", summary="Get details for a single user")
|
||||
def get_user(user_id: str, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Return profile and document statistics for a specific user."""
|
||||
|
||||
+84
-68
@@ -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)
|
||||
|
||||
|
||||
@@ -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,32 @@ 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."""
|
||||
return templates.TemplateResponse(
|
||||
"forgot_password.html",
|
||||
{
|
||||
"request": request,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
"app_version": settings.version,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/reset-password", include_in_schema=False)
|
||||
async def reset_password_page(request: Request) -> Any:
|
||||
"""Render the password reset form page."""
|
||||
@@ -337,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."}
|
||||
|
||||
Reference in New Issue
Block a user