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.
|
||||
|
||||
|
||||
+122
-3
@@ -103,11 +103,18 @@ if AUTH_ENABLED and settings.social_auth_apple_enabled:
|
||||
logger.warning("SOCIAL_AUTH_APPLE_ENABLED=true but client ID/team ID not configured")
|
||||
|
||||
if AUTH_ENABLED and settings.social_auth_dropbox_enabled:
|
||||
if settings.social_auth_dropbox_client_id and settings.social_auth_dropbox_client_secret:
|
||||
# Determine which credentials to use for Dropbox social login
|
||||
_dropbox_client_id = settings.social_auth_dropbox_client_id
|
||||
_dropbox_client_secret = settings.social_auth_dropbox_client_secret
|
||||
if settings.social_auth_dropbox_use_global_credentials and not _dropbox_client_id:
|
||||
_dropbox_client_id = settings.dropbox_app_key
|
||||
_dropbox_client_secret = settings.dropbox_app_secret
|
||||
|
||||
if _dropbox_client_id and _dropbox_client_secret:
|
||||
oauth.register(
|
||||
name="dropbox",
|
||||
client_id=settings.social_auth_dropbox_client_id,
|
||||
client_secret=settings.social_auth_dropbox_client_secret,
|
||||
client_id=_dropbox_client_id,
|
||||
client_secret=_dropbox_client_secret,
|
||||
authorize_url="https://www.dropbox.com/oauth2/authorize",
|
||||
access_token_url="https://api.dropboxapi.com/oauth2/token",
|
||||
userinfo_endpoint="https://api.dropboxapi.com/2/users/get_current_account",
|
||||
@@ -129,6 +136,25 @@ def get_current_user(request: Request):
|
||||
return api_user
|
||||
session_user = request.session.get("user")
|
||||
if session_user:
|
||||
# Validate server-side session if a session token is present
|
||||
session_token = request.session.get("_session_token")
|
||||
if session_token:
|
||||
try:
|
||||
from app.database import SessionLocal
|
||||
from app.utils.session_manager import validate_session
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
valid = validate_session(db, session_token)
|
||||
if not valid:
|
||||
logger.debug("[AUTH] get_current_user: server-side session invalid — clearing")
|
||||
request.session.pop("user", None)
|
||||
request.session.pop("_session_token", None)
|
||||
return None
|
||||
finally:
|
||||
db.close()
|
||||
except Exception:
|
||||
logger.debug("[AUTH] get_current_user: session validation error", exc_info=True)
|
||||
logger.debug(
|
||||
"[AUTH] get_current_user: resolved from session (user=%s)",
|
||||
session_user.get("preferred_username") or session_user.get("email") or session_user.get("id"),
|
||||
@@ -167,6 +193,16 @@ def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
|
||||
logger.debug("[AUTH] _resolve_bearer_user: no active API token matched the provided hash")
|
||||
return None
|
||||
|
||||
# Reject tokens that have passed their optional expiry.
|
||||
if db_token.expires_at is not None:
|
||||
now_utc = datetime.now(timezone.utc)
|
||||
expires_aware = db_token.expires_at
|
||||
if expires_aware.tzinfo is None:
|
||||
expires_aware = expires_aware.replace(tzinfo=timezone.utc)
|
||||
if now_utc > expires_aware:
|
||||
logger.debug("[AUTH] _resolve_bearer_user: API token id=%s has expired", db_token.id)
|
||||
return None
|
||||
|
||||
logger.debug(
|
||||
"[AUTH] _resolve_bearer_user: matched API token id=%s owner=%s",
|
||||
db_token.id,
|
||||
@@ -481,6 +517,27 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
|
||||
|
||||
request.session["user"] = user_data
|
||||
|
||||
# Create server-side session for tracking and revocation
|
||||
try:
|
||||
from app.utils.session_manager import create_session
|
||||
|
||||
_session_user_id = (
|
||||
user_data.get("sub")
|
||||
or user_data.get("preferred_username")
|
||||
or user_data.get("email")
|
||||
or user_data.get("id")
|
||||
)
|
||||
if _session_user_id:
|
||||
user_session = create_session(
|
||||
db,
|
||||
user_id=_session_user_id,
|
||||
ip_address=get_client_ip(request),
|
||||
user_agent=request.headers.get("user-agent"),
|
||||
)
|
||||
request.session["_session_token"] = user_session.session_token
|
||||
except Exception:
|
||||
logger.debug("[AUTH] Failed to create server-side session for social user", exc_info=True)
|
||||
|
||||
# Auto-create or update UserProfile
|
||||
_ensure_user_profile(db, user_data, is_admin=False)
|
||||
|
||||
@@ -665,6 +722,27 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)):
|
||||
|
||||
request.session["user"] = user_data
|
||||
|
||||
# Create server-side session for tracking and revocation
|
||||
try:
|
||||
from app.utils.session_manager import create_session
|
||||
|
||||
_session_user_id = (
|
||||
user_data.get("sub")
|
||||
or user_data.get("preferred_username")
|
||||
or user_data.get("email")
|
||||
or user_data.get("id")
|
||||
)
|
||||
if _session_user_id:
|
||||
user_session = create_session(
|
||||
db,
|
||||
user_id=_session_user_id,
|
||||
ip_address=get_client_ip(request),
|
||||
user_agent=request.headers.get("user-agent"),
|
||||
)
|
||||
request.session["_session_token"] = user_session.session_token
|
||||
except Exception:
|
||||
logger.debug("[AUTH] Failed to create server-side session for OAuth user", exc_info=True)
|
||||
|
||||
# Auto-create or update UserProfile so the user appears in admin user management
|
||||
_ensure_user_profile(db, user_data, is_admin=is_admin)
|
||||
|
||||
@@ -918,6 +996,19 @@ async def auth(request: Request, db: Session = Depends(get_db)):
|
||||
return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302)
|
||||
user_data = _build_session_user(local_user)
|
||||
request.session["user"] = user_data
|
||||
# Create server-side session for tracking and revocation
|
||||
try:
|
||||
from app.utils.session_manager import create_session
|
||||
|
||||
user_session = create_session(
|
||||
db,
|
||||
user_id=local_user.email,
|
||||
ip_address=get_client_ip(request),
|
||||
user_agent=request.headers.get("user-agent"),
|
||||
)
|
||||
request.session["_session_token"] = user_session.session_token
|
||||
except Exception:
|
||||
logger.debug("[AUTH] Failed to create server-side session", exc_info=True)
|
||||
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", local_user.email)
|
||||
_record_login_event(db, request, local_user.email, success=True)
|
||||
_ensure_user_profile(db, user_data, is_admin=bool(local_user.is_admin))
|
||||
@@ -974,6 +1065,20 @@ async def auth(request: Request, db: Session = Depends(get_db)):
|
||||
"is_admin": True,
|
||||
}
|
||||
request.session["user"] = admin_user_data
|
||||
# Create server-side session for tracking and revocation
|
||||
try:
|
||||
from app.utils.session_manager import create_session
|
||||
|
||||
admin_user_id = settings.admin_username or "admin"
|
||||
user_session = create_session(
|
||||
db,
|
||||
user_id=admin_user_id,
|
||||
ip_address=get_client_ip(request),
|
||||
user_agent=request.headers.get("user-agent"),
|
||||
)
|
||||
request.session["_session_token"] = user_session.session_token
|
||||
except Exception:
|
||||
logger.debug("[AUTH] Failed to create server-side session for admin", exc_info=True)
|
||||
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", username)
|
||||
_record_login_event(db, request, username, success=True)
|
||||
_ensure_user_profile(db, admin_user_data, is_admin=True)
|
||||
@@ -1024,6 +1129,20 @@ async def logout(request: Request, db: Session = Depends(get_db)):
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write logout audit event for user=%s", username, exc_info=True)
|
||||
# Revoke server-side session
|
||||
session_token = request.session.get("_session_token")
|
||||
if session_token:
|
||||
try:
|
||||
from app.utils.session_manager import validate_session
|
||||
|
||||
user_session = validate_session(db, session_token)
|
||||
if user_session:
|
||||
user_session.is_revoked = True
|
||||
user_session.revoked_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
except Exception:
|
||||
logger.debug("[AUTH] Failed to revoke server-side session", exc_info=True)
|
||||
request.session.pop("_session_token", None)
|
||||
request.session.pop("user", None)
|
||||
return RedirectResponse(url="/login?message=You+have+been+logged+out+successfully", status_code=302)
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from app.tasks.batch_tasks import ( # noqa: F401
|
||||
sync_search_index,
|
||||
)
|
||||
from app.tasks.check_credentials import check_credentials
|
||||
from app.tasks.classify_document import classify_document_task # noqa: F401
|
||||
from app.tasks.compute_embedding import backfill_missing_embeddings, compute_document_embedding # noqa: F401
|
||||
from app.tasks.convert_to_pdf import convert_to_pdf # noqa: F401
|
||||
from app.tasks.convert_to_pdfa import convert_to_pdfa # noqa: F401
|
||||
@@ -53,6 +54,7 @@ from app.tasks.upload_to_onedrive import upload_to_onedrive # noqa: F401
|
||||
from app.tasks.upload_to_paperless import upload_to_paperless # noqa: F401
|
||||
from app.tasks.upload_to_s3 import upload_to_s3 # noqa: F401
|
||||
from app.tasks.upload_to_sftp import upload_to_sftp # noqa: F401
|
||||
from app.tasks.upload_to_sharepoint import upload_to_sharepoint # noqa: F401
|
||||
from app.tasks.upload_to_user_integration import upload_to_user_integration # noqa: F401
|
||||
from app.tasks.upload_to_webdav import upload_to_webdav # noqa: F401
|
||||
from app.tasks.upload_with_rclone import send_to_all_rclone_destinations, upload_with_rclone # noqa: F401
|
||||
|
||||
+129
@@ -13,6 +13,24 @@ class Settings(BaseSettings):
|
||||
|
||||
database_url: str
|
||||
redis_url: str
|
||||
|
||||
# Database connection-pool tuning (ignored for SQLite, which uses NullPool).
|
||||
db_pool_size: int = Field(
|
||||
default=10,
|
||||
description="Number of persistent connections kept in the pool per worker process.",
|
||||
)
|
||||
db_max_overflow: int = Field(
|
||||
default=20,
|
||||
description="Additional connections allowed beyond db_pool_size under burst load.",
|
||||
)
|
||||
db_pool_timeout: int = Field(
|
||||
default=30,
|
||||
description="Seconds to wait for a connection from the pool before raising a TimeoutError.",
|
||||
)
|
||||
db_pool_recycle: int = Field(
|
||||
default=1800,
|
||||
description="Recycle (close and reopen) connections after this many seconds to avoid stale connections.",
|
||||
)
|
||||
openai_api_key: str
|
||||
openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint
|
||||
openai_model: str = "gpt-4o-mini" # Default model
|
||||
@@ -102,6 +120,16 @@ class Settings(BaseSettings):
|
||||
dropbox_app_secret: Optional[str] = None
|
||||
dropbox_folder: Optional[str] = None
|
||||
dropbox_refresh_token: Optional[str] = None
|
||||
dropbox_allow_global_credentials_for_integrations: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When True, users may authorize their personal Dropbox integrations using the global "
|
||||
"DROPBOX_APP_KEY / DROPBOX_APP_SECRET credentials configured by the admin, without "
|
||||
"needing to create their own Dropbox app. The Dropbox OAuth flow is initiated "
|
||||
"server-side so the app secret is never exposed to the browser. "
|
||||
"Default: False (each user must supply their own app credentials)."
|
||||
),
|
||||
)
|
||||
|
||||
# Making Nextcloud optional
|
||||
nextcloud_enabled: bool = Field(
|
||||
@@ -165,6 +193,35 @@ class Settings(BaseSettings):
|
||||
google_docai_processor_id: Optional[str] = None
|
||||
google_docai_location: str = "us" # Processor location, e.g. "us" or "eu"
|
||||
external_hostname: str = "localhost" # Default to localhost
|
||||
public_base_url: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"The full public base URL of the application, including scheme "
|
||||
"(e.g., 'https://docuelevate.example.com'). "
|
||||
"When set, this overrides the auto-detected URL for OAuth redirect URIs. "
|
||||
"This is required when the application is behind a reverse proxy that does "
|
||||
"not forward X-Forwarded-Proto headers correctly."
|
||||
),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Document Translation Settings
|
||||
# ---------------------------------------------------------------------------
|
||||
# Default target language for automatic document translation (ISO 639-1 code).
|
||||
# After OCR / metadata extraction, if the detected document language differs
|
||||
# from this value the system translates the extracted text into this language
|
||||
# and stores it alongside the original. Other language translations are
|
||||
# generated on the fly via the AI provider and are NOT persisted.
|
||||
# Per-user overrides are stored in UserProfile.default_document_language.
|
||||
default_document_language: str = Field(
|
||||
default="en",
|
||||
description=(
|
||||
"ISO 639-1 language code for the default translation target "
|
||||
"(e.g. 'en', 'de', 'fr'). Documents whose detected language "
|
||||
"differs are automatically translated into this language after "
|
||||
"processing. Default: 'en' (English)."
|
||||
),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Document Translation Settings
|
||||
@@ -190,6 +247,26 @@ class Settings(BaseSettings):
|
||||
admin_username: Optional[str] = None
|
||||
admin_password: Optional[str] = None
|
||||
session_secret: Optional[str] = None
|
||||
session_lifetime_days: int = Field(
|
||||
default=30,
|
||||
description=(
|
||||
"Session lifetime in days. Common values: 30, 60, 90. "
|
||||
"Determines how long a user stays logged in before being required to re-authenticate. "
|
||||
"Applies to both browser sessions and the session cookie max_age."
|
||||
),
|
||||
)
|
||||
session_lifetime_custom_days: int | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Override session_lifetime_days with a custom value. "
|
||||
"When set, this takes precedence over session_lifetime_days. "
|
||||
"Useful for admin-configured non-standard durations."
|
||||
),
|
||||
)
|
||||
qr_login_challenge_ttl_seconds: int = Field(
|
||||
default=120,
|
||||
description="Time-to-live in seconds for QR login challenges (default: 2 minutes).",
|
||||
)
|
||||
admin_group_name: str = "admin"
|
||||
|
||||
# Multi-user settings
|
||||
@@ -279,6 +356,16 @@ class Settings(BaseSettings):
|
||||
social_auth_dropbox_enabled: bool = False
|
||||
social_auth_dropbox_client_id: Optional[str] = None
|
||||
social_auth_dropbox_client_secret: Optional[str] = None
|
||||
social_auth_dropbox_use_global_credentials: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When True, Dropbox social login uses the global DROPBOX_APP_KEY / DROPBOX_APP_SECRET "
|
||||
"credentials (the storage integration credentials) instead of requiring separate "
|
||||
"SOCIAL_AUTH_DROPBOX_CLIENT_ID / SOCIAL_AUTH_DROPBOX_CLIENT_SECRET values. "
|
||||
"Requires SOCIAL_AUTH_DROPBOX_ENABLED=True and the global Dropbox app credentials to be set. "
|
||||
"Default: False."
|
||||
),
|
||||
)
|
||||
|
||||
# Local user signup
|
||||
allow_local_signup: bool = Field(
|
||||
@@ -589,6 +676,15 @@ class Settings(BaseSettings):
|
||||
onedrive_refresh_token: Optional[str] = None # Required for personal accounts
|
||||
onedrive_folder_path: Optional[str] = None
|
||||
|
||||
# SharePoint settings
|
||||
sharepoint_client_id: Optional[str] = None
|
||||
sharepoint_client_secret: Optional[str] = None
|
||||
sharepoint_tenant_id: Optional[str] = "common"
|
||||
sharepoint_refresh_token: Optional[str] = None
|
||||
sharepoint_site_url: Optional[str] = None # e.g. https://tenant.sharepoint.com/sites/sitename
|
||||
sharepoint_document_library: Optional[str] = "Documents" # Document library name
|
||||
sharepoint_folder_path: Optional[str] = None # Subfolder inside the library
|
||||
|
||||
# AWS S3 settings
|
||||
s3_enabled: bool = Field(
|
||||
default=True,
|
||||
@@ -639,6 +735,25 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
# System reset / factory reset settings
|
||||
factory_reset_on_startup: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When enabled, DocuElevate wipes all user data (database rows and "
|
||||
"work-files on disk) on every startup so the instance always comes "
|
||||
"up in a clean, fresh state. Useful for demo or testing environments. "
|
||||
"Default: False."
|
||||
),
|
||||
)
|
||||
enable_factory_reset: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Show the 'System Reset' page in the admin UI. When enabled, "
|
||||
"administrators can trigger a full data wipe or a wipe-and-reimport "
|
||||
"directly from the web interface. Default: False."
|
||||
),
|
||||
)
|
||||
|
||||
# PDF/A archival conversion settings
|
||||
enable_pdfa_conversion: bool = Field(
|
||||
default=False,
|
||||
@@ -1073,6 +1188,20 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
# Per-user upload rate limiting (health-aware, Redis-backed sliding window)
|
||||
upload_rate_limit_per_user: int = Field(
|
||||
default=20,
|
||||
description=(
|
||||
"Maximum number of file uploads allowed per user within the sliding window. "
|
||||
"The effective limit may be reduced dynamically when the system is under heavy load "
|
||||
"(high queue depth or CPU usage). Set to 0 to disable per-user upload rate limiting."
|
||||
),
|
||||
)
|
||||
upload_rate_limit_window: int = Field(
|
||||
default=60,
|
||||
description="Sliding window size in seconds for per-user upload rate limiting (default: 60).",
|
||||
)
|
||||
|
||||
# Rate Limiting Configuration (see SECURITY_AUDIT.md and docs/API.md)
|
||||
# Protects against DoS attacks and API abuse
|
||||
rate_limiting_enabled: bool = Field(
|
||||
|
||||
+31
-2
@@ -10,6 +10,7 @@ from typing import Any
|
||||
from sqlalchemy import create_engine, exc
|
||||
from sqlalchemy.engine.url import make_url
|
||||
from sqlalchemy.orm import Session, declarative_base, sessionmaker
|
||||
from sqlalchemy.pool import NullPool, QueuePool
|
||||
|
||||
from app.config import settings
|
||||
|
||||
@@ -17,9 +18,37 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
Base = declarative_base()
|
||||
|
||||
# Parse the DATABASE_URL
|
||||
# ---------------------------------------------------------------------------
|
||||
# Engine construction
|
||||
# ---------------------------------------------------------------------------
|
||||
DB_URL = settings.database_url
|
||||
engine = create_engine(DB_URL, connect_args={"check_same_thread": False})
|
||||
_parsed_url = make_url(DB_URL)
|
||||
|
||||
_connect_args: dict[str, Any] = {}
|
||||
_engine_kwargs: dict[str, Any] = {
|
||||
"pool_pre_ping": True, # detect stale / dropped connections before use
|
||||
}
|
||||
|
||||
if _parsed_url.get_backend_name() == "sqlite":
|
||||
# SQLite does not benefit from connection pooling and is prone to
|
||||
# QueuePool exhaustion under concurrent access. NullPool opens a fresh
|
||||
# connection for each request and closes it immediately afterwards,
|
||||
# completely avoiding the "QueuePool limit reached" TimeoutError.
|
||||
_connect_args["check_same_thread"] = False
|
||||
_engine_kwargs["poolclass"] = NullPool
|
||||
else:
|
||||
# PostgreSQL / MySQL — use a bounded QueuePool with configurable limits.
|
||||
_engine_kwargs["poolclass"] = QueuePool
|
||||
_engine_kwargs.update(
|
||||
{
|
||||
"pool_size": settings.db_pool_size,
|
||||
"max_overflow": settings.db_max_overflow,
|
||||
"pool_timeout": settings.db_pool_timeout,
|
||||
"pool_recycle": settings.db_pool_recycle,
|
||||
}
|
||||
)
|
||||
|
||||
engine = create_engine(DB_URL, connect_args=_connect_args, **_engine_kwargs)
|
||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||
|
||||
|
||||
|
||||
+18
-1
@@ -170,6 +170,12 @@ async def lifespan(app: FastAPI):
|
||||
# Startup: Initialize database
|
||||
init_db() # Create tables if they don't exist
|
||||
|
||||
# Factory reset on startup — wipe all user data before anything else
|
||||
if settings.factory_reset_on_startup:
|
||||
from app.utils.system_reset import perform_startup_reset
|
||||
|
||||
perform_startup_reset()
|
||||
|
||||
# Load settings from database after DB initialization
|
||||
from app.database import SessionLocal
|
||||
from app.utils.config_loader import load_settings_from_db
|
||||
@@ -318,8 +324,19 @@ app.add_middleware(CSRFMiddleware, config=settings)
|
||||
# See SECURITY_AUDIT.md – Infrastructure Security section
|
||||
app.add_middleware(AuditLogMiddleware, config=settings)
|
||||
|
||||
|
||||
# 3) Session Middleware (for request.session to work)
|
||||
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET)
|
||||
def _get_session_max_age() -> int:
|
||||
"""Compute session max-age at startup time."""
|
||||
try:
|
||||
from app.utils.session_manager import get_session_max_age_seconds
|
||||
|
||||
return get_session_max_age_seconds()
|
||||
except Exception:
|
||||
return 30 * 86400 # 30 days default fallback
|
||||
|
||||
|
||||
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET, max_age=_get_session_max_age())
|
||||
|
||||
# 3a) CORS Middleware - handles cross-origin requests and preflight (OPTIONS) responses.
|
||||
# Disabled by default: set CORS_ENABLED=True only when NOT using a reverse proxy
|
||||
|
||||
@@ -20,6 +20,9 @@ How it works:
|
||||
|
||||
Exempt paths (CSRF is not checked even for state-changing methods):
|
||||
- ``/oauth-callback`` – OAuth 2.0 callback; protected by the ``state`` parameter.
|
||||
- ``/api/qr-auth/claim`` – Called by the unauthenticated mobile app; the
|
||||
cryptographically-random, single-use challenge token provides equivalent
|
||||
protection.
|
||||
"""
|
||||
|
||||
import logging
|
||||
@@ -39,6 +42,10 @@ CSRF_PROTECTED_METHODS = {"POST", "PUT", "DELETE", "PATCH"}
|
||||
# their own replay-protection mechanism).
|
||||
CSRF_EXEMPT_PATHS = {
|
||||
"/oauth-callback",
|
||||
# The mobile app calls this endpoint without a browser session/CSRF token.
|
||||
# The cryptographically-random, single-use challenge token already provides
|
||||
# equivalent protection against cross-site request forgery.
|
||||
"/api/qr-auth/claim",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,290 @@
|
||||
"""Per-user, health-aware upload rate limiter for DocuElevate.
|
||||
|
||||
This module provides a FastAPI dependency that enforces per-user upload rate
|
||||
limits using a Redis-backed sliding window counter. The effective limit is
|
||||
dynamically reduced when the system is under heavy load (high Celery queue
|
||||
depth or elevated CPU load average), ensuring the server remains responsive
|
||||
to all users even during bulk-upload scenarios.
|
||||
|
||||
Usage in an endpoint::
|
||||
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
|
||||
@router.post("/ui-upload")
|
||||
@require_login
|
||||
async def ui_upload(
|
||||
request: Request,
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
...
|
||||
):
|
||||
...
|
||||
|
||||
See ``docs/ConfigurationGuide.md`` for the configuration options
|
||||
(``UPLOAD_RATE_LIMIT_PER_USER``, ``UPLOAD_RATE_LIMIT_WINDOW``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import redis
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
||||
from app.config import settings
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Redis key prefix
|
||||
# ---------------------------------------------------------------------------
|
||||
_KEY_PREFIX = "docuelevate:upload_rate"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health-check queue names (Celery defaults used by DocuElevate)
|
||||
# ---------------------------------------------------------------------------
|
||||
_CELERY_QUEUES = ("document_processor", "default", "celery")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Singleton Redis client (lazy-initialised; fail-open when unavailable)
|
||||
# ---------------------------------------------------------------------------
|
||||
_redis_client: redis.Redis | None = None
|
||||
|
||||
|
||||
def _get_redis() -> redis.Redis | None:
|
||||
"""Return a shared Redis client, or *None* when Redis is unavailable."""
|
||||
global _redis_client
|
||||
if _redis_client is not None:
|
||||
return _redis_client
|
||||
try:
|
||||
_redis_client = redis.Redis.from_url(
|
||||
settings.redis_url,
|
||||
decode_responses=True,
|
||||
socket_connect_timeout=2,
|
||||
socket_timeout=2,
|
||||
)
|
||||
# Quick connectivity check – raises on failure.
|
||||
_redis_client.ping()
|
||||
return _redis_client
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Redis unavailable for upload rate limiter – falling back to allow-all", exc_info=True)
|
||||
_redis_client = None
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health metrics helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_queue_depth(r: redis.Redis) -> int:
|
||||
"""Return the total number of pending tasks across all Celery queues."""
|
||||
total = 0
|
||||
for queue_name in _CELERY_QUEUES:
|
||||
try:
|
||||
total += r.llen(queue_name)
|
||||
except Exception: # noqa: BLE001, S110
|
||||
logger.debug("Could not read queue length for %r", queue_name, exc_info=True)
|
||||
return total
|
||||
|
||||
|
||||
def _get_cpu_load_ratio() -> float:
|
||||
"""Return the 1-minute load average divided by the number of CPU cores.
|
||||
|
||||
Returns ``0.0`` on platforms that do not support :func:`os.getloadavg`
|
||||
(e.g. Windows) so that the limiter never penalises on those systems.
|
||||
"""
|
||||
try:
|
||||
load_1m = os.getloadavg()[0]
|
||||
cpu_count = os.cpu_count() or 1
|
||||
return load_1m / cpu_count
|
||||
except (OSError, AttributeError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def compute_effective_limit(
|
||||
base_limit: int,
|
||||
queue_depth: int = 0,
|
||||
cpu_load_ratio: float = 0.0,
|
||||
) -> tuple[int, float, str]:
|
||||
"""Compute the effective upload rate limit based on system health.
|
||||
|
||||
The function applies a *reduction factor* (``0.0 < factor ≤ 1.0``) to the
|
||||
configured base limit. Both queue depth and CPU load contribute
|
||||
independently; the lowest factor wins.
|
||||
|
||||
Args:
|
||||
base_limit: The configured maximum uploads per window.
|
||||
queue_depth: Total pending tasks in Celery queues.
|
||||
cpu_load_ratio: 1-minute load average divided by CPU count.
|
||||
|
||||
Returns:
|
||||
A 3-tuple of ``(effective_limit, factor, reason)`` where *reason*
|
||||
is a human-readable tag for logging.
|
||||
"""
|
||||
factor = 1.0
|
||||
reason = "normal"
|
||||
|
||||
# --- Queue-depth thresholds ---
|
||||
if queue_depth > 200:
|
||||
factor, reason = min(factor, 0.10), f"critical_queue({queue_depth})"
|
||||
elif queue_depth > 100:
|
||||
factor, reason = min(factor, 0.25), f"high_queue({queue_depth})"
|
||||
elif queue_depth > 50:
|
||||
factor, reason = min(factor, 0.50), f"moderate_queue({queue_depth})"
|
||||
|
||||
# --- CPU-load thresholds ---
|
||||
if cpu_load_ratio > 3.0:
|
||||
new_factor = 0.10
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"critical_cpu({cpu_load_ratio:.1f})"
|
||||
elif cpu_load_ratio > 2.0:
|
||||
new_factor = 0.25
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"high_cpu({cpu_load_ratio:.1f})"
|
||||
elif cpu_load_ratio > 1.5:
|
||||
new_factor = 0.50
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"moderate_cpu({cpu_load_ratio:.1f})"
|
||||
|
||||
effective = max(1, int(base_limit * factor))
|
||||
return effective, factor, reason
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core sliding-window check (Redis sorted set)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _check_and_record(
|
||||
r: redis.Redis,
|
||||
user_id: str,
|
||||
window: int,
|
||||
effective_limit: int,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Atomically check the user's upload count and record the new upload.
|
||||
|
||||
Uses a Redis sorted set where each member is a unique timestamp-based ID
|
||||
and the score is the Unix timestamp. Entries older than *window* seconds
|
||||
are pruned on every call so the set never grows unbounded.
|
||||
|
||||
Returns:
|
||||
``None`` if the request is allowed, or a ``dict`` with ``count``,
|
||||
``limit``, and ``retry_after`` if the limit is exceeded.
|
||||
"""
|
||||
key = f"{_KEY_PREFIX}:{user_id}"
|
||||
now = time.time()
|
||||
window_start = now - window
|
||||
|
||||
pipe = r.pipeline(transaction=True)
|
||||
# 1. Remove entries outside the window
|
||||
pipe.zremrangebyscore(key, "-inf", window_start)
|
||||
# 2. Count current entries
|
||||
pipe.zcard(key)
|
||||
# 3. Retrieve the oldest entry's score (to compute retry_after)
|
||||
pipe.zrange(key, 0, 0, withscores=True)
|
||||
results = pipe.execute()
|
||||
|
||||
current_count: int = results[1]
|
||||
oldest_entries: list = results[2]
|
||||
|
||||
if current_count >= effective_limit:
|
||||
# Compute how long until the oldest entry expires from the window.
|
||||
if oldest_entries:
|
||||
oldest_score = oldest_entries[0][1]
|
||||
retry_after = max(1, int((oldest_score + window) - now))
|
||||
else:
|
||||
retry_after = max(1, window // 2)
|
||||
return {
|
||||
"count": current_count,
|
||||
"limit": effective_limit,
|
||||
"retry_after": retry_after,
|
||||
}
|
||||
|
||||
# 4. Record this upload (unique member = timestamp with random suffix)
|
||||
member = f"{now}:{os.urandom(4).hex()}"
|
||||
pipe2 = r.pipeline(transaction=True)
|
||||
pipe2.zadd(key, {member: now})
|
||||
pipe2.expire(key, window + 60) # TTL slightly longer than window
|
||||
pipe2.execute()
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastAPI dependency
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def require_upload_rate_limit(request: Request) -> None:
|
||||
"""FastAPI dependency that enforces per-user upload rate limits.
|
||||
|
||||
The dependency is designed to **fail open**: if Redis is unavailable the
|
||||
request is allowed through so that uploads are never blocked by a
|
||||
monitoring outage.
|
||||
|
||||
Raises:
|
||||
HTTPException: 429 Too Many Requests when the per-user upload limit
|
||||
is exceeded. The ``Retry-After`` header indicates how many
|
||||
seconds the client should wait before retrying.
|
||||
"""
|
||||
r = _get_redis()
|
||||
if r is None:
|
||||
# Redis unavailable – fail open.
|
||||
return
|
||||
|
||||
# Identify the user (owner_id for multi-user, IP fallback).
|
||||
user_id = get_current_owner_id(request)
|
||||
if not user_id:
|
||||
user_id = f"ip:{request.client.host}" if request.client else "ip:unknown"
|
||||
|
||||
base_limit: int = settings.upload_rate_limit_per_user
|
||||
window: int = settings.upload_rate_limit_window
|
||||
|
||||
# Gather health metrics and compute effective limit.
|
||||
try:
|
||||
queue_depth = _get_queue_depth(r)
|
||||
except Exception: # noqa: BLE001
|
||||
queue_depth = 0
|
||||
|
||||
cpu_load_ratio = _get_cpu_load_ratio()
|
||||
effective_limit, factor, health_reason = compute_effective_limit(base_limit, queue_depth, cpu_load_ratio)
|
||||
|
||||
# Sliding-window check.
|
||||
try:
|
||||
rejection = _check_and_record(r, user_id, window, effective_limit)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("Upload rate-limit check failed (allowing request): %s", exc)
|
||||
return
|
||||
|
||||
if rejection is not None:
|
||||
retry_after = rejection["retry_after"]
|
||||
logger.warning(
|
||||
"Upload rate limit exceeded: user=%s count=%d/%d window=%ds health=%s retry_after=%ds",
|
||||
user_id,
|
||||
rejection["count"],
|
||||
rejection["limit"],
|
||||
window,
|
||||
health_reason,
|
||||
retry_after,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=(
|
||||
f"Upload rate limit exceeded ({rejection['count']}/{rejection['limit']} "
|
||||
f"in {window}s). Retry after {retry_after}s."
|
||||
),
|
||||
headers={"Retry-After": str(retry_after)},
|
||||
)
|
||||
|
||||
if factor < 1.0:
|
||||
logger.info(
|
||||
"Upload allowed with reduced limit: user=%s effective=%d/%d health=%s",
|
||||
user_id,
|
||||
effective_limit,
|
||||
base_limit,
|
||||
health_reason,
|
||||
)
|
||||
+130
@@ -625,6 +625,7 @@ class IntegrationType:
|
||||
EMAIL = "EMAIL"
|
||||
PAPERLESS = "PAPERLESS"
|
||||
RCLONE = "RCLONE"
|
||||
SHAREPOINT = "SHAREPOINT"
|
||||
ICLOUD = "ICLOUD"
|
||||
|
||||
ALL = {
|
||||
@@ -642,6 +643,7 @@ class IntegrationType:
|
||||
EMAIL,
|
||||
PAPERLESS,
|
||||
RCLONE,
|
||||
SHAREPOINT,
|
||||
ICLOUD,
|
||||
}
|
||||
|
||||
@@ -806,6 +808,9 @@ class ApiToken(Base):
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
revoked_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# Optional expiry: if set, the token is rejected after this timestamp.
|
||||
expires_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class SharedLink(Base):
|
||||
"""Shareable, time-limited or view-limited document link.
|
||||
@@ -957,6 +962,51 @@ class ScheduledJob(Base):
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class ClassificationRuleModel(Base):
|
||||
"""Custom document classification rule.
|
||||
|
||||
Rules are evaluated during the ``classify`` pipeline step to assign a
|
||||
category to a document. System-wide rules have ``owner_id IS NULL``;
|
||||
user-specific rules belong to a single owner.
|
||||
"""
|
||||
|
||||
__tablename__ = "classification_rules"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# NULL = system-wide rule visible to all users.
|
||||
owner_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# Human-readable rule name (unique per owner).
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# Target category (e.g. "invoice", "contract", "receipt").
|
||||
category = Column(String(100), nullable=False, index=True)
|
||||
|
||||
# Rule type: "filename_pattern", "content_keyword", or "metadata_match".
|
||||
rule_type = Column(String(50), nullable=False)
|
||||
|
||||
# The matching pattern:
|
||||
# - filename_pattern: a regex
|
||||
# - content_keyword: pipe-separated keywords
|
||||
# - metadata_match: "field=value"
|
||||
pattern = Column(String(1000), nullable=False)
|
||||
|
||||
# Higher priority rules are evaluated first (default 0).
|
||||
priority = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Whether pattern matching is case-sensitive.
|
||||
case_sensitive = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Disabled rules are skipped during classification.
|
||||
enabled = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
__table_args__ = (UniqueConstraint("owner_id", "name", name="uq_classification_rules_owner_name"),)
|
||||
|
||||
|
||||
class MobileDevice(Base):
|
||||
"""Registered mobile device for push notifications.
|
||||
|
||||
@@ -991,6 +1041,86 @@ class MobileDevice(Base):
|
||||
__table_args__ = (UniqueConstraint("owner_id", "push_token", name="uq_mobile_device_owner_token"),)
|
||||
|
||||
|
||||
class UserSession(Base):
|
||||
"""Server-side session tracking for invalidation and device management.
|
||||
|
||||
Each row represents an active browser or app session. The ``session_token``
|
||||
is stored in the user's cookie and validated on every authenticated request.
|
||||
Revoking a row (``is_revoked=True``) immediately terminates that session
|
||||
on the next request.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_sessions"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Cryptographically random token stored in the session cookie.
|
||||
session_token = Column(String(128), unique=True, nullable=False, index=True)
|
||||
|
||||
# Stable owner identifier — matches FileRecord.owner_id.
|
||||
user_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Client metadata for display in the session management UI.
|
||||
ip_address = Column(String(45), nullable=True)
|
||||
user_agent = Column(String(512), nullable=True)
|
||||
device_info = Column(String(255), nullable=True)
|
||||
|
||||
is_revoked = Column(Boolean, nullable=False, default=False)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
last_active_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
expires_at = Column(DateTime(timezone=True), nullable=False)
|
||||
revoked_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class QRLoginChallenge(Base):
|
||||
"""Time-limited QR code login challenge for mobile app authentication.
|
||||
|
||||
A logged-in web user generates a challenge that produces a QR code. The
|
||||
mobile app scans the QR code and calls the claim endpoint with the
|
||||
``challenge_token``. The server verifies the challenge is still valid,
|
||||
unclaimed, and unexpired, then issues an API token for the mobile app.
|
||||
|
||||
Security properties:
|
||||
* Time-bound (default 2 minutes).
|
||||
* Single-use (``is_claimed`` prevents replay).
|
||||
* Cryptographically random 64-byte token.
|
||||
* Bound to the creating user — only that user's mobile device receives a
|
||||
token.
|
||||
"""
|
||||
|
||||
__tablename__ = "qr_login_challenges"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Cryptographically random token encoded in the QR code.
|
||||
challenge_token = Column(String(128), unique=True, nullable=False, index=True)
|
||||
|
||||
# The user who created this challenge (from the web session).
|
||||
user_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Whether the challenge has been successfully claimed by a mobile app.
|
||||
is_claimed = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Whether the challenge has been explicitly cancelled or expired.
|
||||
is_cancelled = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
# IP address of the web client that created the challenge.
|
||||
created_by_ip = Column(String(45), nullable=True)
|
||||
|
||||
# IP address of the mobile client that claimed the challenge.
|
||||
claimed_by_ip = Column(String(45), nullable=True)
|
||||
|
||||
# Device name provided by the mobile app when claiming.
|
||||
device_name = Column(String(255), nullable=True)
|
||||
|
||||
# The API token ID that was issued to the mobile app (for audit trail).
|
||||
issued_token_id = Column(Integer, nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
expires_at = Column(DateTime(timezone=True), nullable=False)
|
||||
claimed_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class ComplianceTemplate(Base):
|
||||
"""Pre-built compliance configuration templates (GDPR, HIPAA, SOC2).
|
||||
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
"""Celery task for rule-based document classification.
|
||||
|
||||
This task is executed as a pipeline step (``step_type="classify"``). It
|
||||
applies built-in and user-defined classification rules against the document's
|
||||
filename, OCR text, and existing AI metadata to assign a ``document_type``
|
||||
category.
|
||||
|
||||
The result is stored in the ``ai_metadata`` JSON blob on the
|
||||
:class:`~app.models.FileRecord` (field ``classification``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.database import SessionLocal
|
||||
from app.models import ClassificationRuleModel, FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.classification_rules import (
|
||||
ClassificationResult,
|
||||
classify_document,
|
||||
db_rule_to_engine_rule,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
STEP_NAME = "classify_document"
|
||||
|
||||
|
||||
def _load_custom_rules(owner_id: str | None) -> list[Any]:
|
||||
"""Load enabled custom classification rules from the database.
|
||||
|
||||
Returns engine-level :class:`ClassificationRule` dataclass instances.
|
||||
Rules are loaded in priority-descending order. System rules
|
||||
(``owner_id IS NULL``) and the user's own rules are both included.
|
||||
"""
|
||||
with SessionLocal() as db:
|
||||
query = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.enabled.is_(True))
|
||||
if owner_id:
|
||||
query = query.filter(
|
||||
(ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == owner_id)
|
||||
)
|
||||
else:
|
||||
query = query.filter(ClassificationRuleModel.owner_id.is_(None))
|
||||
rules = query.order_by(ClassificationRuleModel.priority.desc()).all()
|
||||
return [db_rule_to_engine_rule(r) for r in rules]
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def classify_document_task(
|
||||
self: Any,
|
||||
file_id: int,
|
||||
owner_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Classify a document using rule-based matching.
|
||||
|
||||
This task:
|
||||
1. Loads the :class:`FileRecord` from the database.
|
||||
2. Gathers filename, OCR text, and existing AI metadata.
|
||||
3. Loads built-in + user-defined classification rules.
|
||||
4. Runs the classification engine.
|
||||
5. Persists the result into ``ai_metadata.classification``.
|
||||
|
||||
Args:
|
||||
file_id: Primary key of the :class:`FileRecord` to classify.
|
||||
owner_id: Owner identifier for loading user-specific rules.
|
||||
|
||||
Returns:
|
||||
Dict with ``category``, ``confidence``, and ``matched_rules``.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"in_progress",
|
||||
f"Starting classification for file {file_id}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
file_record: FileRecord | None = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if file_record is None:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"failure",
|
||||
f"FileRecord {file_id} not found",
|
||||
file_id=file_id,
|
||||
)
|
||||
return {"status": "error", "detail": "File not found"}
|
||||
|
||||
# Gather inputs
|
||||
filename = file_record.original_filename or ""
|
||||
text = file_record.ocr_text or ""
|
||||
existing_metadata: dict[str, Any] = {}
|
||||
if file_record.ai_metadata:
|
||||
try:
|
||||
existing_metadata = json.loads(file_record.ai_metadata)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning("Failed to parse ai_metadata for file %s, starting fresh", file_id)
|
||||
existing_metadata = {}
|
||||
|
||||
# Load custom rules
|
||||
effective_owner = owner_id or file_record.owner_id
|
||||
custom_rules = _load_custom_rules(effective_owner)
|
||||
|
||||
# Run classification engine
|
||||
result: ClassificationResult = classify_document(
|
||||
filename=filename,
|
||||
text=text,
|
||||
metadata=existing_metadata,
|
||||
custom_rules=custom_rules,
|
||||
)
|
||||
|
||||
# Persist result into ai_metadata
|
||||
classification_data = {
|
||||
"category": result.category,
|
||||
"confidence": result.confidence,
|
||||
"matched_rules": [
|
||||
{
|
||||
"rule_name": m.rule_name,
|
||||
"rule_type": m.rule_type,
|
||||
"category": m.category,
|
||||
"confidence": m.confidence,
|
||||
}
|
||||
for m in result.matched_rules
|
||||
],
|
||||
}
|
||||
|
||||
existing_metadata["classification"] = classification_data
|
||||
|
||||
# If no document_type was set yet, populate it from the classification
|
||||
if not existing_metadata.get("document_type"):
|
||||
from app.utils.classification_rules import BUILTIN_CATEGORIES
|
||||
|
||||
existing_metadata["document_type"] = BUILTIN_CATEGORIES.get(
|
||||
result.category, result.category.replace("_", " ").title()
|
||||
)
|
||||
|
||||
file_record.ai_metadata = json.dumps(existing_metadata, ensure_ascii=False)
|
||||
db.commit()
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"success",
|
||||
f"Classified as '{result.category}' with confidence {result.confidence}",
|
||||
file_id=file_id,
|
||||
detail=f"Matched {len(result.matched_rules)} rule(s)",
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"category": result.category,
|
||||
"confidence": result.confidence,
|
||||
"matched_rules": len(result.matched_rules),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("Classification failed for file %s: %s", file_id, e)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"failure",
|
||||
f"Classification failed: {e}",
|
||||
file_id=file_id,
|
||||
)
|
||||
raise
|
||||
@@ -205,7 +205,7 @@ def convert_to_pdf(
|
||||
".pdf", # PDF (already in PDF format but can be processed)
|
||||
}
|
||||
|
||||
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".tif", ".webp", ".svg"}
|
||||
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".tif", ".webp", ".svg", ".heic", ".heif"}
|
||||
|
||||
HTML_EXTENSIONS = {".html", ".htm"}
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from app.tasks.upload_to_onedrive import upload_to_onedrive
|
||||
from app.tasks.upload_to_paperless import upload_to_paperless
|
||||
from app.tasks.upload_to_s3 import upload_to_s3
|
||||
from app.tasks.upload_to_sftp import upload_to_sftp
|
||||
from app.tasks.upload_to_sharepoint import upload_to_sharepoint
|
||||
from app.tasks.upload_to_webdav import upload_to_webdav
|
||||
from app.utils.config_validator import get_provider_status
|
||||
from app.utils.logging import log_task_progress
|
||||
@@ -121,6 +122,18 @@ def _should_upload_to_icloud():
|
||||
return bool(getattr(settings, "icloud_enabled", True) and settings.icloud_username and settings.icloud_password)
|
||||
|
||||
|
||||
def _should_upload_to_sharepoint():
|
||||
return bool(
|
||||
settings.sharepoint_client_id
|
||||
and settings.sharepoint_client_secret
|
||||
and settings.sharepoint_site_url
|
||||
and (
|
||||
settings.sharepoint_refresh_token
|
||||
or (settings.sharepoint_tenant_id and settings.sharepoint_tenant_id != "common")
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def get_configured_services_from_validator():
|
||||
"""
|
||||
Use the config validator to determine which services are configured and enabled.
|
||||
@@ -140,6 +153,7 @@ def get_configured_services_from_validator():
|
||||
"Email": "email",
|
||||
"OneDrive": "onedrive",
|
||||
"S3 Storage": "s3",
|
||||
"SharePoint": "sharepoint",
|
||||
"iCloud Drive": "icloud",
|
||||
}
|
||||
|
||||
@@ -250,6 +264,11 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
|
||||
"should_upload": _should_upload_to_s3,
|
||||
"upload_func": upload_to_s3,
|
||||
},
|
||||
{
|
||||
"name": "sharepoint",
|
||||
"should_upload": _should_upload_to_sharepoint,
|
||||
"upload_func": upload_to_sharepoint,
|
||||
},
|
||||
{
|
||||
"name": "icloud",
|
||||
"should_upload": _should_upload_to_icloud,
|
||||
|
||||
@@ -0,0 +1,338 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Upload documents to Microsoft SharePoint via the Microsoft Graph API.
|
||||
|
||||
This module authenticates using MSAL (same OAuth2 flow as OneDrive) and
|
||||
uploads files to a configurable SharePoint Online document library using
|
||||
the chunked upload session approach for reliability with large files.
|
||||
|
||||
Key differences from the OneDrive provider:
|
||||
- Uses ``/sites/{siteId}/drives/{driveId}`` instead of ``/me/drive``
|
||||
- Requires a SharePoint site URL to resolve the site and drive IDs
|
||||
- Targets a named document library (default: ``Documents``)
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import urllib.parse
|
||||
|
||||
import msal
|
||||
import requests
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_sharepoint_token() -> str:
|
||||
"""Acquire a Microsoft Graph API access token for SharePoint.
|
||||
|
||||
Uses MSAL ``ConfidentialClientApplication`` with the refresh-token flow
|
||||
(delegated permissions) or the client-credentials flow (application
|
||||
permissions) depending on configuration.
|
||||
|
||||
Returns:
|
||||
A valid access token string.
|
||||
|
||||
Raises:
|
||||
ValueError: When required settings are missing or token acquisition fails.
|
||||
"""
|
||||
if not settings.sharepoint_client_id or not settings.sharepoint_client_secret:
|
||||
raise ValueError("SharePoint client ID and client secret must be configured")
|
||||
|
||||
tenant = settings.sharepoint_tenant_id or "common"
|
||||
logger.info("Using SharePoint tenant: %s", tenant)
|
||||
|
||||
scopes = ["https://graph.microsoft.com/.default"]
|
||||
|
||||
if settings.sharepoint_refresh_token:
|
||||
app = msal.ConfidentialClientApplication(
|
||||
client_id=settings.sharepoint_client_id,
|
||||
client_credential=settings.sharepoint_client_secret,
|
||||
authority=f"https://login.microsoftonline.com/{tenant}",
|
||||
)
|
||||
|
||||
logger.info("Attempting to acquire SharePoint token using refresh token")
|
||||
token_response = app.acquire_token_by_refresh_token(
|
||||
refresh_token=settings.sharepoint_refresh_token, scopes=scopes
|
||||
)
|
||||
|
||||
if "access_token" not in token_response:
|
||||
error = token_response.get("error", "")
|
||||
error_desc = token_response.get("error_description", "Unknown error")
|
||||
logger.error("Failed to get SharePoint access token: %s - %s", error, error_desc)
|
||||
raise ValueError(f"Failed to get SharePoint access token: {error} - {error_desc}")
|
||||
|
||||
if "refresh_token" in token_response:
|
||||
settings.sharepoint_refresh_token = token_response["refresh_token"]
|
||||
logger.info("Updated SharePoint refresh token in memory")
|
||||
|
||||
return token_response["access_token"]
|
||||
|
||||
elif settings.sharepoint_tenant_id and settings.sharepoint_tenant_id != "common":
|
||||
authority = f"https://login.microsoftonline.com/{settings.sharepoint_tenant_id}"
|
||||
app = msal.ConfidentialClientApplication(
|
||||
client_id=settings.sharepoint_client_id,
|
||||
client_credential=settings.sharepoint_client_secret,
|
||||
authority=authority,
|
||||
)
|
||||
|
||||
token_response = app.acquire_token_for_client(scopes=scopes)
|
||||
|
||||
if "access_token" not in token_response:
|
||||
error = token_response.get("error", "")
|
||||
error_desc = token_response.get("error_description", "Unknown error")
|
||||
raise ValueError(f"Failed to get SharePoint access token: {error} - {error_desc}")
|
||||
|
||||
return token_response["access_token"]
|
||||
|
||||
else:
|
||||
raise ValueError("For SharePoint, either a refresh token or a non-'common' tenant ID is required")
|
||||
|
||||
|
||||
def resolve_sharepoint_drive(access_token: str, site_url: str, library_name: str) -> tuple[str, str]:
|
||||
"""Resolve the Graph API site ID and drive ID for a SharePoint site.
|
||||
|
||||
Args:
|
||||
access_token: Valid Microsoft Graph API token.
|
||||
site_url: Full SharePoint site URL, e.g.
|
||||
``https://tenant.sharepoint.com/sites/sitename``.
|
||||
library_name: Display name of the document library (e.g. ``Documents``).
|
||||
|
||||
Returns:
|
||||
A ``(site_id, drive_id)`` tuple.
|
||||
|
||||
Raises:
|
||||
ValueError: When the site URL cannot be parsed.
|
||||
RuntimeError: When the Graph API call fails.
|
||||
"""
|
||||
parsed = urllib.parse.urlparse(site_url)
|
||||
hostname = parsed.hostname
|
||||
site_path = parsed.path.rstrip("/")
|
||||
|
||||
if not hostname or not site_path:
|
||||
raise ValueError(
|
||||
f"Invalid SharePoint site URL '{site_url}'. Expected format: https://tenant.sharepoint.com/sites/sitename"
|
||||
)
|
||||
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
|
||||
# Resolve site ID
|
||||
site_api_url = f"https://graph.microsoft.com/v1.0/sites/{hostname}:{site_path}"
|
||||
logger.info("Resolving SharePoint site: %s", site_api_url)
|
||||
resp = requests.get(site_api_url, headers=headers, timeout=settings.http_request_timeout)
|
||||
|
||||
if resp.status_code != 200:
|
||||
raise RuntimeError(f"Failed to resolve SharePoint site: {resp.status_code} - {resp.text}")
|
||||
|
||||
site_id = resp.json()["id"]
|
||||
logger.info("Resolved SharePoint site ID: %s", site_id)
|
||||
|
||||
# Resolve drive ID from the document library name
|
||||
drives_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives"
|
||||
resp = requests.get(drives_url, headers=headers, timeout=settings.http_request_timeout)
|
||||
|
||||
if resp.status_code != 200:
|
||||
raise RuntimeError(f"Failed to list SharePoint drives: {resp.status_code} - {resp.text}")
|
||||
|
||||
drives = resp.json().get("value", [])
|
||||
drive_id = None
|
||||
for drive in drives:
|
||||
if drive.get("name", "").lower() == library_name.lower():
|
||||
drive_id = drive["id"]
|
||||
break
|
||||
|
||||
if not drive_id:
|
||||
available = [d.get("name") for d in drives]
|
||||
raise RuntimeError(f"Document library '{library_name}' not found on site. Available libraries: {available}")
|
||||
|
||||
logger.info("Resolved SharePoint drive ID: %s (library: %s)", drive_id, library_name)
|
||||
return site_id, drive_id
|
||||
|
||||
|
||||
def create_sharepoint_upload_session(
|
||||
filename: str, folder_path: str | None, drive_id: str, site_id: str, access_token: str
|
||||
) -> str:
|
||||
"""Create a resumable upload session on a SharePoint document library.
|
||||
|
||||
Args:
|
||||
filename: Name of the file to upload.
|
||||
folder_path: Optional subfolder path inside the library.
|
||||
drive_id: Graph API drive ID of the document library.
|
||||
site_id: Graph API site ID.
|
||||
access_token: Valid access token.
|
||||
|
||||
Returns:
|
||||
The upload session URL for chunked PUT requests.
|
||||
|
||||
Raises:
|
||||
RuntimeError: When session creation fails.
|
||||
"""
|
||||
base_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives/{drive_id}"
|
||||
|
||||
if folder_path:
|
||||
folder_path = folder_path.strip("/")
|
||||
path_components = folder_path.split("/")
|
||||
encoded_path = "/".join(urllib.parse.quote(component) for component in path_components)
|
||||
encoded_filename = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_path}/{encoded_filename}:/createUploadSession"
|
||||
else:
|
||||
encoded_filename = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_filename}:/createUploadSession"
|
||||
|
||||
url = f"{base_url}{item_path}"
|
||||
request_body = {"item": {"@microsoft.graph.conflictBehavior": "replace"}}
|
||||
headers = {"Authorization": f"Bearer {access_token}", "Content-Type": "application/json"}
|
||||
|
||||
logger.info("Creating SharePoint upload session for %s at path %s", filename, folder_path)
|
||||
response = requests.post(url, headers=headers, json=request_body, timeout=settings.http_request_timeout)
|
||||
|
||||
if response.status_code == 200:
|
||||
upload_url = response.json().get("uploadUrl")
|
||||
logger.info("SharePoint upload session created for %s", filename)
|
||||
return upload_url
|
||||
else:
|
||||
raise RuntimeError(f"Failed to create SharePoint upload session: {response.status_code} - {response.text}")
|
||||
|
||||
|
||||
def upload_large_file_sharepoint(file_path: str, upload_url: str) -> dict:
|
||||
"""Upload a file to SharePoint using a chunked upload session.
|
||||
|
||||
Args:
|
||||
file_path: Local path to the file.
|
||||
upload_url: The upload session URL from ``create_sharepoint_upload_session``.
|
||||
|
||||
Returns:
|
||||
The Graph API response dict containing file metadata.
|
||||
|
||||
Raises:
|
||||
RuntimeError: When a chunk upload fails after retries.
|
||||
"""
|
||||
file_size = os.path.getsize(file_path)
|
||||
chunk_size = 10 * 1024 * 1024 # 10 MB
|
||||
|
||||
response = None
|
||||
with open(file_path, "rb") as f:
|
||||
chunk_number = 0
|
||||
while True:
|
||||
chunk = f.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
|
||||
chunk_start = chunk_number * chunk_size
|
||||
chunk_end = chunk_start + len(chunk) - 1
|
||||
content_range = f"bytes {chunk_start}-{chunk_end}/{file_size}"
|
||||
|
||||
headers = {"Content-Length": str(len(chunk)), "Content-Range": content_range}
|
||||
|
||||
max_retries = 3
|
||||
retry_delay = 2
|
||||
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
response = requests.put(
|
||||
upload_url, headers=headers, data=chunk, timeout=settings.http_request_timeout
|
||||
)
|
||||
if response.status_code in (201, 202):
|
||||
break
|
||||
else:
|
||||
logger.warning(
|
||||
"SharePoint chunk upload failed (attempt %d): %d", attempt + 1, response.status_code
|
||||
)
|
||||
if attempt < max_retries - 1:
|
||||
time.sleep(retry_delay * (attempt + 1))
|
||||
except Exception as e:
|
||||
logger.warning("SharePoint chunk upload error (attempt %d): %s", attempt + 1, str(e))
|
||||
if attempt < max_retries - 1:
|
||||
time.sleep(retry_delay * (attempt + 1))
|
||||
|
||||
if response is None or response.status_code not in (201, 202):
|
||||
status = response.status_code if response else "no response"
|
||||
text = response.text if response else ""
|
||||
raise RuntimeError(f"Failed to upload chunk after {max_retries} attempts: {status} - {text}")
|
||||
|
||||
chunk_number += 1
|
||||
|
||||
return response.json() if response else {}
|
||||
|
||||
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_sharepoint(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""Upload a file to SharePoint Online.
|
||||
|
||||
Args:
|
||||
file_path: Path to the file to upload.
|
||||
file_id: Optional file ID to associate with logs.
|
||||
folder_override: Optional folder path override.
|
||||
|
||||
Returns:
|
||||
A dict with upload status and file details.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: When the file does not exist.
|
||||
ValueError: When SharePoint is not configured.
|
||||
RuntimeError: When the upload fails.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info("[%s] Starting SharePoint upload: %s", task_id, file_path)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_to_sharepoint",
|
||||
"in_progress",
|
||||
f"Uploading to SharePoint: {os.path.basename(file_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
if not os.path.exists(file_path):
|
||||
error_msg = f"File not found: {file_path}"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
|
||||
raise FileNotFoundError(error_msg)
|
||||
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not settings.sharepoint_client_id:
|
||||
error_msg = "SharePoint client ID is not configured"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
|
||||
raise ValueError(error_msg)
|
||||
|
||||
if not settings.sharepoint_site_url:
|
||||
error_msg = "SharePoint site URL is not configured"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
|
||||
raise ValueError(error_msg)
|
||||
|
||||
try:
|
||||
access_token = get_sharepoint_token()
|
||||
|
||||
library_name = settings.sharepoint_document_library or "Documents"
|
||||
site_id, drive_id = resolve_sharepoint_drive(access_token, settings.sharepoint_site_url, library_name)
|
||||
|
||||
folder_path = folder_override if folder_override is not None else settings.sharepoint_folder_path
|
||||
|
||||
upload_url = create_sharepoint_upload_session(filename, folder_path, drive_id, site_id, access_token)
|
||||
result = upload_large_file_sharepoint(file_path, upload_url)
|
||||
|
||||
web_url = result.get("webUrl", "Not available")
|
||||
logger.info("[%s] Successfully uploaded %s to SharePoint", task_id, filename)
|
||||
logger.info("[%s] File accessible at: %s", task_id, web_url)
|
||||
log_task_progress(
|
||||
task_id, "upload_to_sharepoint", "success", f"Uploaded to SharePoint: {filename}", file_id=file_id
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "Completed",
|
||||
"file_path": file_path,
|
||||
"sharepoint_path": f"{folder_path or ''}/{filename}",
|
||||
"web_url": web_url,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to upload {filename} to SharePoint: {str(e)}"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
|
||||
raise RuntimeError(error_msg) from e
|
||||
@@ -571,6 +571,113 @@ def _upload_rclone(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], t
|
||||
return {"status": "Completed", "rclone_dest": dest}
|
||||
|
||||
|
||||
def _upload_sharepoint(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to SharePoint using per-user MSAL credentials."""
|
||||
import urllib.parse
|
||||
|
||||
import msal
|
||||
import requests as _requests
|
||||
|
||||
client_id = creds.get("client_id") or ""
|
||||
client_secret = creds.get("client_secret") or ""
|
||||
refresh_token = creds.get("refresh_token") or ""
|
||||
tenant = cfg.get("tenant_id") or "common"
|
||||
site_url = cfg.get("site_url") or ""
|
||||
library_name = cfg.get("document_library") or "Documents"
|
||||
folder_path = cfg.get("folder_path") or ""
|
||||
|
||||
if not (client_id and client_secret):
|
||||
raise ValueError("SharePoint integration is missing client_id or client_secret in credentials")
|
||||
if not site_url:
|
||||
raise ValueError("SharePoint integration is missing site_url in config")
|
||||
|
||||
scopes = ["https://graph.microsoft.com/.default"]
|
||||
msal_app = msal.ConfidentialClientApplication(
|
||||
client_id=client_id,
|
||||
client_credential=client_secret,
|
||||
authority=f"https://login.microsoftonline.com/{tenant}",
|
||||
)
|
||||
|
||||
if refresh_token:
|
||||
token_resp = msal_app.acquire_token_by_refresh_token(refresh_token=refresh_token, scopes=scopes)
|
||||
else:
|
||||
token_resp = msal_app.acquire_token_for_client(scopes=scopes)
|
||||
|
||||
if "access_token" not in token_resp:
|
||||
raise ValueError(f"SharePoint token acquisition failed: {token_resp.get('error_description', 'unknown')}")
|
||||
|
||||
access_token = token_resp["access_token"]
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
|
||||
# Resolve site ID
|
||||
parsed = urllib.parse.urlparse(site_url)
|
||||
hostname = parsed.hostname
|
||||
site_path = parsed.path.rstrip("/")
|
||||
if not hostname or not site_path:
|
||||
raise ValueError(f"Invalid SharePoint site URL: {site_url}")
|
||||
|
||||
resp = _requests.get(f"https://graph.microsoft.com/v1.0/sites/{hostname}:{site_path}", headers=headers, timeout=30)
|
||||
resp.raise_for_status()
|
||||
site_id = resp.json()["id"]
|
||||
|
||||
# Resolve drive ID
|
||||
resp = _requests.get(f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives", headers=headers, timeout=30)
|
||||
resp.raise_for_status()
|
||||
drive_id = None
|
||||
for drive in resp.json().get("value", []):
|
||||
if drive.get("name", "").lower() == library_name.lower():
|
||||
drive_id = drive["id"]
|
||||
break
|
||||
if not drive_id:
|
||||
raise RuntimeError(f"Document library '{library_name}' not found on SharePoint site")
|
||||
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
# Build upload-session URL
|
||||
base_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives/{drive_id}"
|
||||
if folder_path:
|
||||
folder_path = folder_path.strip("/")
|
||||
encoded_path = "/".join(urllib.parse.quote(p) for p in folder_path.split("/"))
|
||||
encoded_file = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_path}/{encoded_file}:/createUploadSession"
|
||||
else:
|
||||
encoded_file = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_file}:/createUploadSession"
|
||||
|
||||
session_url = f"{base_url}{item_path}"
|
||||
session_headers = {"Authorization": f"Bearer {access_token}", "Content-Type": "application/json"}
|
||||
resp = _requests.post(
|
||||
session_url,
|
||||
headers=session_headers,
|
||||
json={"item": {"@microsoft.graph.conflictBehavior": "replace"}},
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
upload_url = resp.json()["uploadUrl"]
|
||||
|
||||
file_size = os.path.getsize(file_path)
|
||||
chunk_size = 10 * 1024 * 1024
|
||||
with open(file_path, "rb") as fh:
|
||||
chunk_num = 0
|
||||
while True:
|
||||
chunk = fh.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
start = chunk_num * chunk_size
|
||||
end = start + len(chunk) - 1
|
||||
upload_headers = {
|
||||
"Content-Length": str(len(chunk)),
|
||||
"Content-Range": f"bytes {start}-{end}/{file_size}",
|
||||
}
|
||||
upload_resp = _requests.put(upload_url, headers=upload_headers, data=chunk, timeout=120)
|
||||
if upload_resp.status_code not in (201, 202):
|
||||
raise RuntimeError(f"SharePoint chunk upload failed: {upload_resp.status_code}")
|
||||
chunk_num += 1
|
||||
|
||||
logger.info("[%s] SharePoint upload complete: %s/%s", task_id, folder_path, filename)
|
||||
return {"status": "Completed", "sharepoint_folder": folder_path, "filename": filename}
|
||||
|
||||
|
||||
def _upload_icloud(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to iCloud Drive using per-user credentials.
|
||||
|
||||
@@ -615,6 +722,7 @@ _UPLOAD_HANDLERS = {
|
||||
IntegrationType.PAPERLESS: _upload_paperless,
|
||||
IntegrationType.EMAIL: _upload_email,
|
||||
IntegrationType.RCLONE: _upload_rclone,
|
||||
IntegrationType.SHAREPOINT: _upload_sharepoint,
|
||||
IntegrationType.ICLOUD: _upload_icloud,
|
||||
}
|
||||
|
||||
|
||||
@@ -68,6 +68,8 @@ IMAGE_MIME_TYPES: set[str] = {
|
||||
"image/tiff",
|
||||
"image/webp",
|
||||
"image/svg+xml",
|
||||
"image/heic",
|
||||
"image/heif",
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -124,6 +126,8 @@ ALLOWED_EXTENSIONS: set[str] = {
|
||||
".tif",
|
||||
".webp",
|
||||
".svg",
|
||||
".heic",
|
||||
".heif",
|
||||
# Web
|
||||
".html",
|
||||
".htm",
|
||||
@@ -234,7 +238,7 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
|
||||
},
|
||||
"images": {
|
||||
"label": "Images",
|
||||
"description": "Image files (.jpg, .png, .gif, .bmp, .tiff, .webp, .svg)",
|
||||
"description": "Image files (.jpg, .png, .gif, .bmp, .tiff, .webp, .svg, .heic, .heif)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"image/jpeg",
|
||||
@@ -245,6 +249,8 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
|
||||
"image/tiff",
|
||||
"image/webp",
|
||||
"image/svg+xml",
|
||||
"image/heic",
|
||||
"image/heif",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset(
|
||||
@@ -258,6 +264,8 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
|
||||
".tif",
|
||||
".webp",
|
||||
".svg",
|
||||
".heic",
|
||||
".heif",
|
||||
}
|
||||
),
|
||||
},
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
"""
|
||||
Rule-based document classification engine.
|
||||
|
||||
Provides pre-built categories and a rule matcher that classifies documents
|
||||
using filename patterns, content keywords, and metadata fields. Custom
|
||||
rules stored in the database are evaluated alongside the built-in defaults.
|
||||
|
||||
Usage::
|
||||
|
||||
from app.utils.classification_rules import classify_document
|
||||
|
||||
result = classify_document(
|
||||
filename="2024-03-01_Invoice_Acme.pdf",
|
||||
text="Invoice total: $1,234.56",
|
||||
metadata={"absender": "Acme Corp"},
|
||||
custom_rules=custom_rules_from_db,
|
||||
)
|
||||
# result -> ClassificationResult(category="invoice", confidence=85, matched_rules=[...])
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-built categories
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Canonical category names recognized by the system. Users may also define
|
||||
#: their own categories via custom rules.
|
||||
BUILTIN_CATEGORIES: dict[str, str] = {
|
||||
"invoice": "Invoice",
|
||||
"contract": "Contract",
|
||||
"receipt": "Receipt",
|
||||
"letter": "Letter",
|
||||
"report": "Report",
|
||||
"bank_statement": "Bank Statement",
|
||||
"tax_document": "Tax Document",
|
||||
"insurance": "Insurance Document",
|
||||
"payslip": "Payslip",
|
||||
"unknown": "Unknown",
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Rule type constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
RULE_TYPE_FILENAME = "filename_pattern"
|
||||
RULE_TYPE_CONTENT = "content_keyword"
|
||||
RULE_TYPE_METADATA = "metadata_match"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Data classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClassificationRule:
|
||||
"""A single classification rule."""
|
||||
|
||||
name: str
|
||||
category: str
|
||||
rule_type: str # filename_pattern | content_keyword | metadata_match
|
||||
pattern: str # regex for filename, keyword(s) for content, "field=value" for metadata
|
||||
priority: int = 0 # higher = evaluated first
|
||||
case_sensitive: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.rule_type not in (RULE_TYPE_FILENAME, RULE_TYPE_CONTENT, RULE_TYPE_METADATA):
|
||||
raise ValueError(f"Invalid rule_type: {self.rule_type!r}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatchedRule:
|
||||
"""Records which rule matched and why."""
|
||||
|
||||
rule_name: str
|
||||
rule_type: str
|
||||
category: str
|
||||
confidence: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClassificationResult:
|
||||
"""The outcome of running the classification engine on a document."""
|
||||
|
||||
category: str
|
||||
confidence: int # 0 – 100
|
||||
matched_rules: list[MatchedRule] = field(default_factory=list)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Built-in rules
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
BUILTIN_RULES: list[ClassificationRule] = [
|
||||
# ── Invoice ───────────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_invoice_filename", "invoice", RULE_TYPE_FILENAME, r"(?i)invoice|rechnung|facture"),
|
||||
ClassificationRule(
|
||||
"builtin_invoice_content",
|
||||
"invoice",
|
||||
RULE_TYPE_CONTENT,
|
||||
"invoice number|invoice total|amount due|rechnung|rechnungsnummer|total amount|bill to",
|
||||
),
|
||||
ClassificationRule("builtin_invoice_metadata", "invoice", RULE_TYPE_METADATA, "document_type=Invoice"),
|
||||
ClassificationRule(
|
||||
"builtin_invoice_kommunikationsart", "invoice", RULE_TYPE_METADATA, "kommunikationsart=Rechnung"
|
||||
),
|
||||
# ── Contract ──────────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_contract_filename", "contract", RULE_TYPE_FILENAME, r"(?i)contract|vertrag|agreement"),
|
||||
ClassificationRule(
|
||||
"builtin_contract_content",
|
||||
"contract",
|
||||
RULE_TYPE_CONTENT,
|
||||
"hereby agrees|terms and conditions|vertrag|agreement between|party agrees|effective date",
|
||||
),
|
||||
ClassificationRule("builtin_contract_metadata", "contract", RULE_TYPE_METADATA, "document_type=Contract"),
|
||||
ClassificationRule(
|
||||
"builtin_contract_kommunikationsart", "contract", RULE_TYPE_METADATA, "kommunikationsart=Vertrag"
|
||||
),
|
||||
# ── Receipt ───────────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_receipt_filename", "receipt", RULE_TYPE_FILENAME, r"(?i)receipt|quittung|beleg"),
|
||||
ClassificationRule(
|
||||
"builtin_receipt_content",
|
||||
"receipt",
|
||||
RULE_TYPE_CONTENT,
|
||||
"receipt|quittung|payment received|thank you for your purchase|transaction id",
|
||||
),
|
||||
ClassificationRule("builtin_receipt_metadata", "receipt", RULE_TYPE_METADATA, "document_type=Receipt"),
|
||||
ClassificationRule(
|
||||
"builtin_receipt_kommunikationsart", "receipt", RULE_TYPE_METADATA, "kommunikationsart=Quittung"
|
||||
),
|
||||
# ── Letter ────────────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_letter_filename", "letter", RULE_TYPE_FILENAME, r"(?i)letter|brief|schreiben"),
|
||||
ClassificationRule(
|
||||
"builtin_letter_content",
|
||||
"letter",
|
||||
RULE_TYPE_CONTENT,
|
||||
"dear sir|dear madam|sehr geehrte|to whom it may concern|sincerely|mit freundlichen",
|
||||
),
|
||||
# ── Report ────────────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_report_filename", "report", RULE_TYPE_FILENAME, r"(?i)report|bericht"),
|
||||
ClassificationRule(
|
||||
"builtin_report_content",
|
||||
"report",
|
||||
RULE_TYPE_CONTENT,
|
||||
"executive summary|table of contents|annual report|quarterly report|findings",
|
||||
),
|
||||
# ── Bank statement ────────────────────────────────────────────────────
|
||||
ClassificationRule(
|
||||
"builtin_bank_filename",
|
||||
"bank_statement",
|
||||
RULE_TYPE_FILENAME,
|
||||
r"(?i)bank.?statement|kontoauszug",
|
||||
),
|
||||
ClassificationRule(
|
||||
"builtin_bank_content",
|
||||
"bank_statement",
|
||||
RULE_TYPE_CONTENT,
|
||||
"account statement|kontoauszug|opening balance|closing balance|account number",
|
||||
),
|
||||
ClassificationRule(
|
||||
"builtin_bank_kommunikationsart", "bank_statement", RULE_TYPE_METADATA, "kommunikationsart=Kontoauszug"
|
||||
),
|
||||
# ── Tax document ──────────────────────────────────────────────────────
|
||||
ClassificationRule("builtin_tax_filename", "tax_document", RULE_TYPE_FILENAME, r"(?i)tax|steuer|steuerbescheid"),
|
||||
ClassificationRule(
|
||||
"builtin_tax_content",
|
||||
"tax_document",
|
||||
RULE_TYPE_CONTENT,
|
||||
"tax return|steuerbescheid|taxable income|finanzamt|tax assessment",
|
||||
),
|
||||
# ── Insurance ─────────────────────────────────────────────────────────
|
||||
ClassificationRule(
|
||||
"builtin_insurance_filename", "insurance", RULE_TYPE_FILENAME, r"(?i)insurance|versicherung|police"
|
||||
),
|
||||
ClassificationRule(
|
||||
"builtin_insurance_content",
|
||||
"insurance",
|
||||
RULE_TYPE_CONTENT,
|
||||
"insurance policy|versicherung|policennummer|coverage|premium|deductible",
|
||||
),
|
||||
# ── Payslip ───────────────────────────────────────────────────────────
|
||||
ClassificationRule(
|
||||
"builtin_payslip_filename", "payslip", RULE_TYPE_FILENAME, r"(?i)payslip|gehaltsabrechnung|lohnabrechnung"
|
||||
),
|
||||
ClassificationRule(
|
||||
"builtin_payslip_content",
|
||||
"payslip",
|
||||
RULE_TYPE_CONTENT,
|
||||
"gross salary|net salary|gehaltsabrechnung|lohnabrechnung|bruttolohn|nettolohn",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Confidence scoring
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Base confidence for each rule type when it matches.
|
||||
_CONFIDENCE_MAP: dict[str, int] = {
|
||||
RULE_TYPE_FILENAME: 60,
|
||||
RULE_TYPE_CONTENT: 70,
|
||||
RULE_TYPE_METADATA: 90,
|
||||
}
|
||||
|
||||
#: Extra confidence per additional matching rule of the same category (capped).
|
||||
_CONFIDENCE_BONUS_PER_EXTRA_RULE = 10
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Matching helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _match_filename(rule: ClassificationRule, filename: str) -> bool:
|
||||
"""Return True if *rule.pattern* (regex) matches anywhere in *filename*."""
|
||||
if not filename:
|
||||
return False
|
||||
flags = 0 if rule.case_sensitive else re.IGNORECASE
|
||||
return bool(re.search(rule.pattern, filename, flags))
|
||||
|
||||
|
||||
def _match_content(rule: ClassificationRule, text: str) -> bool:
|
||||
"""Return True if any keyword in *rule.pattern* appears in *text*.
|
||||
|
||||
Keywords are separated by ``|`` (pipe).
|
||||
"""
|
||||
if not text:
|
||||
return False
|
||||
keywords = [kw.strip() for kw in rule.pattern.split("|") if kw.strip()]
|
||||
text_lower = text if rule.case_sensitive else text.lower()
|
||||
return any((kw if rule.case_sensitive else kw.lower()) in text_lower for kw in keywords)
|
||||
|
||||
|
||||
def _match_metadata(rule: ClassificationRule, metadata: dict[str, Any] | None) -> bool:
|
||||
"""Return True if *rule.pattern* (``field=value``) matches *metadata*.
|
||||
|
||||
Pattern format: ``field_name=expected_value``.
|
||||
"""
|
||||
if not metadata:
|
||||
return False
|
||||
if "=" not in rule.pattern:
|
||||
return False
|
||||
field_name, expected_value = rule.pattern.split("=", 1)
|
||||
actual = metadata.get(field_name.strip())
|
||||
if actual is None:
|
||||
return False
|
||||
if rule.case_sensitive:
|
||||
return str(actual) == expected_value.strip()
|
||||
return str(actual).lower() == expected_value.strip().lower()
|
||||
|
||||
|
||||
_MATCHERS: dict[str, tuple] = {
|
||||
RULE_TYPE_FILENAME: (_match_filename, "filename"),
|
||||
RULE_TYPE_CONTENT: (_match_content, "text"),
|
||||
RULE_TYPE_METADATA: (_match_metadata, "metadata"),
|
||||
}
|
||||
|
||||
|
||||
def _evaluate_rule(
|
||||
rule: ClassificationRule,
|
||||
filename: str,
|
||||
text: str,
|
||||
metadata: dict[str, Any] | None,
|
||||
) -> MatchedRule | None:
|
||||
"""Evaluate a single rule against the document. Return a :class:`MatchedRule` on match."""
|
||||
entry = _MATCHERS.get(rule.rule_type)
|
||||
if entry is None:
|
||||
return None
|
||||
|
||||
matcher, arg_key = entry
|
||||
arg_map = {"filename": filename, "text": text, "metadata": metadata}
|
||||
matched = matcher(rule, arg_map[arg_key])
|
||||
|
||||
if matched:
|
||||
return MatchedRule(
|
||||
rule_name=rule.name,
|
||||
rule_type=rule.rule_type,
|
||||
category=rule.category,
|
||||
confidence=_CONFIDENCE_MAP.get(rule.rule_type, 50),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def classify_document(
|
||||
filename: str = "",
|
||||
text: str = "",
|
||||
metadata: dict[str, Any] | None = None,
|
||||
custom_rules: list[ClassificationRule] | None = None,
|
||||
) -> ClassificationResult:
|
||||
"""Classify a document by evaluating built-in and custom rules.
|
||||
|
||||
Rules are evaluated in priority order (highest first, then built-in before
|
||||
custom for the same priority). The category with the most rule matches
|
||||
wins; ties are broken by cumulative confidence.
|
||||
|
||||
Args:
|
||||
filename: Original filename of the document.
|
||||
text: Extracted / OCR text of the document.
|
||||
metadata: Previously-extracted AI metadata dict (e.g. from ``ai_metadata``).
|
||||
custom_rules: Optional list of user-defined :class:`ClassificationRule` objects.
|
||||
|
||||
Returns:
|
||||
A :class:`ClassificationResult` with the best matching category,
|
||||
overall confidence score, and the list of rules that fired.
|
||||
"""
|
||||
all_rules = list(BUILTIN_RULES)
|
||||
if custom_rules:
|
||||
all_rules.extend(custom_rules)
|
||||
|
||||
# Sort by priority descending (higher priority first)
|
||||
all_rules.sort(key=lambda r: r.priority, reverse=True)
|
||||
|
||||
matches: list[MatchedRule] = []
|
||||
for rule in all_rules:
|
||||
result = _evaluate_rule(rule, filename, text, metadata)
|
||||
if result is not None:
|
||||
matches.append(result)
|
||||
|
||||
if not matches:
|
||||
return ClassificationResult(category="unknown", confidence=0, matched_rules=[])
|
||||
|
||||
# Aggregate by category: pick the one with the most matches, then highest
|
||||
# cumulative confidence as tiebreaker.
|
||||
category_scores: dict[str, list[MatchedRule]] = {}
|
||||
for m in matches:
|
||||
category_scores.setdefault(m.category, []).append(m)
|
||||
|
||||
best_category = max(
|
||||
category_scores,
|
||||
key=lambda cat: (len(category_scores[cat]), sum(m.confidence for m in category_scores[cat])),
|
||||
)
|
||||
|
||||
best_matches = category_scores[best_category]
|
||||
base_confidence = max(m.confidence for m in best_matches)
|
||||
bonus = min(
|
||||
(len(best_matches) - 1) * _CONFIDENCE_BONUS_PER_EXTRA_RULE,
|
||||
100 - base_confidence,
|
||||
)
|
||||
final_confidence = min(base_confidence + bonus, 100)
|
||||
|
||||
return ClassificationResult(
|
||||
category=best_category,
|
||||
confidence=final_confidence,
|
||||
matched_rules=best_matches,
|
||||
)
|
||||
|
||||
|
||||
def db_rule_to_engine_rule(db_rule: Any) -> ClassificationRule:
|
||||
"""Convert a database ``ClassificationRuleModel`` row to an engine :class:`ClassificationRule`.
|
||||
|
||||
Args:
|
||||
db_rule: A SQLAlchemy model instance with ``name``, ``category``,
|
||||
``rule_type``, ``pattern``, ``priority``, and ``case_sensitive`` attributes.
|
||||
|
||||
Returns:
|
||||
A :class:`ClassificationRule` dataclass instance.
|
||||
"""
|
||||
return ClassificationRule(
|
||||
name=db_rule.name,
|
||||
category=db_rule.category,
|
||||
rule_type=db_rule.rule_type,
|
||||
pattern=db_rule.pattern,
|
||||
priority=db_rule.priority,
|
||||
case_sensitive=getattr(db_rule, "case_sensitive", False),
|
||||
)
|
||||
@@ -296,6 +296,28 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
},
|
||||
}
|
||||
|
||||
# Check SharePoint configuration
|
||||
providers["SharePoint"] = {
|
||||
"name": "SharePoint",
|
||||
"icon": "fa-brands fa-microsoft",
|
||||
"configured": bool(
|
||||
getattr(settings, "sharepoint_client_id", None)
|
||||
and getattr(settings, "sharepoint_client_secret", None)
|
||||
and getattr(settings, "sharepoint_site_url", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"description": "Store documents in Microsoft SharePoint Online",
|
||||
"details": {
|
||||
"client_id": getattr(settings, "sharepoint_client_id", "Not set"),
|
||||
"client_secret": mask_sensitive_value(getattr(settings, "sharepoint_client_secret", None)),
|
||||
"tenant_id": getattr(settings, "sharepoint_tenant_id", "Not set"),
|
||||
"refresh_token": mask_sensitive_value(getattr(settings, "sharepoint_refresh_token", None)),
|
||||
"site_url": getattr(settings, "sharepoint_site_url", "Not set"),
|
||||
"document_library": getattr(settings, "sharepoint_document_library", "Not set"),
|
||||
"folder_path": getattr(settings, "sharepoint_folder_path", "Not set"),
|
||||
},
|
||||
}
|
||||
|
||||
# Check S3 configuration
|
||||
providers["S3 Storage"] = {
|
||||
"name": "S3 Storage",
|
||||
|
||||
@@ -0,0 +1,505 @@
|
||||
"""Server-side session management utilities.
|
||||
|
||||
Provides helpers for creating, validating, and revoking user sessions.
|
||||
Sessions are tracked in the ``user_sessions`` table and referenced by a
|
||||
cryptographically random token stored in the browser cookie. This enables
|
||||
the "log off everywhere" feature and per-session revocation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.models import ApiToken, QRLoginChallenge, UserSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _ensure_tz_aware(dt: datetime | None) -> datetime | None:
|
||||
"""Return *dt* with UTC tzinfo if it is naive, or unchanged if already aware.
|
||||
|
||||
SQLite does not persist timezone information, so datetimes read back from
|
||||
the database are offset-naive. This helper normalises them for safe
|
||||
comparison with ``datetime.now(timezone.utc)``.
|
||||
"""
|
||||
if dt is not None and dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=timezone.utc)
|
||||
return dt
|
||||
|
||||
|
||||
def get_session_lifetime_days() -> int:
|
||||
"""Return the effective session lifetime in days.
|
||||
|
||||
If ``session_lifetime_custom_days`` is set it takes precedence over
|
||||
``session_lifetime_days``.
|
||||
"""
|
||||
custom = getattr(settings, "session_lifetime_custom_days", None)
|
||||
if custom is not None and isinstance(custom, int) and custom > 0:
|
||||
return custom
|
||||
return max(1, getattr(settings, "session_lifetime_days", 30))
|
||||
|
||||
|
||||
def get_session_max_age_seconds() -> int:
|
||||
"""Return the session max-age in seconds for the cookie."""
|
||||
return get_session_lifetime_days() * 86400
|
||||
|
||||
|
||||
def create_session(
|
||||
db: Session,
|
||||
user_id: str,
|
||||
ip_address: str | None = None,
|
||||
user_agent: str | None = None,
|
||||
) -> UserSession:
|
||||
"""Create a new server-side session record.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
user_id: Stable owner identifier.
|
||||
ip_address: Client IP address.
|
||||
user_agent: Client User-Agent header.
|
||||
|
||||
Returns:
|
||||
The newly created ``UserSession`` instance.
|
||||
"""
|
||||
session_token = secrets.token_urlsafe(64)
|
||||
now = datetime.now(timezone.utc)
|
||||
lifetime_days = get_session_lifetime_days()
|
||||
expires_at = now + timedelta(days=lifetime_days)
|
||||
|
||||
device_info = _parse_device_info(user_agent)
|
||||
|
||||
user_session = UserSession(
|
||||
session_token=session_token,
|
||||
user_id=user_id,
|
||||
ip_address=ip_address,
|
||||
user_agent=(user_agent or "")[:512],
|
||||
device_info=device_info,
|
||||
created_at=now,
|
||||
last_active_at=now,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
try:
|
||||
db.add(user_session)
|
||||
db.commit()
|
||||
db.refresh(user_session)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create session for user_id=%s", user_id)
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"[SESSION] Created session id=%s user=%s device=%r expires=%s",
|
||||
user_session.id,
|
||||
user_id,
|
||||
device_info,
|
||||
expires_at.isoformat(),
|
||||
)
|
||||
return user_session
|
||||
|
||||
|
||||
def validate_session(db: Session, session_token: str) -> UserSession | None:
|
||||
"""Validate a session token and return the session if valid.
|
||||
|
||||
A session is valid when:
|
||||
* It exists in the database.
|
||||
* ``is_revoked`` is ``False``.
|
||||
* ``expires_at`` is in the future.
|
||||
|
||||
Side-effect: updates ``last_active_at`` on valid sessions.
|
||||
|
||||
Returns:
|
||||
The ``UserSession`` if valid, else ``None``.
|
||||
"""
|
||||
if not session_token:
|
||||
return None
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
user_session = db.query(UserSession).filter(UserSession.session_token == session_token).first()
|
||||
|
||||
if not user_session:
|
||||
logger.debug("[SESSION] Token not found in database")
|
||||
return None
|
||||
|
||||
if user_session.is_revoked:
|
||||
logger.debug("[SESSION] Session id=%s is revoked", user_session.id)
|
||||
return None
|
||||
|
||||
if user_session.expires_at:
|
||||
expires = _ensure_tz_aware(user_session.expires_at)
|
||||
if expires < now:
|
||||
logger.debug("[SESSION] Session id=%s has expired", user_session.id)
|
||||
return None
|
||||
|
||||
# Update last_active_at (throttled to avoid excessive writes)
|
||||
last_active = _ensure_tz_aware(user_session.last_active_at)
|
||||
if not last_active or (now - last_active).total_seconds() > 60:
|
||||
try:
|
||||
user_session.last_active_at = now
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.debug("[SESSION] Failed to update last_active_at for session id=%s", user_session.id)
|
||||
|
||||
return user_session
|
||||
|
||||
|
||||
def revoke_session(db: Session, session_id: int, user_id: str) -> bool:
|
||||
"""Revoke a single session by ID.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
session_id: The session record ID to revoke.
|
||||
user_id: The owner — ensures a user can only revoke their own sessions.
|
||||
|
||||
Returns:
|
||||
``True`` if the session was found and revoked, ``False`` otherwise.
|
||||
"""
|
||||
user_session = db.get(UserSession, session_id)
|
||||
if not user_session or user_session.user_id != user_id:
|
||||
return False
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
user_session.is_revoked = True
|
||||
user_session.revoked_at = now
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("[SESSION] Revoked session id=%s user=%s", session_id, user_id)
|
||||
return True
|
||||
|
||||
|
||||
def revoke_all_sessions(
|
||||
db: Session,
|
||||
user_id: str,
|
||||
*,
|
||||
except_session_id: int | None = None,
|
||||
revoke_api_tokens: bool = True,
|
||||
) -> int:
|
||||
"""Revoke all active sessions for a user ("log off everywhere").
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
user_id: The owner whose sessions should be revoked.
|
||||
except_session_id: If provided, keep this session active (the
|
||||
current browser session).
|
||||
revoke_api_tokens: If ``True``, also revoke all active API tokens.
|
||||
|
||||
Returns:
|
||||
Number of sessions revoked.
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
query = db.query(UserSession).filter(
|
||||
UserSession.user_id == user_id,
|
||||
UserSession.is_revoked.is_(False),
|
||||
)
|
||||
if except_session_id is not None:
|
||||
query = query.filter(UserSession.id != except_session_id)
|
||||
|
||||
sessions = query.all()
|
||||
count = 0
|
||||
for s in sessions:
|
||||
s.is_revoked = True
|
||||
s.revoked_at = now
|
||||
count += 1
|
||||
|
||||
if revoke_api_tokens:
|
||||
tokens = (
|
||||
db.query(ApiToken)
|
||||
.filter(
|
||||
ApiToken.owner_id == user_id,
|
||||
ApiToken.is_active.is_(True),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for t in tokens:
|
||||
t.is_active = False
|
||||
t.revoked_at = now
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"[SESSION] Revoked all sessions for user=%s (count=%d, except_session_id=%s, tokens_revoked=%s)",
|
||||
user_id,
|
||||
count,
|
||||
except_session_id,
|
||||
revoke_api_tokens,
|
||||
)
|
||||
return count
|
||||
|
||||
|
||||
def list_user_sessions(db: Session, user_id: str) -> list[UserSession]:
|
||||
"""Return all non-revoked, non-expired sessions for a user.
|
||||
|
||||
Results are ordered by most recently active first.
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
sessions = (
|
||||
db.query(UserSession)
|
||||
.filter(
|
||||
UserSession.user_id == user_id,
|
||||
UserSession.is_revoked.is_(False),
|
||||
)
|
||||
.order_by(UserSession.last_active_at.desc())
|
||||
.all()
|
||||
)
|
||||
# Filter expired sessions in Python to handle timezone-naive datetimes (SQLite)
|
||||
result = []
|
||||
for s in sessions:
|
||||
expires = _ensure_tz_aware(s.expires_at)
|
||||
if expires and expires > now:
|
||||
result.append(s)
|
||||
return result
|
||||
|
||||
|
||||
def cleanup_expired_sessions(db: Session) -> int:
|
||||
"""Delete sessions that expired more than 7 days ago.
|
||||
|
||||
Intended to be called periodically (e.g. via Celery beat) to keep the
|
||||
table from growing unbounded.
|
||||
|
||||
Returns:
|
||||
Number of rows deleted.
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=7)
|
||||
count = db.query(UserSession).filter(UserSession.expires_at < cutoff).delete(synchronize_session=False)
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
if count:
|
||||
logger.info("[SESSION] Cleaned up %d expired sessions", count)
|
||||
return count
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# QR login helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def create_qr_challenge(db: Session, user_id: str, ip_address: str | None = None) -> QRLoginChallenge:
|
||||
"""Create a new QR login challenge.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
user_id: The authenticated web user creating the challenge.
|
||||
ip_address: IP address of the web client.
|
||||
|
||||
Returns:
|
||||
The newly created ``QRLoginChallenge``.
|
||||
"""
|
||||
token = secrets.token_urlsafe(64)
|
||||
ttl = getattr(settings, "qr_login_challenge_ttl_seconds", 120)
|
||||
now = datetime.now(timezone.utc)
|
||||
expires_at = now + timedelta(seconds=ttl)
|
||||
|
||||
challenge = QRLoginChallenge(
|
||||
challenge_token=token,
|
||||
user_id=user_id,
|
||||
created_by_ip=ip_address,
|
||||
created_at=now,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
try:
|
||||
db.add(challenge)
|
||||
db.commit()
|
||||
db.refresh(challenge)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create QR login challenge for user_id=%s", user_id)
|
||||
raise
|
||||
|
||||
logger.info("[QR_AUTH] Challenge created: id=%s user=%s expires=%s", challenge.id, user_id, expires_at.isoformat())
|
||||
return challenge
|
||||
|
||||
|
||||
def validate_qr_challenge(db: Session, challenge_token: str) -> QRLoginChallenge | None:
|
||||
"""Validate a QR challenge token without claiming it.
|
||||
|
||||
Returns the challenge if it exists, is not expired, not claimed,
|
||||
and not cancelled. Returns ``None`` otherwise.
|
||||
"""
|
||||
if not challenge_token:
|
||||
return None
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
challenge = db.query(QRLoginChallenge).filter(QRLoginChallenge.challenge_token == challenge_token).first()
|
||||
|
||||
if not challenge:
|
||||
return None
|
||||
if challenge.is_claimed or challenge.is_cancelled:
|
||||
return None
|
||||
|
||||
expires = _ensure_tz_aware(challenge.expires_at)
|
||||
if expires and expires < now:
|
||||
return None
|
||||
|
||||
return challenge
|
||||
|
||||
|
||||
def claim_qr_challenge(
|
||||
db: Session,
|
||||
challenge_token: str,
|
||||
device_name: str = "Mobile App",
|
||||
ip_address: str | None = None,
|
||||
) -> dict | None:
|
||||
"""Claim a QR challenge and issue an API token.
|
||||
|
||||
This is the critical security path. The challenge is validated,
|
||||
marked as claimed atomically, and an API token is issued for the
|
||||
user who created the challenge.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
challenge_token: The token from the QR code.
|
||||
device_name: Name provided by the mobile app.
|
||||
ip_address: IP address of the claiming mobile device.
|
||||
|
||||
Returns:
|
||||
Dict with ``token`` (plaintext), ``token_id``, ``name``, ``owner_id``
|
||||
and ``created_at`` on success, or ``None`` if the challenge is invalid.
|
||||
"""
|
||||
from app.api.api_tokens import generate_api_token, hash_token
|
||||
|
||||
challenge = validate_qr_challenge(db, challenge_token)
|
||||
if not challenge:
|
||||
logger.warning("[QR_AUTH] Invalid or expired challenge token attempted")
|
||||
return None
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# Mark as claimed first to prevent race conditions
|
||||
challenge.is_claimed = True
|
||||
challenge.claimed_at = now
|
||||
challenge.claimed_by_ip = ip_address
|
||||
challenge.device_name = device_name
|
||||
|
||||
# Generate API token for the mobile app
|
||||
token_name = f"Mobile App (QR) – {device_name}"
|
||||
plaintext = generate_api_token()
|
||||
token_hash_value = hash_token(plaintext)
|
||||
prefix = plaintext[:12]
|
||||
|
||||
db_token = ApiToken(
|
||||
owner_id=challenge.user_id,
|
||||
name=token_name,
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
)
|
||||
|
||||
try:
|
||||
db.add(db_token)
|
||||
db.flush()
|
||||
challenge.issued_token_id = db_token.id
|
||||
db.commit()
|
||||
db.refresh(db_token)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("[QR_AUTH] Failed to issue token for challenge id=%s", challenge.id)
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"[QR_AUTH] Challenge claimed: id=%s user=%s device=%r token_id=%s",
|
||||
challenge.id,
|
||||
challenge.user_id,
|
||||
device_name,
|
||||
db_token.id,
|
||||
)
|
||||
|
||||
return {
|
||||
"token": plaintext,
|
||||
"token_id": db_token.id,
|
||||
"name": token_name,
|
||||
"owner_id": challenge.user_id,
|
||||
"created_at": db_token.created_at,
|
||||
}
|
||||
|
||||
|
||||
def get_challenge_status(db: Session, challenge_id: int, user_id: str) -> dict | None:
|
||||
"""Get the current status of a QR challenge (for polling from the web UI).
|
||||
|
||||
Returns:
|
||||
Dict with ``status`` ("pending", "claimed", "expired", "cancelled")
|
||||
and metadata, or ``None`` if the challenge doesn't belong to the user.
|
||||
"""
|
||||
challenge = db.get(QRLoginChallenge, challenge_id)
|
||||
if not challenge or challenge.user_id != user_id:
|
||||
return None
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
expires = _ensure_tz_aware(challenge.expires_at)
|
||||
|
||||
if challenge.is_claimed:
|
||||
status = "claimed"
|
||||
elif challenge.is_cancelled:
|
||||
status = "cancelled"
|
||||
elif expires and expires < now:
|
||||
status = "expired"
|
||||
else:
|
||||
status = "pending"
|
||||
|
||||
return {
|
||||
"id": challenge.id,
|
||||
"status": status,
|
||||
"device_name": challenge.device_name,
|
||||
"claimed_at": challenge.claimed_at,
|
||||
"expires_at": challenge.expires_at,
|
||||
}
|
||||
|
||||
|
||||
def _parse_device_info(user_agent: str | None) -> str | None:
|
||||
"""Extract a human-readable device description from User-Agent.
|
||||
|
||||
This is a lightweight parser — not a full UA library — that covers
|
||||
the most common browsers and platforms.
|
||||
"""
|
||||
if not user_agent:
|
||||
return None
|
||||
|
||||
ua = user_agent.lower()
|
||||
|
||||
# Platform detection
|
||||
platform = "Unknown"
|
||||
if "iphone" in ua:
|
||||
platform = "iPhone"
|
||||
elif "ipad" in ua:
|
||||
platform = "iPad"
|
||||
elif "android" in ua:
|
||||
platform = "Android"
|
||||
elif "macintosh" in ua or "mac os" in ua:
|
||||
platform = "macOS"
|
||||
elif "windows" in ua:
|
||||
platform = "Windows"
|
||||
elif "linux" in ua:
|
||||
platform = "Linux"
|
||||
elif "cros" in ua:
|
||||
platform = "ChromeOS"
|
||||
|
||||
# Browser detection
|
||||
browser = "Unknown Browser"
|
||||
if "edg/" in ua or "edge/" in ua:
|
||||
browser = "Edge"
|
||||
elif "opr/" in ua or "opera" in ua:
|
||||
browser = "Opera"
|
||||
elif "chrome/" in ua and "safari/" in ua:
|
||||
browser = "Chrome"
|
||||
elif "safari/" in ua and "chrome/" not in ua:
|
||||
browser = "Safari"
|
||||
elif "firefox/" in ua:
|
||||
browser = "Firefox"
|
||||
elif "docuelevate" in ua:
|
||||
browser = "DocuElevate App"
|
||||
|
||||
return f"{browser} on {platform}"
|
||||
@@ -39,6 +39,50 @@ SETTING_METADATA = {
|
||||
"required": True,
|
||||
"restart_required": True,
|
||||
},
|
||||
"db_pool_size": {
|
||||
"category": "Core",
|
||||
"description": (
|
||||
"Number of persistent database connections kept in the pool per worker process. "
|
||||
"Ignored for SQLite (which uses NullPool). Default: 10."
|
||||
),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"db_max_overflow": {
|
||||
"category": "Core",
|
||||
"description": (
|
||||
"Additional database connections allowed beyond db_pool_size under burst load. "
|
||||
"Ignored for SQLite. Default: 20."
|
||||
),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"db_pool_timeout": {
|
||||
"category": "Core",
|
||||
"description": (
|
||||
"Seconds to wait for a database connection from the pool before raising a TimeoutError. "
|
||||
"Ignored for SQLite. Default: 30."
|
||||
),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"db_pool_recycle": {
|
||||
"category": "Core",
|
||||
"description": (
|
||||
"Recycle (close and reopen) database connections after this many seconds "
|
||||
"to avoid stale connections. Ignored for SQLite. Default: 1800."
|
||||
),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"workdir": {
|
||||
"category": "Core",
|
||||
"description": "Working directory for file storage and processing",
|
||||
@@ -55,6 +99,18 @@ SETTING_METADATA = {
|
||||
"required": True, # Required for OAuth redirects and external URLs
|
||||
"restart_required": True,
|
||||
},
|
||||
"public_base_url": {
|
||||
"category": "Core",
|
||||
"description": (
|
||||
"Full public base URL including scheme (e.g., https://docuelevate.example.com). "
|
||||
"When set, overrides auto-detected URLs for OAuth redirect URIs. "
|
||||
"Required when behind a reverse proxy that does not forward X-Forwarded-Proto."
|
||||
),
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"debug": {
|
||||
"category": "Core",
|
||||
"description": "Enable debug mode for verbose logging",
|
||||
@@ -134,6 +190,30 @@ SETTING_METADATA = {
|
||||
"required": True, # Required when auth_enabled=True (validated in config.py)
|
||||
"restart_required": True,
|
||||
},
|
||||
"session_lifetime_days": {
|
||||
"category": "Authentication",
|
||||
"description": "Session lifetime in days (default 30). Determines how long a user stays logged in.",
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"session_lifetime_custom_days": {
|
||||
"category": "Authentication",
|
||||
"description": "Override session_lifetime_days with a custom value. Takes precedence when set.",
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"qr_login_challenge_ttl_seconds": {
|
||||
"category": "Authentication",
|
||||
"description": "Time-to-live in seconds for QR login challenges (default 120).",
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"admin_username": {
|
||||
"category": "Authentication",
|
||||
"description": "Admin username for local authentication",
|
||||
@@ -302,6 +382,19 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"social_auth_dropbox_use_global_credentials": {
|
||||
"category": "Social Login",
|
||||
"description": (
|
||||
"When True, Dropbox social login uses the global DROPBOX_APP_KEY / DROPBOX_APP_SECRET "
|
||||
"credentials instead of requiring separate SOCIAL_AUTH_DROPBOX_CLIENT_ID / "
|
||||
"SOCIAL_AUTH_DROPBOX_CLIENT_SECRET values. "
|
||||
"Requires SOCIAL_AUTH_DROPBOX_ENABLED=True and global Dropbox credentials to be set."
|
||||
),
|
||||
"type": "boolean",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"social_auth_dropbox_enabled": {
|
||||
"category": "Social Login",
|
||||
"description": (
|
||||
@@ -695,6 +788,18 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"dropbox_allow_global_credentials_for_integrations": {
|
||||
"category": "Storage Providers",
|
||||
"description": (
|
||||
"When True, users may authorize their personal Dropbox integrations using the global "
|
||||
"DROPBOX_APP_KEY / DROPBOX_APP_SECRET credentials configured by the admin, without "
|
||||
"needing to create their own Dropbox app."
|
||||
),
|
||||
"type": "boolean",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Storage Providers - Nextcloud
|
||||
"nextcloud_enabled": {
|
||||
"category": "Storage Providers",
|
||||
@@ -875,6 +980,63 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Storage Providers - SharePoint
|
||||
"sharepoint_client_id": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint Azure AD application (client) ID",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_client_secret": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint Azure AD client secret",
|
||||
"type": "string",
|
||||
"sensitive": True,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_tenant_id": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint Azure AD tenant ID (use 'common' for multi-tenant apps)",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_refresh_token": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint OAuth refresh token",
|
||||
"type": "string",
|
||||
"sensitive": True,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_site_url": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint site URL (e.g. https://tenant.sharepoint.com/sites/sitename)",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_document_library": {
|
||||
"category": "Storage Providers",
|
||||
"description": "SharePoint document library name (default: 'Documents')",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"sharepoint_folder_path": {
|
||||
"category": "Storage Providers",
|
||||
"description": "Subfolder path inside the SharePoint document library",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Storage Providers - WebDAV
|
||||
"webdav_enabled": {
|
||||
"category": "Storage Providers",
|
||||
@@ -1912,6 +2074,28 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"factory_reset_on_startup": {
|
||||
"category": "Feature Flags",
|
||||
"description": (
|
||||
"Wipe all user data on every startup so the instance always starts fresh. "
|
||||
"Useful for demo/testing environments. Default: False."
|
||||
),
|
||||
"type": "boolean",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"enable_factory_reset": {
|
||||
"category": "Feature Flags",
|
||||
"description": (
|
||||
"Show the System Reset page in the admin UI. Allows administrators to "
|
||||
"trigger a full data wipe or a wipe-and-reimport from the web interface. Default: False."
|
||||
),
|
||||
"type": "boolean",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Backup / Restore
|
||||
"backup_enabled": {
|
||||
"category": "Backup",
|
||||
@@ -1935,14 +2119,26 @@ SETTING_METADATA = {
|
||||
"category": "Backup",
|
||||
"description": (
|
||||
"Storage provider for remote backup copies. "
|
||||
"Accepted values: s3, dropbox, google_drive, onedrive, nextcloud, webdav, ftp, sftp, email. "
|
||||
"Accepted values: s3, dropbox, google_drive, onedrive, sharepoint, nextcloud, webdav, ftp, sftp, email. "
|
||||
"Leave empty to keep backups local only."
|
||||
),
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
"options": ["", "s3", "dropbox", "google_drive", "onedrive", "nextcloud", "webdav", "ftp", "sftp", "email"],
|
||||
"options": [
|
||||
"",
|
||||
"s3",
|
||||
"dropbox",
|
||||
"google_drive",
|
||||
"onedrive",
|
||||
"sharepoint",
|
||||
"nextcloud",
|
||||
"webdav",
|
||||
"ftp",
|
||||
"sftp",
|
||||
"email",
|
||||
],
|
||||
},
|
||||
"backup_remote_folder": {
|
||||
"category": "Backup",
|
||||
@@ -2454,6 +2650,27 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Per-user upload rate limiting
|
||||
"upload_rate_limit_per_user": {
|
||||
"category": "Security",
|
||||
"description": (
|
||||
"Maximum number of uploads a single user may submit within upload_rate_limit_window seconds. "
|
||||
"The health-aware limiter may reduce this dynamically under high Redis queue depth or CPU load. "
|
||||
"Default: 20."
|
||||
),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
"upload_rate_limit_window": {
|
||||
"category": "Security",
|
||||
"description": ("Sliding window in seconds over which upload_rate_limit_per_user is enforced. Default: 60."),
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Rate Limiting
|
||||
"rate_limiting_enabled": {
|
||||
"category": "Security",
|
||||
|
||||
@@ -0,0 +1,296 @@
|
||||
"""
|
||||
System reset utilities for DocuElevate.
|
||||
|
||||
Provides functions to:
|
||||
- Wipe all user data (database rows + work-files on disk) for a fresh start.
|
||||
- Wipe with re-import: move original files to a dedicated folder, wipe
|
||||
everything, then let the watch-folder mechanism re-ingest the files.
|
||||
|
||||
Security: All public functions in this module require admin-level access.
|
||||
They MUST only be invoked from admin-guarded API/view endpoints.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Subdirectories inside *workdir* that contain user-generated data.
|
||||
# Everything else (app code, static assets, config) is left untouched.
|
||||
_USER_DATA_SUBDIRS = ("original", "processed", "tmp", "pdfa", "backups")
|
||||
|
||||
# JSON cache files written by watch-folder / ingest tasks.
|
||||
_CACHE_FILES = (
|
||||
"watch_folder_processed.json",
|
||||
"ftp_ingest_processed.json",
|
||||
"sftp_ingest_processed.json",
|
||||
"dropbox_ingest_processed.json",
|
||||
"gdrive_ingest_processed.json",
|
||||
"onedrive_ingest_processed.json",
|
||||
"nextcloud_ingest_processed.json",
|
||||
"s3_ingest_processed.json",
|
||||
"webdav_ingest_processed.json",
|
||||
"processed_mails.json",
|
||||
"credential_failures.json",
|
||||
)
|
||||
|
||||
# The folder name used for storing files prior to re-import.
|
||||
REIMPORT_FOLDER_NAME = "reimport"
|
||||
|
||||
|
||||
def _wipe_workdir_data(workdir: str) -> dict[str, int]:
|
||||
"""Delete user data subdirectories and cache files inside *workdir*.
|
||||
|
||||
Leaves the workdir directory itself intact so the application can
|
||||
continue to write into it. Also leaves any files that do not belong
|
||||
to the known data subdirectories or caches.
|
||||
|
||||
Returns:
|
||||
A dict with counts of deleted directories and files.
|
||||
"""
|
||||
workdir_path = Path(workdir)
|
||||
deleted_dirs = 0
|
||||
deleted_files = 0
|
||||
|
||||
# Remove data subdirectories
|
||||
for subdir in _USER_DATA_SUBDIRS:
|
||||
target = workdir_path / subdir
|
||||
if target.is_dir():
|
||||
shutil.rmtree(target)
|
||||
logger.info("Deleted data directory: %s", target)
|
||||
deleted_dirs += 1
|
||||
|
||||
# Remove cache / state JSON files
|
||||
for cache_file in _CACHE_FILES:
|
||||
target = workdir_path / cache_file
|
||||
if target.is_file():
|
||||
target.unlink()
|
||||
logger.info("Deleted cache file: %s", target)
|
||||
deleted_files += 1
|
||||
|
||||
# Also remove user_wf_*.json files (per-user watch folder caches)
|
||||
for f in workdir_path.glob("user_wf_*.json"):
|
||||
f.unlink()
|
||||
logger.info("Deleted user watch-folder cache: %s", f)
|
||||
deleted_files += 1
|
||||
|
||||
# Remove loose files in workdir root that are user uploads (uuid-named
|
||||
# files like "a1b2c3d4-…pdf") but NOT application config files.
|
||||
for entry in workdir_path.iterdir():
|
||||
if entry.is_file() and entry.suffix.lower() in {
|
||||
".pdf",
|
||||
".png",
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".tiff",
|
||||
".tif",
|
||||
".docx",
|
||||
".doc",
|
||||
".xlsx",
|
||||
".xls",
|
||||
".pptx",
|
||||
".heic",
|
||||
".heif",
|
||||
".webp",
|
||||
".bmp",
|
||||
".gif",
|
||||
".txt",
|
||||
".rtf",
|
||||
".odt",
|
||||
".ods",
|
||||
".odp",
|
||||
".csv",
|
||||
".pages",
|
||||
".numbers",
|
||||
".keynote",
|
||||
}:
|
||||
entry.unlink()
|
||||
logger.info("Deleted loose workdir file: %s", entry)
|
||||
deleted_files += 1
|
||||
|
||||
return {"deleted_dirs": deleted_dirs, "deleted_files": deleted_files}
|
||||
|
||||
|
||||
def _wipe_database(db: Session) -> dict[str, int]:
|
||||
"""Delete all user-generated rows from the database.
|
||||
|
||||
Preserves schema (tables, migrations) and system-seeded rows that will
|
||||
be re-created on the next startup (subscription plans, default pipeline,
|
||||
scheduled jobs, compliance templates).
|
||||
|
||||
Returns:
|
||||
A dict mapping table name → number of rows deleted.
|
||||
"""
|
||||
from app.models import (
|
||||
AuditLog,
|
||||
BackupRecord,
|
||||
DocumentMetadata,
|
||||
FileProcessingStep,
|
||||
FileRecord,
|
||||
InAppNotification,
|
||||
ProcessingLog,
|
||||
SavedSearch,
|
||||
SettingsAuditLog,
|
||||
SharedLink,
|
||||
UserImapAccount,
|
||||
UserIntegration,
|
||||
UserNotificationPreference,
|
||||
UserNotificationTarget,
|
||||
)
|
||||
|
||||
# Order matters: delete children before parents to respect FK constraints.
|
||||
tables_to_wipe: list[tuple[str, type]] = [
|
||||
("file_processing_steps", FileProcessingStep),
|
||||
("processing_logs", ProcessingLog),
|
||||
("shared_links", SharedLink),
|
||||
("in_app_notifications", InAppNotification),
|
||||
("user_notification_preferences", UserNotificationPreference),
|
||||
("user_notification_targets", UserNotificationTarget),
|
||||
("user_imap_accounts", UserImapAccount),
|
||||
("user_integrations", UserIntegration),
|
||||
("saved_searches", SavedSearch),
|
||||
("settings_audit_log", SettingsAuditLog),
|
||||
("audit_logs", AuditLog),
|
||||
("backup_records", BackupRecord),
|
||||
("document_metadata", DocumentMetadata),
|
||||
("files", FileRecord),
|
||||
]
|
||||
|
||||
result: dict[str, int] = {}
|
||||
for table_name, model in tables_to_wipe:
|
||||
try:
|
||||
count = db.query(model).delete()
|
||||
result[table_name] = count
|
||||
logger.info("Wiped %d rows from %s", count, table_name)
|
||||
except Exception:
|
||||
logger.exception("Failed to wipe table %s during system reset", table_name)
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
db.commit()
|
||||
return result
|
||||
|
||||
|
||||
def perform_full_reset(db: Session) -> dict:
|
||||
"""Perform a complete system reset: wipe database rows + work-files.
|
||||
|
||||
Args:
|
||||
db: An active SQLAlchemy session.
|
||||
|
||||
Returns:
|
||||
Summary dict with ``database`` and ``filesystem`` sub-dicts.
|
||||
"""
|
||||
logger.warning(">>> SYSTEM RESET: wiping all user data <<<")
|
||||
|
||||
db_result = _wipe_database(db)
|
||||
fs_result = _wipe_workdir_data(settings.workdir)
|
||||
|
||||
logger.warning(">>> SYSTEM RESET complete <<<")
|
||||
return {"database": db_result, "filesystem": fs_result}
|
||||
|
||||
|
||||
def perform_reset_and_reimport(db: Session) -> dict:
|
||||
"""Move original files to a reimport folder, wipe everything, then
|
||||
configure the reimport folder as a watch folder for re-ingestion.
|
||||
|
||||
The watch-folder scanner (``scan_all_watch_folders``) will pick up
|
||||
the files on its next periodic run and process them exactly as if
|
||||
they had been freshly uploaded — respecting the same backoff
|
||||
strategy, size limits, and rate limits.
|
||||
|
||||
Args:
|
||||
db: An active SQLAlchemy session.
|
||||
|
||||
Returns:
|
||||
Summary dict with ``database``, ``filesystem``, and ``reimport`` sub-dicts.
|
||||
"""
|
||||
workdir_path = Path(settings.workdir)
|
||||
reimport_dir = workdir_path / REIMPORT_FOLDER_NAME
|
||||
original_dir = workdir_path / "original"
|
||||
|
||||
# 1. Collect original files
|
||||
files_moved = 0
|
||||
reimport_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if original_dir.is_dir():
|
||||
for entry in original_dir.iterdir():
|
||||
if entry.is_file():
|
||||
# Validate the resolved path stays within original_dir (path traversal guard)
|
||||
try:
|
||||
entry.resolve().relative_to(original_dir.resolve())
|
||||
except ValueError:
|
||||
logger.warning("Skipping file outside original dir: %s", entry)
|
||||
continue
|
||||
dest = reimport_dir / entry.name
|
||||
# Avoid overwriting: append counter if name clash
|
||||
if dest.exists():
|
||||
stem = dest.stem
|
||||
suffix = dest.suffix
|
||||
counter = 1
|
||||
while dest.exists():
|
||||
dest = reimport_dir / f"{stem}_{counter}{suffix}"
|
||||
counter += 1
|
||||
shutil.copy2(str(entry), str(dest))
|
||||
files_moved += 1
|
||||
|
||||
logger.info("Copied %d original files to reimport folder: %s", files_moved, reimport_dir)
|
||||
|
||||
# 2. Perform the full reset (wipe DB + other workdir data)
|
||||
reset_result = perform_full_reset(db)
|
||||
|
||||
# 3. Ensure the reimport folder survived the wipe (it's not in _USER_DATA_SUBDIRS)
|
||||
# and set up watch folder config to point at it.
|
||||
_configure_reimport_watch_folder(str(reimport_dir))
|
||||
|
||||
reset_result["reimport"] = {
|
||||
"files_moved": files_moved,
|
||||
"reimport_folder": str(reimport_dir),
|
||||
}
|
||||
logger.warning(">>> SYSTEM RESET with re-import configured — %d files staged <<<", files_moved)
|
||||
return reset_result
|
||||
|
||||
|
||||
def _configure_reimport_watch_folder(reimport_path: str) -> None:
|
||||
"""Append *reimport_path* to the application's watch-folder list.
|
||||
|
||||
The watch-folder scanner uses ``settings.watch_folders`` (a
|
||||
comma-separated string). We mutate the runtime setting so the
|
||||
next scan picks up the folder. We also set
|
||||
``watch_folder_delete_after_process = True`` so files are cleaned
|
||||
up after successful processing.
|
||||
"""
|
||||
current = getattr(settings, "watch_folders", None) or ""
|
||||
folders = [f.strip() for f in current.split(",") if f.strip()]
|
||||
|
||||
if reimport_path not in folders:
|
||||
folders.append(reimport_path)
|
||||
|
||||
# Mutate runtime settings (not persisted to .env — ephemeral)
|
||||
object.__setattr__(settings, "watch_folders", ",".join(folders))
|
||||
object.__setattr__(settings, "watch_folder_delete_after_process", True)
|
||||
logger.info("Configured reimport watch folder: %s", reimport_path)
|
||||
|
||||
|
||||
def perform_startup_reset() -> None:
|
||||
"""Called during application startup when ``FACTORY_RESET_ON_STARTUP=True``.
|
||||
|
||||
Wipes database and filesystem data so the instance starts completely
|
||||
fresh. Uses its own DB session so it runs before the normal lifespan
|
||||
seeding logic.
|
||||
"""
|
||||
from app.database import SessionLocal
|
||||
|
||||
logger.warning("FACTORY_RESET_ON_STARTUP is enabled — wiping all data")
|
||||
db = SessionLocal()
|
||||
try:
|
||||
perform_full_reset(db)
|
||||
except Exception:
|
||||
logger.exception("Factory reset on startup failed")
|
||||
db.rollback()
|
||||
finally:
|
||||
db.close()
|
||||
@@ -10,6 +10,7 @@ from app.views.audit_logs import router as audit_logs_router
|
||||
from app.views.backup import router as backup_router
|
||||
from app.views.compliance import router as compliance_router
|
||||
from app.views.db_wizard import router as db_wizard_router
|
||||
from app.views.devices import router as devices_router # Mobile devices dashboard
|
||||
from app.views.dropbox import router as dropbox_router
|
||||
from app.views.filemanager import router as filemanager_router
|
||||
|
||||
@@ -26,6 +27,7 @@ from app.views.onedrive import router as onedrive_router
|
||||
from app.views.pipelines import router as pipelines_router # Processing pipelines
|
||||
from app.views.plans import router as plans_router # Admin Plan Designer
|
||||
from app.views.profile import router as profile_router # User self-service profile
|
||||
from app.views.qr_login import router as qr_login_router # QR code mobile login
|
||||
from app.views.queue import router as queue_router
|
||||
from app.views.scheduled_jobs import router as scheduled_jobs_router # Scheduled batch jobs
|
||||
from app.views.search import router as search_router
|
||||
@@ -34,6 +36,7 @@ from app.views.share import router as share_router
|
||||
from app.views.shared_links import router as shared_links_router
|
||||
from app.views.status import router as status_router
|
||||
from app.views.subscriptions import router as subscriptions_router # Pricing + subscription pages
|
||||
from app.views.system_reset import router as system_reset_router # System reset / factory reset
|
||||
from app.views.wizard import router as wizard_router
|
||||
|
||||
# Create a main router that includes all the view routers
|
||||
@@ -60,6 +63,7 @@ router.include_router(plans_router) # Admin Plan Designer
|
||||
router.include_router(onboarding_router) # User onboarding wizard
|
||||
router.include_router(pipelines_router) # Processing pipelines
|
||||
router.include_router(profile_router) # User self-service profile settings
|
||||
router.include_router(qr_login_router) # QR code mobile login page
|
||||
router.include_router(imap_accounts_router) # Per-user IMAP ingestion accounts
|
||||
router.include_router(integrations_router) # Unified integrations dashboard
|
||||
router.include_router(notifications_router) # User notification dashboard
|
||||
@@ -67,3 +71,5 @@ router.include_router(scheduled_jobs_router) # Admin scheduled batch jobs
|
||||
router.include_router(audit_logs_router) # Comprehensive audit log viewer
|
||||
router.include_router(help_router) # Built-in help / How-To docs
|
||||
router.include_router(compliance_router) # Compliance templates dashboard
|
||||
router.include_router(devices_router) # Mobile devices dashboard
|
||||
router.include_router(system_reset_router) # System reset / factory reset
|
||||
|
||||
@@ -94,6 +94,7 @@ def _inject_global_context(ctx: dict) -> None:
|
||||
"allow_signup",
|
||||
getattr(settings, "multi_user_enabled", False) and getattr(settings, "allow_local_signup", False),
|
||||
)
|
||||
ctx.setdefault("enable_factory_reset", getattr(settings, "enable_factory_reset", False))
|
||||
|
||||
req = ctx.get("request")
|
||||
if req is not None:
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""View route for the Devices management page.
|
||||
|
||||
Renders the ``devices.html`` template where users can see their registered
|
||||
mobile devices, mobile API tokens (created via the mobile SSO flow or QR
|
||||
code login), and revoke access per-device.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Request
|
||||
|
||||
from app.views.base import require_login, templates
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/devices", include_in_schema=False)
|
||||
@require_login
|
||||
async def devices_page(request: Request):
|
||||
"""Render the Devices management page."""
|
||||
return templates.TemplateResponse(
|
||||
"devices.html",
|
||||
{"request": request, "page_title": "Devices"},
|
||||
)
|
||||
+27
-1
@@ -14,6 +14,19 @@ from app.views.base import APIRouter, Depends, get_db, require_login, settings,
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _get_dropbox_callback_url(request: Request) -> str:
|
||||
"""Return the Dropbox OAuth callback URL.
|
||||
|
||||
Uses ``PUBLIC_BASE_URL`` when configured so that the redirect URI displayed
|
||||
to the user (and registered in the Dropbox developer console) matches the
|
||||
one used in the OAuth authorization request. Falls back to deriving the URL
|
||||
from the incoming request when ``PUBLIC_BASE_URL`` is not set.
|
||||
"""
|
||||
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-setup")
|
||||
@require_login
|
||||
async def dropbox_setup_page(
|
||||
@@ -30,6 +43,8 @@ async def dropbox_setup_page(
|
||||
path from the integration's existing config is pre-populated; global
|
||||
admin credentials are never exposed in this mode.
|
||||
"""
|
||||
callback_url = _get_dropbox_callback_url(request)
|
||||
|
||||
if integration_id is not None:
|
||||
owner_id = get_current_owner_id(request)
|
||||
integration = (
|
||||
@@ -46,6 +61,12 @@ async def dropbox_setup_page(
|
||||
cfg = {}
|
||||
# Support both "folder" (DROPBOX destination) and "folder_path" (WATCH_FOLDER source)
|
||||
folder_path = cfg.get("folder", cfg.get("folder_path", ""))
|
||||
# Determine if global credentials are available for users to reuse
|
||||
global_creds_available = bool(
|
||||
settings.dropbox_allow_global_credentials_for_integrations
|
||||
and settings.dropbox_app_key
|
||||
and settings.dropbox_app_secret
|
||||
)
|
||||
return templates.TemplateResponse(
|
||||
"dropbox.html",
|
||||
{
|
||||
@@ -56,9 +77,12 @@ async def dropbox_setup_page(
|
||||
"integration_name": integration.name,
|
||||
"integration_type": integration.integration_type,
|
||||
"folder_path": folder_path,
|
||||
"app_key_value": "",
|
||||
# Only expose the public app key (not the secret) when global creds are allowed
|
||||
"app_key_value": settings.dropbox_app_key if global_creds_available else "",
|
||||
"app_secret_value": "",
|
||||
"refresh_token_value": "",
|
||||
"global_creds_available": global_creds_available,
|
||||
"callback_url": callback_url,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -78,6 +102,7 @@ async def dropbox_setup_page(
|
||||
"integration_id": integration_id,
|
||||
"integration_name": None,
|
||||
"integration_type": None,
|
||||
"callback_url": callback_url,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -108,5 +133,6 @@ async def dropbox_callback(request: Request, code: str = None, error: str = None
|
||||
"app_key_value": "", # The callback will prioritize sessionStorage values
|
||||
"app_secret_value": "", # The callback will prioritize sessionStorage values
|
||||
"folder_path": "", # The callback will prioritize sessionStorage values
|
||||
"callback_url": _get_dropbox_callback_url(request),
|
||||
},
|
||||
)
|
||||
|
||||
+1
-3
@@ -396,9 +396,7 @@ _STEP_TYPE_TO_STAGES: dict[str, list[str]] = {
|
||||
"embed_metadata": ["embed_metadata_into_pdf"],
|
||||
"compute_embedding": ["compute_embedding"],
|
||||
"send_to_destinations": ["finalize_document_storage", "send_to_all_destinations"],
|
||||
# "classify" is defined in PIPELINE_STEP_TYPES but has no Celery log stages yet.
|
||||
# When a classify task is implemented, add its stage key(s) here.
|
||||
"classify": [],
|
||||
"classify": ["classify_document"],
|
||||
}
|
||||
|
||||
# These internal bookkeeping stages are always shown in the flow regardless of
|
||||
|
||||
@@ -45,6 +45,9 @@ async def google_drive_setup_page(
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
cfg = {}
|
||||
folder_id = cfg.get("folder_id", "")
|
||||
# Provide system-wide OAuth credentials when available so users can
|
||||
# authorize without registering their own Google Cloud app.
|
||||
has_system_credentials = bool(settings.google_drive_client_id and settings.google_drive_client_secret)
|
||||
return templates.TemplateResponse(
|
||||
"google_drive.html",
|
||||
{
|
||||
@@ -58,10 +61,13 @@ async def google_drive_setup_page(
|
||||
"use_oauth": True,
|
||||
"oauth_configured": bool(integration.credentials),
|
||||
"sa_configured": False,
|
||||
"client_id": False,
|
||||
"client_id_value": "",
|
||||
"client_secret": False,
|
||||
"client_secret_value": "",
|
||||
"has_system_credentials": has_system_credentials,
|
||||
"client_id": bool(settings.google_drive_client_id) if has_system_credentials else False,
|
||||
"client_id_value": (settings.google_drive_client_id or "" if has_system_credentials else ""),
|
||||
"client_secret": bool(settings.google_drive_client_secret) if has_system_credentials else False,
|
||||
"client_secret_value": (
|
||||
settings.google_drive_client_secret or "" if has_system_credentials else ""
|
||||
),
|
||||
"refresh_token": False,
|
||||
"refresh_token_value": "",
|
||||
"has_credentials_json": False,
|
||||
@@ -90,6 +96,7 @@ async def google_drive_setup_page(
|
||||
"use_oauth": use_oauth,
|
||||
"oauth_configured": oauth_configured,
|
||||
"sa_configured": sa_configured,
|
||||
"has_system_credentials": bool(settings.google_drive_client_id and settings.google_drive_client_secret),
|
||||
"client_id": bool(settings.google_drive_client_id),
|
||||
"client_id_value": settings.google_drive_client_id or "",
|
||||
"client_secret": bool(settings.google_drive_client_secret),
|
||||
|
||||
+10
-5
@@ -44,6 +44,9 @@ async def onedrive_setup_page(
|
||||
cfg = {}
|
||||
# Support both "folder_path" (WATCH_FOLDER / ONEDRIVE destination)
|
||||
folder_path = cfg.get("folder_path", cfg.get("folder", ""))
|
||||
# Provide system-wide app credentials when available so users can
|
||||
# authorize without registering their own Azure/OneDrive app.
|
||||
has_system_credentials = bool(settings.onedrive_client_id and settings.onedrive_client_secret)
|
||||
return templates.TemplateResponse(
|
||||
"onedrive.html",
|
||||
{
|
||||
@@ -54,11 +57,12 @@ async def onedrive_setup_page(
|
||||
"integration_name": integration.name,
|
||||
"integration_type": integration.integration_type,
|
||||
"folder_path": folder_path,
|
||||
"client_id": False,
|
||||
"client_id_value": "",
|
||||
"client_secret": False,
|
||||
"client_secret_value": "",
|
||||
"tenant_id": "common",
|
||||
"has_system_credentials": has_system_credentials,
|
||||
"client_id": bool(settings.onedrive_client_id) if has_system_credentials else False,
|
||||
"client_id_value": settings.onedrive_client_id or "" if has_system_credentials else "",
|
||||
"client_secret": bool(settings.onedrive_client_secret) if has_system_credentials else False,
|
||||
"client_secret_value": (settings.onedrive_client_secret or "" if has_system_credentials else ""),
|
||||
"tenant_id": settings.onedrive_tenant_id or "common",
|
||||
"refresh_token": False,
|
||||
"refresh_token_value": "",
|
||||
},
|
||||
@@ -75,6 +79,7 @@ async def onedrive_setup_page(
|
||||
"request": request,
|
||||
"user_mode": False,
|
||||
"is_configured": is_configured,
|
||||
"has_system_credentials": bool(settings.onedrive_client_id and settings.onedrive_client_secret),
|
||||
"client_id": bool(settings.onedrive_client_id),
|
||||
"client_id_value": settings.onedrive_client_id or "",
|
||||
"client_secret": bool(settings.onedrive_client_secret),
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
"""View route for the QR code mobile login page.
|
||||
|
||||
Route:
|
||||
GET /qr-login — renders the QR login page (requires login)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from app.views.base import APIRouter, require_login, templates
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/qr-login", include_in_schema=False)
|
||||
@require_login
|
||||
async def qr_login_page(request: Request):
|
||||
"""Serve the QR code login page for mobile app authentication."""
|
||||
return templates.TemplateResponse(
|
||||
"qr_login.html",
|
||||
{"request": request},
|
||||
)
|
||||
@@ -0,0 +1,40 @@
|
||||
"""
|
||||
System reset view — admin-only UI page.
|
||||
|
||||
Renders a confirmation-heavy page that allows administrators to:
|
||||
1. **Full Reset** — wipe all user data (DB + disk) for a fresh start.
|
||||
2. **Reset & Re-import** — move originals to a reimport folder, wipe,
|
||||
and let the watch-folder mechanism re-ingest them.
|
||||
|
||||
Both options are gated behind the ``ENABLE_FACTORY_RESET`` feature flag.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import Depends, Request
|
||||
from fastapi.responses import RedirectResponse, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.views.base import APIRouter, get_db, require_login, templates
|
||||
from app.views.settings import require_admin_access
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/admin/system-reset")
|
||||
@require_login
|
||||
@require_admin_access
|
||||
async def system_reset_page(request: Request, db: Session = Depends(get_db)) -> Response:
|
||||
"""Render the system reset administration page."""
|
||||
if not settings.enable_factory_reset:
|
||||
return RedirectResponse(url="/settings", status_code=302)
|
||||
|
||||
return templates.TemplateResponse(
|
||||
"system_reset.html",
|
||||
{
|
||||
"request": request,
|
||||
"factory_reset_on_startup": settings.factory_reset_on_startup,
|
||||
},
|
||||
)
|
||||
Reference in New Issue
Block a user