fix: merge main branch and renumber migration 037→040
Resolve all merge conflicts between our automation feature branch and current main (v0.163.0, 920 commits ahead). Conflicts resolved: - app/api/__init__.py: add automation_router alongside main's new routers (classification_rules, qr_auth, sessions, system_reset) - app/config.py: add main's new settings (dropbox_use_global_credentials, factory_reset_on_startup, enable_factory_reset) - app/models.py: add main's new models (ClassificationRuleModel, UserSession, QRLoginChallenge, SharePoint integration type) - app/utils/settings_service.py: merge automation_hooks_enabled with main's new metadata entries - docs/API.md: merge automation API docs with main's classification rules docs - docs/ConfigurationGuide.md: add factory reset settings - tests/conftest.py: import both AutomationHook and new main models Migration renumbered: - 037_add_automation_hooks → 040_add_automation_hooks - down_revision: 039_add_classification_rules (was 036_add_document_translation_fields) - Chain: 036 → 037 → 038 → 039 → 040 (automation hooks) For all non-automation files with conflicts, main's version was taken since our branch did not modify those files (conflicts were from a stale prior merge). Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com> Agent-Logs-Url: https://github.com/christianlouis/DocuElevate/sessions/cb62f012-3b69-4415-835e-3857ce3e9f45
This commit is contained in:
@@ -13,6 +13,7 @@ from app.api.automation import router as automation_router
|
||||
from app.api.azure import router as azure_router
|
||||
from app.api.backup import router as backup_router
|
||||
from app.api.billing import router as billing_router
|
||||
from app.api.classification_rules import router as classification_rules_router
|
||||
from app.api.compliance import router as compliance_router
|
||||
from app.api.database import router as database_router
|
||||
from app.api.diagnostic import router as diagnostic_router
|
||||
@@ -34,16 +35,19 @@ from app.api.pipelines import router as pipelines_router
|
||||
from app.api.plans import router as plans_router
|
||||
from app.api.process import router as process_router
|
||||
from app.api.profile import router as profile_router
|
||||
from app.api.qr_auth import router as qr_auth_router
|
||||
from app.api.queue import router as queue_router
|
||||
from app.api.routing_rules import router as routing_rules_router
|
||||
from app.api.saved_searches import router as saved_searches_router
|
||||
from app.api.scheduled_jobs import router as scheduled_jobs_router
|
||||
from app.api.search import router as search_router
|
||||
from app.api.sessions import router as sessions_router
|
||||
from app.api.settings import router as settings_router
|
||||
from app.api.shared_links import public_router as shared_links_public_router
|
||||
from app.api.shared_links import router as shared_links_router
|
||||
from app.api.similarity import router as similarity_router
|
||||
from app.api.subscriptions import router as subscriptions_router
|
||||
from app.api.system_reset import router as system_reset_router
|
||||
from app.api.translation import router as translation_router
|
||||
from app.api.url_upload import router as url_upload_router
|
||||
|
||||
@@ -97,6 +101,10 @@ router.include_router(scheduled_jobs_router)
|
||||
router.include_router(audit_logs_router)
|
||||
router.include_router(i18n_router)
|
||||
router.include_router(mobile_router)
|
||||
router.include_router(sessions_router)
|
||||
router.include_router(qr_auth_router)
|
||||
router.include_router(compliance_router)
|
||||
router.include_router(system_reset_router)
|
||||
router.include_router(translation_router)
|
||||
router.include_router(classification_rules_router)
|
||||
router.include_router(automation_router)
|
||||
|
||||
+123
-26
@@ -13,7 +13,7 @@ plaintext is returned exactly once at creation time.
|
||||
import hashlib
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
@@ -42,6 +42,9 @@ TOKEN_HASH_ITERATIONS = 100_000
|
||||
#: PBKDF2 salt for API token hashing (not secret, but fixed for determinism).
|
||||
TOKEN_HASH_SALT = b"api-token-v1"
|
||||
|
||||
#: Name prefix used for tokens created by the mobile app flow.
|
||||
MOBILE_TOKEN_PREFIX = "Mobile App"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
@@ -91,6 +94,21 @@ def hash_token(token: str) -> str:
|
||||
return dk.hex()
|
||||
|
||||
|
||||
def _token_to_dict(t: ApiToken) -> dict[str, Any]:
|
||||
"""Convert an ``ApiToken`` ORM instance to a serialisable dict."""
|
||||
return {
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"token_prefix": t.token_prefix,
|
||||
"is_active": t.is_active,
|
||||
"last_used_at": t.last_used_at,
|
||||
"last_used_ip": t.last_used_ip,
|
||||
"created_at": t.created_at,
|
||||
"revoked_at": t.revoked_at,
|
||||
"expires_at": t.expires_at,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -100,6 +118,12 @@ class TokenCreate(BaseModel):
|
||||
"""Schema for creating a new API token."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255, description="Human-readable label for the token")
|
||||
expires_in_days: int | None = Field(
|
||||
default=None,
|
||||
ge=1,
|
||||
le=3650, # Maximum 10 years; keeps tokens from being effectively permanent while allowing long-lived CI/CD tokens.
|
||||
description="Optional lifetime in days. If omitted the token never expires.",
|
||||
)
|
||||
|
||||
|
||||
class TokenResponse(BaseModel):
|
||||
@@ -113,6 +137,7 @@ class TokenResponse(BaseModel):
|
||||
last_used_ip: str | None
|
||||
created_at: datetime | None
|
||||
revoked_at: datetime | None
|
||||
expires_at: datetime | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
@@ -143,11 +168,16 @@ async def create_token(
|
||||
token_hash_value = hash_token(plaintext)
|
||||
prefix = plaintext[:12] # "de_" prefix + 9 random chars = 12 chars total
|
||||
|
||||
expires_at = None
|
||||
if body.expires_in_days is not None:
|
||||
expires_at = datetime.now(timezone.utc) + timedelta(days=body.expires_in_days)
|
||||
|
||||
db_token = ApiToken(
|
||||
owner_id=owner_id,
|
||||
name=body.name,
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
try:
|
||||
db.add(db_token)
|
||||
@@ -168,6 +198,7 @@ async def create_token(
|
||||
"last_used_ip": db_token.last_used_ip,
|
||||
"created_at": db_token.created_at,
|
||||
"revoked_at": db_token.revoked_at,
|
||||
"expires_at": db_token.expires_at,
|
||||
"token": plaintext,
|
||||
}
|
||||
|
||||
@@ -177,48 +208,114 @@ async def list_tokens(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all API tokens for the authenticated user."""
|
||||
tokens = db.query(ApiToken).filter(ApiToken.owner_id == owner_id).order_by(ApiToken.created_at.desc()).all()
|
||||
return [
|
||||
{
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"token_prefix": t.token_prefix,
|
||||
"is_active": t.is_active,
|
||||
"last_used_at": t.last_used_at,
|
||||
"last_used_ip": t.last_used_ip,
|
||||
"created_at": t.created_at,
|
||||
"revoked_at": t.revoked_at,
|
||||
}
|
||||
for t in tokens
|
||||
]
|
||||
"""List non-mobile API tokens for the authenticated user.
|
||||
|
||||
Mobile tokens (whose names start with ``"Mobile App"``) are excluded
|
||||
from this list; they are managed on the dedicated Devices page via
|
||||
``GET /api/api-tokens/mobile``.
|
||||
"""
|
||||
tokens = (
|
||||
db.query(ApiToken)
|
||||
.filter(
|
||||
ApiToken.owner_id == owner_id,
|
||||
~ApiToken.name.startswith(MOBILE_TOKEN_PREFIX),
|
||||
)
|
||||
.order_by(ApiToken.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_token_to_dict(t) for t in tokens]
|
||||
|
||||
|
||||
@router.get("/mobile", response_model=list[TokenResponse])
|
||||
async def list_mobile_tokens(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List mobile API tokens for the authenticated user.
|
||||
|
||||
Returns tokens whose names start with ``"Mobile App"`` — these are
|
||||
created via the mobile SSO flow or QR code login.
|
||||
"""
|
||||
tokens = (
|
||||
db.query(ApiToken)
|
||||
.filter(
|
||||
ApiToken.owner_id == owner_id,
|
||||
ApiToken.name.startswith(MOBILE_TOKEN_PREFIX),
|
||||
)
|
||||
.order_by(ApiToken.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_token_to_dict(t) for t in tokens]
|
||||
|
||||
|
||||
@router.delete("/{token_id}", status_code=status.HTTP_200_OK)
|
||||
async def revoke_token(
|
||||
async def revoke_or_delete_token(
|
||||
token_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Revoke (soft-delete) an API token.
|
||||
"""Revoke or permanently delete an API token.
|
||||
|
||||
The token row is kept for audit purposes but marked inactive with a
|
||||
``revoked_at`` timestamp.
|
||||
* **Active token** – soft-revoked: the row is kept for audit purposes
|
||||
but marked inactive with a ``revoked_at`` timestamp.
|
||||
* **Already-revoked token** – hard-deleted: the row is permanently
|
||||
removed from the database.
|
||||
"""
|
||||
db_token = db.query(ApiToken).filter(ApiToken.id == token_id, ApiToken.owner_id == owner_id).first()
|
||||
if not db_token:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Token not found")
|
||||
|
||||
if not db_token.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already revoked")
|
||||
if db_token.is_active:
|
||||
# Soft-revoke the active token.
|
||||
try:
|
||||
db_token.is_active = False
|
||||
db_token.revoked_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
|
||||
return {"detail": "Token revoked"}
|
||||
|
||||
# Hard-delete an already-revoked token.
|
||||
try:
|
||||
db_token.is_active = False
|
||||
db_token.revoked_at = datetime.now(timezone.utc)
|
||||
db.delete(db_token)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("API token permanently deleted: id=%s owner=%s", token_id, owner_id)
|
||||
return {"detail": "Token deleted"}
|
||||
|
||||
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
|
||||
return {"detail": "Token revoked"}
|
||||
|
||||
@router.post("/{token_id}/reactivate", status_code=status.HTTP_200_OK, response_model=TokenResponse)
|
||||
async def reactivate_token(
|
||||
token_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Reactivate a previously revoked API token.
|
||||
|
||||
Clears the ``revoked_at`` timestamp and sets ``is_active`` back to
|
||||
``True``. The token can be used for authentication again immediately.
|
||||
If the token had an ``expires_at`` in the past the caller should
|
||||
consider re-creating a new token instead.
|
||||
"""
|
||||
db_token = db.query(ApiToken).filter(ApiToken.id == token_id, ApiToken.owner_id == owner_id).first()
|
||||
if not db_token:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Token not found")
|
||||
|
||||
if db_token.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already active")
|
||||
|
||||
try:
|
||||
db_token.is_active = True
|
||||
db_token.revoked_at = None
|
||||
db.commit()
|
||||
db.refresh(db_token)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("API token reactivated: id=%s owner=%s", token_id, owner_id)
|
||||
return _token_to_dict(db_token)
|
||||
|
||||
@@ -0,0 +1,325 @@
|
||||
"""Classification Rules API endpoints.
|
||||
|
||||
Provides CRUD operations for managing custom document classification rules.
|
||||
System-wide rules (``owner_id IS NULL``) can only be managed by admins.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import ClassificationRuleModel
|
||||
from app.utils.classification_rules import (
|
||||
BUILTIN_CATEGORIES,
|
||||
RULE_TYPE_CONTENT,
|
||||
RULE_TYPE_FILENAME,
|
||||
RULE_TYPE_METADATA,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/classification-rules", tags=["classification"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
_VALID_RULE_TYPES = {RULE_TYPE_FILENAME, RULE_TYPE_CONTENT, RULE_TYPE_METADATA}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_user_id(request: Request) -> str:
|
||||
"""Extract the user identifier from the request session."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if user and hasattr(user, "get"):
|
||||
return user.get("sub") or user.get("email") or "anonymous"
|
||||
return "anonymous"
|
||||
|
||||
|
||||
def _is_admin(request: Request) -> bool:
|
||||
"""Check whether the current user is an admin."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if user and hasattr(user, "get"):
|
||||
groups = user.get("groups", [])
|
||||
return "admin" in groups or "Admin" in groups
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RuleCreate(BaseModel):
|
||||
"""Schema for creating a classification rule."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
category: str = Field(..., min_length=1, max_length=100)
|
||||
rule_type: str = Field(..., description="One of: filename_pattern, content_keyword, metadata_match")
|
||||
pattern: str = Field(..., min_length=1, max_length=1000)
|
||||
priority: int = Field(default=0, ge=0, le=1000)
|
||||
case_sensitive: bool = False
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class RuleUpdate(BaseModel):
|
||||
"""Schema for updating a classification rule."""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
category: str | None = Field(default=None, min_length=1, max_length=100)
|
||||
rule_type: str | None = Field(default=None)
|
||||
pattern: str | None = Field(default=None, min_length=1, max_length=1000)
|
||||
priority: int | None = Field(default=None, ge=0, le=1000)
|
||||
case_sensitive: bool | None = None
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
class RuleResponse(BaseModel):
|
||||
"""Schema for a classification rule response."""
|
||||
|
||||
id: int
|
||||
owner_id: str | None
|
||||
name: str
|
||||
category: str
|
||||
rule_type: str
|
||||
pattern: str
|
||||
priority: int
|
||||
case_sensitive: bool
|
||||
enabled: bool
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/categories")
|
||||
@require_login
|
||||
async def list_categories(request: Request) -> dict[str, str]:
|
||||
"""Return all built-in classification categories.
|
||||
|
||||
Custom categories created via rules are not included here; they are
|
||||
discovered dynamically when rules are evaluated.
|
||||
"""
|
||||
return BUILTIN_CATEGORIES
|
||||
|
||||
|
||||
@router.get("/rule-types")
|
||||
@require_login
|
||||
async def list_rule_types(request: Request) -> list[dict[str, str]]:
|
||||
"""Return the supported rule types with descriptions."""
|
||||
return [
|
||||
{
|
||||
"type": RULE_TYPE_FILENAME,
|
||||
"label": "Filename Pattern",
|
||||
"description": "Regex pattern matched against the original filename.",
|
||||
},
|
||||
{
|
||||
"type": RULE_TYPE_CONTENT,
|
||||
"label": "Content Keyword",
|
||||
"description": "Pipe-separated keywords matched against the OCR text.",
|
||||
},
|
||||
{
|
||||
"type": RULE_TYPE_METADATA,
|
||||
"label": "Metadata Match",
|
||||
"description": "field=value pattern matched against existing AI metadata.",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@router.get("/")
|
||||
@require_login
|
||||
async def list_rules(request: Request, db: DbSession) -> list[dict[str, Any]]:
|
||||
"""List classification rules visible to the current user.
|
||||
|
||||
Returns both system rules (``owner_id IS NULL``) and the user's own rules.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rules = (
|
||||
db.query(ClassificationRuleModel)
|
||||
.filter((ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == user_id))
|
||||
.order_by(ClassificationRuleModel.priority.desc(), ClassificationRuleModel.id)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": r.id,
|
||||
"owner_id": r.owner_id,
|
||||
"name": r.name,
|
||||
"category": r.category,
|
||||
"rule_type": r.rule_type,
|
||||
"pattern": r.pattern,
|
||||
"priority": r.priority,
|
||||
"case_sensitive": r.case_sensitive,
|
||||
"enabled": r.enabled,
|
||||
}
|
||||
for r in rules
|
||||
]
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
async def create_rule(request: Request, body: RuleCreate, db: DbSession) -> dict[str, Any]:
|
||||
"""Create a new custom classification rule.
|
||||
|
||||
The rule is owned by the current user. Admins may create system-wide
|
||||
rules by setting ``owner_id`` to ``null`` (not yet exposed).
|
||||
"""
|
||||
if body.rule_type not in _VALID_RULE_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
|
||||
)
|
||||
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
# Check for duplicate name within the user's scope
|
||||
existing = (
|
||||
db.query(ClassificationRuleModel)
|
||||
.filter(ClassificationRuleModel.owner_id == user_id, ClassificationRuleModel.name == body.name)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"A rule named '{body.name}' already exists.",
|
||||
)
|
||||
|
||||
rule = ClassificationRuleModel(
|
||||
owner_id=user_id,
|
||||
name=body.name,
|
||||
category=body.category,
|
||||
rule_type=body.rule_type,
|
||||
pattern=body.pattern,
|
||||
priority=body.priority,
|
||||
case_sensitive=body.case_sensitive,
|
||||
enabled=body.enabled,
|
||||
)
|
||||
try:
|
||||
db.add(rule)
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule created: id=%s, user=%s", rule.id, user_id)
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/{rule_id}")
|
||||
@require_login
|
||||
async def get_rule(request: Request, rule_id: int, db: DbSession) -> dict[str, Any]:
|
||||
"""Get a single classification rule by ID."""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
# Users can see system rules and their own rules
|
||||
if rule.owner_id is not None and rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{rule_id}")
|
||||
@require_login
|
||||
async def update_rule(request: Request, rule_id: int, body: RuleUpdate, db: DbSession) -> dict[str, Any]:
|
||||
"""Update an existing classification rule.
|
||||
|
||||
Users can only update their own rules. Admins can update any rule.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
if rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this rule")
|
||||
|
||||
if body.rule_type is not None and body.rule_type not in _VALID_RULE_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
|
||||
)
|
||||
|
||||
update_data = body.model_dump(exclude_unset=True)
|
||||
for field_name, value in update_data.items():
|
||||
setattr(rule, field_name, value)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule updated: id=%s, user=%s", rule.id, user_id)
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{rule_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
async def delete_rule(request: Request, rule_id: int, db: DbSession) -> None:
|
||||
"""Delete a classification rule.
|
||||
|
||||
Users can only delete their own rules. Admins can delete any rule.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
if rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot delete this rule")
|
||||
|
||||
try:
|
||||
db.delete(rule)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule deleted: id=%s, user=%s", rule_id, user_id)
|
||||
@@ -21,6 +21,67 @@ _DEFAULT_REDIS_URL = "redis://localhost:6379/0"
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unauthenticated probe endpoints for Kubernetes liveness / readiness checks.
|
||||
# These intentionally skip authentication so that kubelet can reach them
|
||||
# without credentials. They live under /diagnostic/healthz/* so that the
|
||||
# existing authenticated /diagnostic/health endpoint is unaffected.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/diagnostic/healthz/live")
|
||||
async def liveness_probe() -> JSONResponse:
|
||||
"""Lightweight liveness probe for Kubernetes.
|
||||
|
||||
Returns **200 OK** as long as the process is running. Kubernetes uses
|
||||
this to decide whether to *restart* the container — it should therefore
|
||||
be as cheap as possible and **never** check external dependencies.
|
||||
|
||||
**Authentication:** None (designed for kubelet probes).
|
||||
"""
|
||||
return JSONResponse(content={"status": "ok"}, status_code=200)
|
||||
|
||||
|
||||
@router.get("/diagnostic/healthz/ready")
|
||||
async def readiness_probe() -> JSONResponse:
|
||||
"""Readiness probe for Kubernetes.
|
||||
|
||||
Verifies that the application can serve traffic by checking the database
|
||||
and Redis. Kubernetes uses this to decide whether to *route traffic* to
|
||||
the pod.
|
||||
|
||||
Returns **200 OK** when all critical subsystems are reachable, or
|
||||
**503 Service Unavailable** when the database is down.
|
||||
|
||||
**Authentication:** None (designed for kubelet probes).
|
||||
"""
|
||||
checks: dict[str, dict[str, str]] = {}
|
||||
db_ok = False
|
||||
|
||||
# ── Database check ─────────────────────────────────────────────────
|
||||
try:
|
||||
with engine.connect() as conn:
|
||||
conn.execute(text("SELECT 1"))
|
||||
checks["database"] = {"status": "ok"}
|
||||
db_ok = True
|
||||
except Exception as exc:
|
||||
logger.warning("Readiness probe: database check failed: %s", exc)
|
||||
checks["database"] = {"status": "error", "detail": str(exc)}
|
||||
|
||||
# ── Redis check ────────────────────────────────────────────────────
|
||||
try:
|
||||
redis_url = settings.redis_url or _DEFAULT_REDIS_URL
|
||||
r = redis_lib.from_url(redis_url, socket_connect_timeout=2, socket_timeout=2)
|
||||
r.ping()
|
||||
checks["redis"] = {"status": "ok"}
|
||||
except Exception as exc:
|
||||
logger.warning("Readiness probe: Redis check failed: %s", exc)
|
||||
checks["redis"] = {"status": "error", "detail": str(exc)}
|
||||
|
||||
http_status = 503 if not db_ok else 200
|
||||
overall = "ready" if db_ok else "not_ready"
|
||||
return JSONResponse(content={"status": overall, "checks": checks}, status_code=http_status)
|
||||
|
||||
|
||||
@router.get("/diagnostic/health")
|
||||
@require_login
|
||||
|
||||
@@ -5,6 +5,7 @@ Dropbox API endpoints
|
||||
import logging
|
||||
import os
|
||||
from typing import Annotated, Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
|
||||
@@ -23,6 +24,93 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _build_dropbox_redirect_uri(request: Request) -> str:
|
||||
"""Build the Dropbox OAuth callback redirect URI.
|
||||
|
||||
Uses ``PUBLIC_BASE_URL`` when configured (recommended for deployments behind
|
||||
a reverse proxy that doesn't forward ``X-Forwarded-Proto``). Falls back to
|
||||
deriving the URI from the incoming request's scheme and host headers.
|
||||
"""
|
||||
if settings.public_base_url:
|
||||
return settings.public_base_url.rstrip("/") + "/dropbox-callback"
|
||||
return f"{request.url.scheme}://{request.url.netloc}/dropbox-callback"
|
||||
|
||||
|
||||
@router.get("/dropbox/global-authorize-url")
|
||||
@require_login
|
||||
async def dropbox_global_authorize_url(request: Request):
|
||||
"""Return the Dropbox OAuth authorization URL using the global app credentials.
|
||||
|
||||
This endpoint is used when ``DROPBOX_ALLOW_GLOBAL_CREDENTIALS_FOR_INTEGRATIONS``
|
||||
is enabled so that users can authorize their personal Dropbox integration without
|
||||
needing to supply their own app key/secret. Only the public ``app_key`` is
|
||||
embedded in the URL; the ``app_secret`` is never sent to the browser.
|
||||
"""
|
||||
if not settings.dropbox_allow_global_credentials_for_integrations:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Global credentials for integrations are not enabled",
|
||||
)
|
||||
if not settings.dropbox_app_key or not settings.dropbox_app_secret:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Global Dropbox credentials are not configured",
|
||||
)
|
||||
redirect_uri = _build_dropbox_redirect_uri(request)
|
||||
authorize_url = (
|
||||
"https://www.dropbox.com/oauth2/authorize"
|
||||
f"?client_id={settings.dropbox_app_key}"
|
||||
"&response_type=code"
|
||||
"&token_access_type=offline"
|
||||
f"&redirect_uri={quote(redirect_uri, safe='')}"
|
||||
)
|
||||
return {"authorize_url": authorize_url}
|
||||
|
||||
|
||||
@router.post("/dropbox/exchange-token-global")
|
||||
@require_login
|
||||
async def exchange_dropbox_token_global(
|
||||
request: Request,
|
||||
code: Annotated[str, Form(...)],
|
||||
redirect_uri: Annotated[str, Form(...)],
|
||||
):
|
||||
"""Exchange an authorization code using the global Dropbox app credentials.
|
||||
|
||||
Used when ``DROPBOX_ALLOW_GLOBAL_CREDENTIALS_FOR_INTEGRATIONS`` is enabled so
|
||||
that the ``app_secret`` is never exposed to the browser. Only the OAuth code
|
||||
and redirect URI need to be supplied by the client.
|
||||
"""
|
||||
if not settings.dropbox_allow_global_credentials_for_integrations:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Global credentials for integrations are not enabled",
|
||||
)
|
||||
if not settings.dropbox_app_key or not settings.dropbox_app_secret:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Global Dropbox credentials are not configured",
|
||||
)
|
||||
|
||||
token_url = "https://api.dropboxapi.com/oauth2/token"
|
||||
payload = {
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
}
|
||||
|
||||
token_data = exchange_oauth_token(provider_name="Dropbox", token_url=token_url, payload=payload)
|
||||
|
||||
return {
|
||||
"refresh_token": token_data["refresh_token"],
|
||||
"access_token": token_data["access_token"],
|
||||
"expires_in": token_data.get("expires_in", 14400),
|
||||
# Return the public app_key so the callback can store it in the integration
|
||||
"app_key": settings.dropbox_app_key,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/dropbox/exchange-token")
|
||||
@require_login
|
||||
async def exchange_dropbox_token(
|
||||
@@ -210,6 +298,91 @@ async def test_dropbox_token(request: Request):
|
||||
return {"status": "error", "message": f"Connection error: {str(e)}"}
|
||||
|
||||
|
||||
@router.post("/dropbox/list-folders")
|
||||
@require_login
|
||||
async def list_dropbox_folders(
|
||||
request: Request,
|
||||
access_token: Annotated[str, Form(...)],
|
||||
path: Annotated[str, Form()] = "",
|
||||
):
|
||||
"""
|
||||
List folders in a Dropbox account for the directory selector.
|
||||
|
||||
Accepts an OAuth access token (short-lived) and a path to list.
|
||||
Returns a flat list of folder entries under the given path.
|
||||
"""
|
||||
try:
|
||||
# Normalize path: Dropbox API uses "" for root, otherwise "/path"
|
||||
folder_path = path.strip()
|
||||
if folder_path == "/":
|
||||
folder_path = ""
|
||||
elif folder_path and not folder_path.startswith("/"):
|
||||
folder_path = f"/{folder_path}"
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
payload = {
|
||||
"path": folder_path,
|
||||
"recursive": False,
|
||||
"include_deleted": False,
|
||||
"include_has_explicit_shared_members": False,
|
||||
"include_mounted_folders": True,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
"https://api.dropboxapi.com/2/files/list_folder",
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Access token is invalid or expired. Please re-authorize.",
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox list_folder failed: {response.status_code} {response.text}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"Failed to list Dropbox folders: {response.text}",
|
||||
)
|
||||
|
||||
data = response.json()
|
||||
folders = []
|
||||
for entry in data.get("entries", []):
|
||||
if entry.get(".tag") == "folder":
|
||||
folders.append(
|
||||
{
|
||||
"name": entry["name"],
|
||||
"path": entry["path_display"],
|
||||
"id": entry.get("id", ""),
|
||||
}
|
||||
)
|
||||
|
||||
# Sort folders alphabetically
|
||||
folders.sort(key=lambda f: f["name"].lower())
|
||||
|
||||
return {
|
||||
"folders": folders,
|
||||
"path": folder_path or "/",
|
||||
"has_more": data.get("has_more", False),
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.exception(f"Error listing Dropbox folders: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to list folders: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/dropbox/save-settings")
|
||||
@require_login
|
||||
async def save_dropbox_settings(
|
||||
|
||||
+64
-22
@@ -20,6 +20,7 @@ from sqlalchemy.orm import Session
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
from app.models import FileProcessingStep, FileRecord, ProcessingLog
|
||||
from app.tasks.convert_to_pdf import convert_to_pdf
|
||||
from app.tasks.process_document import process_document
|
||||
@@ -346,7 +347,9 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
|
||||
try:
|
||||
# Find all file records
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -385,7 +388,9 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find all file records
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -457,7 +462,9 @@ def bulk_reprocess_files_cloud_ocr(request: Request, file_ids: List[int], db: Db
|
||||
Useful for re-running OCR on files with poor text quality or missing OCR text.
|
||||
"""
|
||||
try:
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -537,7 +544,9 @@ def bulk_download_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
Files not found on disk are silently skipped.
|
||||
"""
|
||||
try:
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -619,7 +628,9 @@ def reprocess_single_file(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
||||
@@ -675,7 +686,9 @@ def reprocess_with_cloud_ocr(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
||||
@@ -938,7 +951,9 @@ def retry_subtask(
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
||||
@@ -1080,7 +1095,9 @@ def get_file_preview(
|
||||
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
||||
@@ -1160,7 +1177,9 @@ def download_file(
|
||||
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
||||
@@ -1249,7 +1268,12 @@ async def _save_upload_file_chunks(file: UploadFile, target_path: str, max_size:
|
||||
|
||||
|
||||
def _check_for_exact_duplicate(db: DbSession, target_path: str, safe_filename: str) -> dict | None:
|
||||
"""Check for an exact duplicate of the uploaded file and return a warning if found."""
|
||||
"""Check for an exact duplicate of the uploaded file.
|
||||
|
||||
Returns a dict with duplicate info when the file's SHA-256 hash matches an
|
||||
already-processed document, or ``None`` when no duplicate is found (or
|
||||
deduplication is disabled).
|
||||
"""
|
||||
if not settings.enable_deduplication:
|
||||
return None
|
||||
|
||||
@@ -1268,8 +1292,8 @@ def _check_for_exact_duplicate(db: DbSession, target_path: str, safe_filename: s
|
||||
"original_file_id": existing.id,
|
||||
"original_filename": existing.original_filename,
|
||||
"message": (
|
||||
"This file appears to be an exact duplicate of an already-processed document. "
|
||||
"It will still be queued but will be flagged as a duplicate."
|
||||
"This file is an exact duplicate of an already-processed document. "
|
||||
"It has not been queued for processing again."
|
||||
),
|
||||
}
|
||||
except Exception as e:
|
||||
@@ -1280,7 +1304,12 @@ def _check_for_exact_duplicate(db: DbSession, target_path: str, safe_filename: s
|
||||
|
||||
@router.post("/ui-upload")
|
||||
@require_login
|
||||
async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...)):
|
||||
async def ui_upload(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
file: UploadFile = File(...),
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
):
|
||||
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
|
||||
workdir = settings.workdir
|
||||
|
||||
@@ -1366,6 +1395,25 @@ async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...
|
||||
logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'")
|
||||
file_size = written_size
|
||||
|
||||
# ── Early duplicate rejection ──────────────────────────────────────────
|
||||
# Check for exact duplicates (same SHA-256 hash) BEFORE enqueuing a
|
||||
# processing task. When deduplication is enabled and the file already
|
||||
# exists, we skip processing entirely, clean up the temp file, and
|
||||
# return the existing file's information to the caller.
|
||||
exact_duplicate = _check_for_exact_duplicate(db, target_path, safe_filename)
|
||||
if exact_duplicate:
|
||||
# Remove the just-saved temp file — it's a duplicate.
|
||||
try:
|
||||
os.remove(target_path)
|
||||
except OSError:
|
||||
pass
|
||||
return {
|
||||
"status": "duplicate",
|
||||
"original_filename": safe_filename,
|
||||
"stored_filename": target_filename,
|
||||
"duplicate_of": exact_duplicate,
|
||||
}
|
||||
|
||||
# Determine if the file is a PDF or needs conversion
|
||||
mime_type, _ = mimetypes.guess_type(target_path)
|
||||
file_ext = os.path.splitext(target_path)[1].lower()
|
||||
@@ -1429,6 +1477,8 @@ async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...
|
||||
".tif",
|
||||
".webp",
|
||||
".svg",
|
||||
".heic",
|
||||
".heif",
|
||||
}:
|
||||
# If it's an image, convert to PDF first
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
@@ -1442,20 +1492,12 @@ async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...
|
||||
logger.warning(f"Unsupported MIME type {mime_type} for {target_path}, attempting conversion")
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
|
||||
# Check for exact duplicates (same SHA-256 hash) before returning.
|
||||
# This gives the caller an immediate warning without waiting for the pipeline.
|
||||
# Only performed when deduplication is enabled in settings.
|
||||
exact_duplicate_warning = _check_for_exact_duplicate(db, target_path, safe_filename)
|
||||
|
||||
response: dict = {
|
||||
return {
|
||||
"task_id": task.id,
|
||||
"status": "queued",
|
||||
"original_filename": safe_filename,
|
||||
"stored_filename": target_filename,
|
||||
}
|
||||
if exact_duplicate_warning:
|
||||
response["duplicate_warning"] = exact_duplicate_warning
|
||||
return response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -32,6 +32,21 @@ from app.utils.encryption import decrypt_value, encrypt_value
|
||||
from app.utils.subscription import get_tier, get_user_tier_id
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
# Optional Dropbox SDK — imported at module level so tests can patch it cleanly.
|
||||
try:
|
||||
import dropbox as dbx_lib
|
||||
from dropbox.exceptions import AuthError as _DropboxAuthError
|
||||
from dropbox.exceptions import BadInputError as _DropboxBadInputError
|
||||
except ImportError: # pragma: no cover
|
||||
dbx_lib = None # type: ignore[assignment]
|
||||
|
||||
class _DropboxAuthError(Exception): # type: ignore[no-redef]
|
||||
"""Stub — only used when the dropbox package is missing."""
|
||||
|
||||
class _DropboxBadInputError(Exception): # type: ignore[no-redef]
|
||||
"""Stub — only used when the dropbox package is missing."""
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/integrations", tags=["integrations"])
|
||||
|
||||
@@ -550,6 +565,47 @@ def _test_s3_connection(config: dict[str, Any] | None, credentials: dict[str, An
|
||||
return {"success": False, "message": "S3 connection failed"}
|
||||
|
||||
|
||||
def _test_dropbox_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test a Dropbox connection by verifying OAuth credentials via the Dropbox API."""
|
||||
if dbx_lib is None:
|
||||
return {"success": False, "message": "dropbox package is not installed"} # pragma: no cover
|
||||
|
||||
creds = credentials or {}
|
||||
app_key = creds.get("app_key", "")
|
||||
app_secret = creds.get("app_secret", "")
|
||||
refresh_token = creds.get("refresh_token", "")
|
||||
|
||||
if not refresh_token:
|
||||
return {"success": False, "message": "Missing required credential: refresh_token"}
|
||||
if not app_key or not app_secret:
|
||||
return {"success": False, "message": "Missing required credentials: app_key and app_secret"}
|
||||
|
||||
try:
|
||||
dbx = dbx_lib.Dropbox(
|
||||
app_key=app_key,
|
||||
app_secret=app_secret,
|
||||
oauth2_refresh_token=refresh_token,
|
||||
)
|
||||
account = dbx.users_get_current_account()
|
||||
display_name = getattr(account, "name", None)
|
||||
name_str = ""
|
||||
if display_name:
|
||||
name_str = f" ({getattr(display_name, 'display_name', '') or ''})"
|
||||
return {"success": True, "message": f"Dropbox connection successful{name_str}"}
|
||||
except _DropboxAuthError as exc:
|
||||
logger.warning("Dropbox auth error: %s", exc)
|
||||
return {
|
||||
"success": False,
|
||||
"message": "Dropbox authentication failed — check app_key, app_secret, and refresh_token",
|
||||
}
|
||||
except _DropboxBadInputError as exc:
|
||||
logger.warning("Dropbox bad input error: %s", exc)
|
||||
return {"success": False, "message": "Dropbox connection failed — invalid credentials format"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("Dropbox connection error: %s", exc)
|
||||
return {"success": False, "message": "Dropbox connection failed — check credentials and network connectivity"}
|
||||
|
||||
|
||||
def _test_webdav_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test a WebDAV/Nextcloud connection by issuing an HTTP PROPFIND."""
|
||||
import urllib.request
|
||||
@@ -596,6 +652,7 @@ def _test_webdav_connection(config: dict[str, Any] | None, credentials: dict[str
|
||||
|
||||
|
||||
_CONNECTION_TESTERS: dict[str, Any] = {
|
||||
IntegrationType.DROPBOX: _test_dropbox_connection,
|
||||
IntegrationType.IMAP: _test_imap_connection,
|
||||
IntegrationType.S3: _test_s3_connection,
|
||||
IntegrationType.WEBDAV: _test_webdav_connection,
|
||||
|
||||
+23
-8
@@ -120,6 +120,7 @@ class WhoAmIResponse(BaseModel):
|
||||
email: str | None
|
||||
avatar_url: str | None
|
||||
is_admin: bool
|
||||
preferred_language: str | None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -273,31 +274,44 @@ async def list_devices(
|
||||
return [_device_to_response(d) for d in devices]
|
||||
|
||||
|
||||
@router.delete("/devices/{device_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@router.delete("/devices/{device_id}", status_code=status.HTTP_200_OK)
|
||||
@require_login
|
||||
async def deactivate_device(
|
||||
request: Request,
|
||||
device_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> None:
|
||||
"""Deactivate a push-notification device registration.
|
||||
) -> dict[str, str]:
|
||||
"""Deactivate or permanently delete a push-notification device registration.
|
||||
|
||||
The device record is kept for audit purposes but will no longer receive
|
||||
push notifications.
|
||||
* **Active device** – soft-deactivated: the record is kept for audit
|
||||
purposes but will no longer receive push notifications.
|
||||
* **Already-inactive device** – hard-deleted: the record is permanently
|
||||
removed from the database.
|
||||
"""
|
||||
device = db.get(MobileDevice, device_id)
|
||||
if not device or device.owner_id != owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Device not found")
|
||||
|
||||
device.is_active = False
|
||||
if device.is_active:
|
||||
device.is_active = False
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Mobile device deactivated: id=%s owner=%s", device_id, owner_id)
|
||||
return {"detail": "Device deactivated"}
|
||||
|
||||
# Hard-delete an already-inactive device.
|
||||
try:
|
||||
db.delete(device)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Mobile device deactivated: id=%s owner=%s", device_id, owner_id)
|
||||
logger.info("Mobile device permanently deleted: id=%s owner=%s", device_id, owner_id)
|
||||
return {"detail": "Device deleted"}
|
||||
|
||||
|
||||
@router.get("/whoami", response_model=WhoAmIResponse)
|
||||
@@ -344,4 +358,5 @@ async def whoami(
|
||||
"email": email,
|
||||
"avatar_url": avatar_url,
|
||||
"is_admin": is_admin,
|
||||
"preferred_language": profile.preferred_language if profile else None,
|
||||
}
|
||||
|
||||
@@ -56,6 +56,7 @@ async def exchange_onedrive_token(
|
||||
# Return just what's needed by the frontend
|
||||
return {
|
||||
"refresh_token": token_data["refresh_token"],
|
||||
"access_token": token_data.get("access_token", ""),
|
||||
"expires_in": token_data.get("expires_in", 3600),
|
||||
}
|
||||
|
||||
@@ -183,6 +184,102 @@ async def test_onedrive_token(request: Request):
|
||||
return {"status": "error", "message": f"Connection error: {str(e)}"}
|
||||
|
||||
|
||||
@router.post("/onedrive/list-folders")
|
||||
@require_login
|
||||
async def list_onedrive_folders(
|
||||
request: Request,
|
||||
access_token: Annotated[str, Form(...)],
|
||||
path: Annotated[str, Form()] = "",
|
||||
):
|
||||
"""
|
||||
List folders in a OneDrive account for the directory selector.
|
||||
|
||||
Accepts an OAuth access token (short-lived) and a path to list.
|
||||
Returns a flat list of folder entries under the given path.
|
||||
"""
|
||||
try:
|
||||
folder_path = path.strip().strip("/")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
}
|
||||
|
||||
# Build the Graph API URL for listing children
|
||||
if not folder_path or folder_path == "root":
|
||||
url = "https://graph.microsoft.com/v1.0/me/drive/root/children"
|
||||
else:
|
||||
url = f"https://graph.microsoft.com/v1.0/me/drive/root:/{folder_path}:/children"
|
||||
|
||||
# Only request folders and minimal fields
|
||||
params = {
|
||||
"$filter": "folder ne null",
|
||||
"$select": "name,id,parentReference,folder",
|
||||
"$top": "200",
|
||||
}
|
||||
|
||||
response = requests.get(
|
||||
url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Access token is invalid or expired. Please re-authorize.",
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"OneDrive list children failed: {response.status_code} {response.text}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"Failed to list OneDrive folders: {response.text}",
|
||||
)
|
||||
|
||||
data = response.json()
|
||||
folders = []
|
||||
for item in data.get("value", []):
|
||||
if "folder" in item:
|
||||
parent_path = ""
|
||||
if item.get("parentReference", {}).get("path"):
|
||||
# parentReference.path looks like /drive/root:/some/path
|
||||
raw_parent = item["parentReference"]["path"]
|
||||
prefix = "/drive/root:"
|
||||
if raw_parent.startswith(prefix):
|
||||
parent_path = raw_parent[len(prefix) :]
|
||||
elif raw_parent == "/drive/root":
|
||||
parent_path = ""
|
||||
|
||||
item_path = f"{parent_path}/{item['name']}" if parent_path else f"/{item['name']}"
|
||||
|
||||
folders.append(
|
||||
{
|
||||
"name": item["name"],
|
||||
"path": item_path,
|
||||
"id": item.get("id", ""),
|
||||
"child_count": item.get("folder", {}).get("childCount", 0),
|
||||
}
|
||||
)
|
||||
|
||||
# Sort folders alphabetically
|
||||
folders.sort(key=lambda f: f["name"].lower())
|
||||
|
||||
return {
|
||||
"folders": folders,
|
||||
"path": f"/{folder_path}" if folder_path else "/",
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.exception(f"Error listing OneDrive folders: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to list folders: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
def format_time_remaining(time_delta):
|
||||
"""Format a timedelta into a human-readable string."""
|
||||
if time_delta.total_seconds() <= 0:
|
||||
|
||||
@@ -117,8 +117,14 @@ PIPELINE_STEP_TYPES: dict[str, dict[str, Any]] = {
|
||||
},
|
||||
"classify": {
|
||||
"label": "Document Classification",
|
||||
"description": "Classify the document type using AI without full metadata extraction.",
|
||||
"config_schema": {},
|
||||
"description": "Classify the document type using built-in and custom rules (filename patterns, content keywords, metadata matching).",
|
||||
"config_schema": {
|
||||
"use_builtin_rules": {
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "Include the pre-built classification rules (invoice, contract, receipt, etc.).",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
"""QR code login API endpoints for mobile app authentication.
|
||||
|
||||
Provides a secure challenge-response flow for logging into the mobile app
|
||||
by scanning a QR code displayed in the web interface:
|
||||
|
||||
1. **Web user** calls ``POST /qr-auth/challenge`` → receives a time-limited
|
||||
challenge token (encoded in the QR code).
|
||||
2. **Web UI** polls ``GET /qr-auth/challenge/{id}/status`` to detect when
|
||||
the mobile app has claimed the challenge.
|
||||
3. **Mobile app** scans the QR code and calls ``POST /qr-auth/claim`` with
|
||||
the challenge token + device name → receives an API token.
|
||||
|
||||
Security properties:
|
||||
* Challenges expire after a configurable TTL (default 2 minutes).
|
||||
* Single-use: once claimed, a challenge cannot be reused (replay-safe).
|
||||
* Cryptographically random 64-byte tokens.
|
||||
* IP addresses are logged for audit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
import segno
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.middleware.audit_log import get_client_ip
|
||||
from app.utils.session_manager import (
|
||||
claim_qr_challenge,
|
||||
create_qr_challenge,
|
||||
get_challenge_status,
|
||||
)
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/qr-auth", tags=["qr-auth"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request / Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CreateChallengeResponse(BaseModel):
|
||||
"""Response after creating a QR login challenge."""
|
||||
|
||||
challenge_id: int
|
||||
challenge_token: str
|
||||
expires_at: datetime
|
||||
ttl_seconds: int = Field(description="Seconds until the challenge expires (use for client-side countdown).")
|
||||
qr_payload: str = Field(description="The string to encode in the QR code.")
|
||||
qr_code_svg: str = Field(description="Base64-encoded SVG data URI of the QR code, ready for use in an <img> src.")
|
||||
|
||||
|
||||
class ChallengeStatusResponse(BaseModel):
|
||||
"""Response for polling the status of a QR challenge."""
|
||||
|
||||
id: int
|
||||
status: str # "pending", "claimed", "expired", "cancelled"
|
||||
device_name: str | None = None
|
||||
claimed_at: datetime | None = None
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class ClaimChallengeRequest(BaseModel):
|
||||
"""Request body for claiming a QR login challenge."""
|
||||
|
||||
challenge_token: str = Field(min_length=1, max_length=256)
|
||||
device_name: str = Field(
|
||||
default="Mobile App",
|
||||
min_length=1,
|
||||
max_length=120,
|
||||
description="Human-readable device name.",
|
||||
)
|
||||
|
||||
|
||||
class ClaimChallengeResponse(BaseModel):
|
||||
"""Response after successfully claiming a QR challenge."""
|
||||
|
||||
token: str
|
||||
token_id: int
|
||||
name: str
|
||||
owner_id: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# QR code rendering parameters
|
||||
_QR_ERROR_LEVEL = "M" # Medium error correction (~15% recovery); sufficient for on-screen display
|
||||
_QR_SCALE = 4 # Each QR module is rendered as 4×4 SVG pixels
|
||||
|
||||
|
||||
def _generate_qr_svg(payload: str) -> str:
|
||||
"""Generate a QR code for *payload* and return it as a base64 SVG data URI.
|
||||
|
||||
Using ``segno`` (pure-Python, no Pillow dependency) and SVG output so the
|
||||
QR code scales crisply at any resolution without requiring a canvas or any
|
||||
client-side JavaScript library.
|
||||
"""
|
||||
qr = segno.make(payload, error=_QR_ERROR_LEVEL)
|
||||
buf = io.BytesIO()
|
||||
qr.save(buf, kind="svg", scale=_QR_SCALE, xmldecl=False, svgclass=None, lineclass=None, omitsize=True)
|
||||
svg_bytes = buf.getvalue()
|
||||
return "data:image/svg+xml;base64," + base64.b64encode(svg_bytes).decode("ascii")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/challenge", status_code=status.HTTP_201_CREATED, response_model=CreateChallengeResponse)
|
||||
@require_login
|
||||
async def create_challenge(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new QR login challenge.
|
||||
|
||||
The returned ``qr_payload`` should be encoded into a QR code and
|
||||
displayed to the user. The mobile app scans this QR code and
|
||||
calls the ``/claim`` endpoint.
|
||||
"""
|
||||
ip = get_client_ip(request)
|
||||
challenge = create_qr_challenge(db, owner_id, ip_address=ip)
|
||||
|
||||
# The QR payload is a JSON-like string with enough info for the mobile
|
||||
# app to know the server URL and challenge token.
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
qr_payload = f"docuelevate://qr-login?token={challenge.challenge_token}&server={base_url}"
|
||||
|
||||
# Compute the TTL in seconds so the client can run a countdown timer
|
||||
# without comparing absolute timestamps (which breaks when client and
|
||||
# server clocks are out of sync).
|
||||
ttl_seconds = max(0, int((challenge.expires_at - challenge.created_at).total_seconds()))
|
||||
|
||||
return {
|
||||
"challenge_id": challenge.id,
|
||||
"challenge_token": challenge.challenge_token,
|
||||
"expires_at": challenge.expires_at,
|
||||
"ttl_seconds": ttl_seconds,
|
||||
"qr_payload": qr_payload,
|
||||
"qr_code_svg": _generate_qr_svg(qr_payload),
|
||||
}
|
||||
|
||||
|
||||
@router.get("/challenge/{challenge_id}/status", response_model=ChallengeStatusResponse)
|
||||
@require_login
|
||||
async def poll_challenge_status(
|
||||
request: Request,
|
||||
challenge_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Poll the status of a QR login challenge.
|
||||
|
||||
The web UI calls this endpoint every few seconds to check if the
|
||||
mobile app has scanned the QR code and claimed the challenge.
|
||||
"""
|
||||
result = get_challenge_status(db, challenge_id, owner_id)
|
||||
if not result:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Challenge not found")
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/claim", response_model=ClaimChallengeResponse)
|
||||
async def claim_challenge(
|
||||
request: Request,
|
||||
body: ClaimChallengeRequest,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Claim a QR login challenge and receive an API token.
|
||||
|
||||
This endpoint is called by the mobile app after scanning a QR code.
|
||||
It does **not** require authentication — the challenge token itself
|
||||
serves as proof that the user authorized this login from their web
|
||||
session.
|
||||
"""
|
||||
ip = get_client_ip(request)
|
||||
result = claim_qr_challenge(db, body.challenge_token, device_name=body.device_name, ip_address=ip)
|
||||
|
||||
if not result:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid, expired, or already claimed challenge.",
|
||||
)
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="qr_login_claimed",
|
||||
user=result["owner_id"],
|
||||
resource_type="session",
|
||||
ip_address=ip,
|
||||
details={"device_name": body.device_name, "token_id": result["token_id"]},
|
||||
severity="info",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write QR login audit event", exc_info=True)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,196 @@
|
||||
"""API endpoints for managing user sessions.
|
||||
|
||||
Provides endpoints for listing active sessions, revoking individual sessions,
|
||||
and the "log off everywhere" feature that invalidates all sessions and API
|
||||
tokens across all devices.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.middleware.audit_log import get_client_ip
|
||||
from app.utils.session_manager import (
|
||||
get_session_lifetime_days,
|
||||
list_user_sessions,
|
||||
revoke_all_sessions,
|
||||
revoke_session,
|
||||
)
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionResponse(BaseModel):
|
||||
"""Serialised user session for the management UI."""
|
||||
|
||||
id: int
|
||||
device_info: str | None
|
||||
ip_address: str | None
|
||||
created_at: datetime
|
||||
last_active_at: datetime
|
||||
expires_at: datetime
|
||||
is_current: bool = False
|
||||
|
||||
|
||||
class SessionListResponse(BaseModel):
|
||||
"""Response for listing active sessions."""
|
||||
|
||||
sessions: list[SessionResponse]
|
||||
session_lifetime_days: int
|
||||
|
||||
|
||||
class RevokeAllResponse(BaseModel):
|
||||
"""Response after revoking all sessions."""
|
||||
|
||||
revoked_count: int
|
||||
message: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", response_model=SessionListResponse)
|
||||
@require_login
|
||||
async def list_sessions(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""List all active sessions for the current user."""
|
||||
sessions = list_user_sessions(db, owner_id)
|
||||
|
||||
# Determine which session is the current one
|
||||
current_token = request.session.get("_session_token")
|
||||
|
||||
session_list = []
|
||||
for s in sessions:
|
||||
session_list.append(
|
||||
{
|
||||
"id": s.id,
|
||||
"device_info": s.device_info,
|
||||
"ip_address": s.ip_address,
|
||||
"created_at": s.created_at,
|
||||
"last_active_at": s.last_active_at,
|
||||
"expires_at": s.expires_at,
|
||||
"is_current": s.session_token == current_token if current_token else False,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"sessions": session_list,
|
||||
"session_lifetime_days": get_session_lifetime_days(),
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{session_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
async def revoke_single_session(
|
||||
request: Request,
|
||||
session_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> None:
|
||||
"""Revoke a specific session by ID."""
|
||||
success = revoke_session(db, session_id, owner_id)
|
||||
if not success:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="session_revoked",
|
||||
user=owner_id,
|
||||
resource_type="session",
|
||||
resource_id=str(session_id),
|
||||
ip_address=get_client_ip(request),
|
||||
severity="info",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write session revocation audit event", exc_info=True)
|
||||
|
||||
|
||||
@router.post("/revoke-all", response_model=RevokeAllResponse)
|
||||
@require_login
|
||||
async def revoke_all(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Revoke all sessions except the current one ("log off everywhere").
|
||||
|
||||
Also revokes all active API tokens for the user, which invalidates
|
||||
mobile app sessions and any programmatic access.
|
||||
"""
|
||||
# Find current session to preserve it
|
||||
current_token = request.session.get("_session_token")
|
||||
current_session_id = None
|
||||
if current_token:
|
||||
from app.models import UserSession
|
||||
|
||||
current = db.query(UserSession).filter(UserSession.session_token == current_token).first()
|
||||
if current:
|
||||
current_session_id = current.id
|
||||
|
||||
count = revoke_all_sessions(
|
||||
db,
|
||||
owner_id,
|
||||
except_session_id=current_session_id,
|
||||
revoke_api_tokens=True,
|
||||
)
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="revoke_all_sessions",
|
||||
user=owner_id,
|
||||
resource_type="session",
|
||||
ip_address=get_client_ip(request),
|
||||
details={"revoked_count": count},
|
||||
severity="warning",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write revoke-all audit event", exc_info=True)
|
||||
|
||||
return {
|
||||
"revoked_count": count,
|
||||
"message": f"Successfully revoked {count} session(s) and all API tokens.",
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
"""
|
||||
System reset API endpoints for DocuElevate.
|
||||
|
||||
Provides admin-only REST endpoints for:
|
||||
- Full system reset (wipe all user data)
|
||||
- Reset with re-import (move originals → reimport folder, wipe, re-ingest)
|
||||
|
||||
Both operations require the ``ENABLE_FACTORY_RESET=True`` feature flag and
|
||||
admin privileges.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/system-reset", tags=["system-reset"])
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
def _require_feature_enabled() -> None:
|
||||
"""Raise 404 when the factory-reset feature flag is off."""
|
||||
if not settings.enable_factory_reset:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="System reset is not enabled. Set ENABLE_FACTORY_RESET=True to activate.",
|
||||
)
|
||||
|
||||
|
||||
class ResetRequest(BaseModel):
|
||||
"""Body for system reset endpoints. Requires explicit confirmation."""
|
||||
|
||||
confirmation: str
|
||||
|
||||
|
||||
@router.post("/full")
|
||||
async def full_reset(
|
||||
body: ResetRequest,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Wipe all user data (database + work-files).
|
||||
|
||||
The caller must send ``{"confirmation": "DELETE"}`` to proceed.
|
||||
"""
|
||||
_require_feature_enabled()
|
||||
|
||||
if body.confirmation != "DELETE":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail='Confirmation required: send {"confirmation": "DELETE"} to proceed.',
|
||||
)
|
||||
|
||||
from app.utils.system_reset import perform_full_reset
|
||||
|
||||
try:
|
||||
result = perform_full_reset(db)
|
||||
except Exception as exc:
|
||||
logger.exception("Full system reset failed")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"System reset failed: {exc}",
|
||||
) from exc
|
||||
|
||||
return {"status": "ok", "result": result}
|
||||
|
||||
|
||||
@router.post("/reimport")
|
||||
async def reset_and_reimport(
|
||||
body: ResetRequest,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Move original files to a reimport folder, wipe everything, and
|
||||
configure the reimport folder as a watch folder for automatic
|
||||
re-ingestion.
|
||||
|
||||
The caller must send ``{"confirmation": "REIMPORT"}`` to proceed.
|
||||
"""
|
||||
_require_feature_enabled()
|
||||
|
||||
if body.confirmation != "REIMPORT":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail='Confirmation required: send {"confirmation": "REIMPORT"} to proceed.',
|
||||
)
|
||||
|
||||
from app.utils.system_reset import perform_reset_and_reimport
|
||||
|
||||
try:
|
||||
result = perform_reset_and_reimport(db)
|
||||
except Exception as exc:
|
||||
logger.exception("Reset-and-reimport failed")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Reset and reimport failed: {exc}",
|
||||
) from exc
|
||||
|
||||
return {"status": "ok", "result": result}
|
||||
|
||||
|
||||
@router.get("/status")
|
||||
async def reset_status(_admin: AdminUser) -> dict:
|
||||
"""Return whether the system reset feature is enabled."""
|
||||
return {
|
||||
"enabled": settings.enable_factory_reset,
|
||||
"factory_reset_on_startup": settings.factory_reset_on_startup,
|
||||
}
|
||||
@@ -11,11 +11,12 @@ from typing import Optional
|
||||
|
||||
import aiofiles
|
||||
import httpx
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel, HttpUrl, field_validator
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
from app.tasks.process_document import process_document
|
||||
from app.utils.allowed_types import ALLOWED_MIME_TYPES
|
||||
from app.utils.filename_utils import sanitize_filename
|
||||
@@ -107,7 +108,11 @@ def validate_file_type(content_type: str, filename: str) -> bool:
|
||||
|
||||
@router.post("/process-url")
|
||||
@require_login
|
||||
async def process_url(request: Request, url_request: URLUploadRequest):
|
||||
async def process_url(
|
||||
request: Request,
|
||||
url_request: URLUploadRequest,
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
):
|
||||
"""
|
||||
Download a file from a URL and enqueue it for processing.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user