Merge branch 'main' into sentinel-fix-ssrf-dns-resolution-16520734505214840647

This commit is contained in:
Christian Krakau-Louis
2026-03-23 17:16:53 +01:00
committed by GitHub
185 changed files with 1937 additions and 25684 deletions
-8
View File
@@ -9,12 +9,9 @@ from fastapi import APIRouter
from app.api.admin_users import router as admin_users_router
from app.api.api_tokens import router as api_tokens_router
from app.api.audit_logs import router as audit_logs_router
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.comments import router as comments_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
@@ -46,7 +43,6 @@ 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.sharing import router as sharing_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
@@ -108,7 +104,3 @@ 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)
router.include_router(comments_router)
router.include_router(sharing_router)
+11 -68
View File
@@ -13,7 +13,7 @@ plaintext is returned exactly once at creation time.
import hashlib
import logging
import secrets
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from typing import Annotated, Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
@@ -105,7 +105,6 @@ def _token_to_dict(t: ApiToken) -> dict[str, Any]:
"last_used_ip": t.last_used_ip,
"created_at": t.created_at,
"revoked_at": t.revoked_at,
"expires_at": t.expires_at,
}
@@ -118,12 +117,6 @@ 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):
@@ -137,7 +130,6 @@ 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}
@@ -168,16 +160,11 @@ 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)
@@ -198,7 +185,6 @@ 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,
}
@@ -249,73 +235,30 @@ async def list_mobile_tokens(
@router.delete("/{token_id}", status_code=status.HTTP_200_OK)
async def revoke_or_delete_token(
async def revoke_token(
token_id: int,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, str]:
"""Revoke or permanently delete an API token.
"""Revoke (soft-delete) an API token.
* **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.
The token row is kept for audit purposes but marked inactive with a
``revoked_at`` timestamp.
"""
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:
# 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"}
if not db_token.is_active:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already revoked")
# Hard-delete an already-revoked token.
try:
db.delete(db_token)
db_token.is_active = False
db_token.revoked_at = datetime.now(timezone.utc)
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"}
@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)
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
return {"detail": "Token revoked"}
-311
View File
@@ -1,311 +0,0 @@
"""API endpoints for Zapier / Make.com automation integration.
Provides a REST hooks subscription interface for outgoing triggers and
incoming action endpoints that external automation platforms can call.
Outgoing triggers:
External platforms subscribe to DocuElevate events via
``POST /api/automation/hooks/subscribe``. When a subscribed event
fires, DocuElevate POSTs a flat Zapier-compatible JSON payload to the
registered ``target_url``.
Incoming actions:
``POST /api/automation/actions/upload`` allows automation platforms to
push documents into DocuElevate for processing.
Authentication:
All endpoints require a valid API token via ``Authorization: Bearer``
header.
"""
import json
import logging
import os
import tempfile
from typing import Annotated, Any
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from app.config import settings
from app.database import get_db
from app.models import AutomationHook
from app.utils.automation_hooks import SAMPLE_PAYLOADS
from app.utils.webhook import VALID_EVENTS
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/automation", tags=["automation"])
DbSession = Annotated[Session, Depends(get_db)]
# ---------------------------------------------------------------------------
# Auth helper require a valid API token (Bearer)
# ---------------------------------------------------------------------------
def _require_api_user(request: Request) -> dict:
"""Ensure the caller is authenticated via session or API token.
Raises:
HTTPException: 401 if not authenticated, 403 if automation hooks are disabled.
"""
if not settings.automation_hooks_enabled:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Automation hooks are disabled",
)
# Check for API-token user first (set by auth middleware)
user = getattr(request.state, "api_token_user", None)
if user:
return user
# Fall back to session user
user = request.session.get("user")
if user:
return user
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Authentication required (Bearer token or session)",
)
AuthUser = Annotated[dict, Depends(_require_api_user)]
# ---------------------------------------------------------------------------
# Pydantic schemas
# ---------------------------------------------------------------------------
class HookSubscribe(BaseModel):
"""Schema for subscribing to automation hook events."""
target_url: str = Field(..., min_length=1, max_length=2048, description="URL to POST event payloads to")
events: list[str] = Field(..., min_length=1, description="Event types to subscribe to")
secret: str | None = Field(default=None, max_length=512, description="Optional HMAC-SHA256 signing secret")
hook_type: str = Field(
default="generic",
max_length=50,
description="Platform identifier (zapier, make, generic)",
)
description: str | None = Field(default=None, max_length=500, description="Optional human-readable label")
class HookResponse(BaseModel):
"""Schema returned when listing or creating hooks."""
id: int
target_url: str
events: list[str]
is_active: bool
hook_type: str
description: str | None
has_secret: bool
model_config = {"from_attributes": True}
class ActionUploadResponse(BaseModel):
"""Response after an automation action uploads a document."""
status: str
filename: str
task_id: str | None = None
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _validate_events(events: list[str]) -> None:
"""Raise 422 if any event name is not recognised."""
invalid = set(events) - VALID_EVENTS
if invalid:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Invalid event(s): {', '.join(sorted(invalid))}. Valid: {', '.join(sorted(VALID_EVENTS))}",
)
def _hook_to_response(hook: AutomationHook) -> dict[str, Any]:
"""Convert a DB model instance to a response dict."""
try:
events = json.loads(hook.events)
except (json.JSONDecodeError, TypeError):
events = []
return {
"id": hook.id,
"target_url": hook.target_url,
"events": events,
"is_active": hook.is_active,
"hook_type": hook.hook_type,
"description": hook.description,
"has_secret": hook.secret is not None and len(hook.secret) > 0,
}
# ---------------------------------------------------------------------------
# Outgoing triggers REST hooks subscription endpoints
# ---------------------------------------------------------------------------
@router.post(
"/hooks/subscribe",
status_code=status.HTTP_201_CREATED,
summary="Subscribe to automation events (REST hooks)",
)
def subscribe_hook(body: HookSubscribe, db: DbSession, user: AuthUser) -> dict[str, Any]:
"""Register a new automation hook subscription.
Zapier and Make.com call this endpoint to subscribe to DocuElevate
events. When an event fires, a flat JSON payload is POSTed to
``target_url``.
"""
_validate_events(body.events)
hook = AutomationHook(
target_url=body.target_url,
secret=body.secret,
events=json.dumps(sorted(body.events)),
is_active=True,
hook_type=body.hook_type or "generic",
description=body.description,
)
try:
db.add(hook)
db.commit()
db.refresh(hook)
except Exception:
db.rollback()
raise
logger.info("Automation hook %d created (type=%s) for events %s", hook.id, hook.hook_type, body.events)
return _hook_to_response(hook)
@router.delete(
"/hooks/{hook_id}",
status_code=status.HTTP_204_NO_CONTENT,
summary="Unsubscribe an automation hook",
)
def unsubscribe_hook(hook_id: int, db: DbSession, user: AuthUser) -> None:
"""Remove an automation hook subscription.
Zapier calls this endpoint when a Zap is turned off or deleted.
"""
hook = db.query(AutomationHook).filter(AutomationHook.id == hook_id).first()
if not hook:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Hook not found")
try:
db.delete(hook)
db.commit()
except Exception:
db.rollback()
raise
logger.info("Automation hook %d deleted", hook_id)
@router.get("/hooks", summary="List automation hook subscriptions")
def list_hooks(db: DbSession, user: AuthUser) -> list[dict[str, Any]]:
"""Return all active automation hook subscriptions."""
hooks = db.query(AutomationHook).order_by(AutomationHook.id).all()
return [_hook_to_response(h) for h in hooks]
# ---------------------------------------------------------------------------
# Outgoing triggers sample data for Zapier field mapping
# ---------------------------------------------------------------------------
@router.get("/triggers/sample/{event}", summary="Get sample trigger data")
def get_trigger_sample(event: str, user: AuthUser) -> list[dict[str, Any]]:
"""Return sample payload data for the given event type.
Zapier uses this during Zap setup to discover available fields and
provide a mapping interface. The response is wrapped in an array
as Zapier expects.
"""
if event not in VALID_EVENTS:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Unknown event: {event}. Valid: {', '.join(sorted(VALID_EVENTS))}",
)
sample = SAMPLE_PAYLOADS.get(event, {"id": "evt_sample", "event": event, "timestamp": 0})
return [sample]
# ---------------------------------------------------------------------------
# Outgoing triggers list valid events
# ---------------------------------------------------------------------------
@router.get("/events", summary="List valid automation event types")
def list_events(user: AuthUser) -> list[str]:
"""Return the list of valid event types that automation hooks can subscribe to."""
return sorted(VALID_EVENTS)
# ---------------------------------------------------------------------------
# Incoming actions endpoints that Zapier / Make.com can call
# ---------------------------------------------------------------------------
@router.post("/actions/upload", summary="Upload a document (incoming action)")
def action_upload(
request: Request,
db: DbSession,
user: AuthUser,
file: UploadFile = File(...),
) -> dict[str, Any]:
"""Accept a document upload from an automation platform.
This endpoint allows Zapier or Make.com to push a document into
DocuElevate for processing. The file is saved to the work directory
and a background processing task is queued.
"""
if not file.filename:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
# Sanitise filename to prevent path traversal attacks
safe_filename = os.path.basename(file.filename)
if not safe_filename:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
owner_id = user.get("preferred_username") or user.get("email") or user.get("id", "automation")
workdir = settings.workdir or tempfile.gettempdir()
upload_dir = os.path.join(workdir, "uploads")
os.makedirs(upload_dir, exist_ok=True)
dest_path = os.path.join(upload_dir, safe_filename)
try:
contents = file.file.read()
with open(dest_path, "wb") as f:
f.write(contents)
except Exception as exc:
logger.error("Failed to save uploaded file: %s", exc)
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to save file")
# Queue background processing
task_id = None
try:
from app.tasks.process_document import process_document
result = process_document.delay(dest_path, owner_id)
task_id = result.id
logger.info("Automation upload queued: file=%s, task=%s, owner=%s", safe_filename, task_id, owner_id)
except Exception as exc:
logger.warning("Could not queue processing task (Celery may be unavailable): %s", exc)
return {
"status": "accepted",
"filename": safe_filename,
"task_id": task_id,
}
-325
View File
@@ -1,325 +0,0 @@
"""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)
-751
View File
@@ -1,751 +0,0 @@
"""Document comments and annotations API endpoints.
Provides CRUD operations for threaded comments on documents,
text annotations on PDF pages, and a list of mentionable users
for the @mention feature.
"""
import json
import logging
import re
from typing import Annotated, Any
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session
from app.auth import get_current_user_id, require_login
from app.database import get_db
from app.models import (
FILE_SHARE_ROLE_VIEWER,
DocumentAnnotation,
DocumentComment,
FileRecord,
FileShare,
UserProfile,
)
from app.utils.user_scope import get_current_owner_id, has_file_role
logger = logging.getLogger(__name__)
router = APIRouter(tags=["comments"])
DbSession = Annotated[Session, Depends(get_db)]
# Constraints
MAX_COMMENT_BODY_LENGTH = 10_000
MAX_ANNOTATION_CONTENT_LENGTH = 5_000
# Allowed annotation types
ALLOWED_ANNOTATION_TYPES = frozenset({"note", "highlight", "underline", "strikethrough"})
# Simple pattern for @mentions matches @username tokens inside comment body
_MENTION_PATTERN = re.compile(r"@([\w.\-]+)")
def _extract_mentions(body: str) -> list[str]:
"""Extract unique @mentioned usernames from a comment body.
Args:
body: The raw comment text.
Returns:
A deduplicated list of mentioned usernames (without the ``@`` prefix).
"""
return list(dict.fromkeys(_MENTION_PATTERN.findall(body)))
def _serialize_comment(c: DocumentComment) -> dict[str, Any]:
"""Serialize a DocumentComment to a JSON-friendly dict.
Args:
c: The comment model instance.
Returns:
A dictionary representation of the comment.
"""
mentions: list[str] = []
if c.mentions:
try:
mentions = json.loads(c.mentions)
except (json.JSONDecodeError, TypeError):
pass
return {
"id": c.id,
"file_id": c.file_id,
"user_id": c.user_id,
"parent_id": c.parent_id,
"body": c.body,
"mentions": mentions,
"is_resolved": c.is_resolved,
"created_at": c.created_at.isoformat() if c.created_at else None,
"updated_at": c.updated_at.isoformat() if c.updated_at else None,
}
def _serialize_annotation(a: DocumentAnnotation) -> dict[str, Any]:
"""Serialize a DocumentAnnotation to a JSON-friendly dict.
Args:
a: The annotation model instance.
Returns:
A dictionary representation of the annotation.
"""
return {
"id": a.id,
"file_id": a.file_id,
"user_id": a.user_id,
"page": a.page,
"x": a.x,
"y": a.y,
"width": a.width,
"height": a.height,
"content": a.content,
"annotation_type": a.annotation_type,
"color": a.color,
"created_at": a.created_at.isoformat() if a.created_at else None,
"updated_at": a.updated_at.isoformat() if a.updated_at else None,
}
def _build_thread_tree(comments: list[DocumentComment]) -> list[dict[str, Any]]:
"""Organize a flat list of comments into a threaded tree structure.
Top-level comments (``parent_id is None``) appear as root nodes.
Replies are nested inside their parent's ``replies`` list.
Args:
comments: All comments for a given document, ordered by ``created_at``.
Returns:
A list of root-level comment dicts, each with a ``replies`` key.
"""
by_id: dict[int, dict[str, Any]] = {}
roots: list[dict[str, Any]] = []
for c in comments:
node = _serialize_comment(c)
node["replies"] = []
by_id[c.id] = node
for c in comments:
node = by_id[c.id]
if c.parent_id and c.parent_id in by_id:
by_id[c.parent_id]["replies"].append(node)
else:
roots.append(node)
return roots
# ---------------------------------------------------------------------------
# Comments endpoints
# ---------------------------------------------------------------------------
@router.get("/files/{file_id}/comments")
@require_login
def list_comments(request: Request, file_id: int, db: DbSession):
"""List all comments for a document, organized into threads.
Returns a threaded tree where top-level comments contain nested
``replies``. Requires at least viewer access.
Path Parameters:
file_id: The ID of the document.
Returns:
A dict with ``file_id``, ``comments`` (threaded), and ``total``.
"""
user_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not is_admin and not has_file_role(file_record, user_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
comments = (
db.query(DocumentComment).filter(DocumentComment.file_id == file_id).order_by(DocumentComment.created_at).all()
)
return {
"file_id": file_id,
"comments": _build_thread_tree(comments),
"total": len(comments),
}
@router.post("/files/{file_id}/comments", status_code=status.HTTP_201_CREATED)
@require_login
def create_comment(
request: Request,
file_id: int,
db: DbSession,
body: str = Body(..., embed=True),
parent_id: int | None = Body(None, embed=True),
):
"""Create a new comment on a document.
Automatically extracts @mentions from the comment body and stores
them for later notification or UI highlighting. When multi-user
mode is enabled, any mentioned user that does not already have
access to the document is automatically granted ``viewer`` access by
the file owner so they can read the file and continue the discussion.
Path Parameters:
file_id: The ID of the document to comment on.
Request body (JSON):
body: Comment text (required, max 10 000 characters).
parent_id: ID of the parent comment for threaded replies (optional).
Returns:
The created comment object.
"""
user_id = get_current_user_id(request)
owner_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not is_admin and not has_file_role(file_record, owner_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not isinstance(body, str) or not body.strip():
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="body is required and must be non-empty",
)
body = body.strip()
if len(body) > MAX_COMMENT_BODY_LENGTH:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"body must be at most {MAX_COMMENT_BODY_LENGTH} characters",
)
if parent_id is not None:
parent = (
db.query(DocumentComment)
.filter(DocumentComment.id == parent_id, DocumentComment.file_id == file_id)
.first()
)
if not parent:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Parent comment not found",
)
mentions = _extract_mentions(body)
comment = DocumentComment(
file_id=file_id,
user_id=user_id,
parent_id=parent_id,
body=body,
mentions=json.dumps(mentions) if mentions else None,
)
try:
db.add(comment)
db.flush() # write comment so we can get its id before committing
# Auto-share the file with mentioned users that don't have access yet.
# Only do this in multi-user mode and only when the file has an owner
# (unowned files are already visible to all authenticated users).
if mentions and file_record.owner_id is not None:
from app.config import settings as _settings
if _settings.multi_user_enabled:
for mentioned_user in mentions:
# Skip the file owner (already has full access) and the commenter
# themselves (they already have access to be posting a comment).
if mentioned_user in {file_record.owner_id, owner_id}:
continue
existing_share = (
db.query(FileShare)
.filter(
FileShare.file_id == file_id,
FileShare.shared_with_user_id == mentioned_user,
)
.first()
)
if not existing_share:
auto_share = FileShare(
file_id=file_id,
owner_id=file_record.owner_id,
shared_with_user_id=mentioned_user,
role=FILE_SHARE_ROLE_VIEWER,
)
db.add(auto_share)
logger.info(
"Auto-shared file_id=%s with mentioned user=%s as viewer",
file_id,
mentioned_user,
)
db.commit()
db.refresh(comment)
except HTTPException:
raise
except Exception:
db.rollback()
logger.exception("Failed to create comment on file_id=%s", file_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to create comment",
)
logger.info("Comment created: id=%s, file_id=%s, user=%s", comment.id, file_id, user_id)
return _serialize_comment(comment)
@router.put("/files/{file_id}/comments/{comment_id}")
@require_login
def update_comment(
request: Request,
file_id: int,
comment_id: int,
db: DbSession,
body: str = Body(..., embed=True),
):
"""Update the body of an existing comment.
Only the comment author may update the comment. Mentions are
re-extracted from the updated body.
Path Parameters:
file_id: The ID of the document.
comment_id: The ID of the comment to update.
Request body (JSON):
body: New comment text (required).
Returns:
The updated comment object.
"""
user_id = get_current_user_id(request)
comment = (
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
)
if not comment:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
if comment.user_id != user_id:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only edit your own comments")
if not isinstance(body, str) or not body.strip():
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="body is required and must be non-empty",
)
body = body.strip()
if len(body) > MAX_COMMENT_BODY_LENGTH:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"body must be at most {MAX_COMMENT_BODY_LENGTH} characters",
)
mentions = _extract_mentions(body)
comment.body = body
comment.mentions = json.dumps(mentions) if mentions else None
try:
db.commit()
db.refresh(comment)
except Exception:
db.rollback()
logger.exception("Failed to update comment id=%s", comment_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to update comment",
)
logger.info("Comment updated: id=%s, user=%s", comment_id, user_id)
return _serialize_comment(comment)
@router.delete("/files/{file_id}/comments/{comment_id}", status_code=status.HTTP_204_NO_CONTENT)
@require_login
def delete_comment(request: Request, file_id: int, comment_id: int, db: DbSession):
"""Delete a comment.
Only the comment author may delete the comment. Replies to the
deleted comment are **not** removed — they become orphaned root
comments so that conversation context is preserved.
Path Parameters:
file_id: The ID of the document.
comment_id: The ID of the comment to delete.
"""
user_id = get_current_user_id(request)
comment = (
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
)
if not comment:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
if comment.user_id != user_id:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only delete your own comments")
try:
db.delete(comment)
db.commit()
except Exception:
db.rollback()
logger.exception("Failed to delete comment id=%s", comment_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to delete comment",
)
logger.info("Comment deleted: id=%s, user=%s", comment_id, user_id)
@router.patch("/files/{file_id}/comments/{comment_id}/resolve")
@require_login
def resolve_comment(
request: Request,
file_id: int,
comment_id: int,
db: DbSession,
is_resolved: bool = Body(..., embed=True),
):
"""Mark a top-level comment thread as resolved or unresolved.
Path Parameters:
file_id: The ID of the document.
comment_id: The ID of the comment to resolve / unresolve.
Request body (JSON):
is_resolved: ``true`` to resolve, ``false`` to unresolve.
Returns:
The updated comment object.
"""
comment = (
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
)
if not comment:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
comment.is_resolved = is_resolved
try:
db.commit()
db.refresh(comment)
except Exception:
db.rollback()
logger.exception("Failed to resolve comment id=%s", comment_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to update comment",
)
logger.info("Comment %s: id=%s", "resolved" if is_resolved else "unresolved", comment_id)
return _serialize_comment(comment)
# ---------------------------------------------------------------------------
# Annotations endpoints
# ---------------------------------------------------------------------------
@router.get("/files/{file_id}/annotations")
@require_login
def list_annotations(request: Request, file_id: int, db: DbSession):
"""List all annotations for a document.
Requires at least viewer access.
Path Parameters:
file_id: The ID of the document.
Returns:
A dict with ``file_id``, ``annotations``, and ``total``.
"""
user_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not is_admin and not has_file_role(file_record, user_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
annotations = (
db.query(DocumentAnnotation)
.filter(DocumentAnnotation.file_id == file_id)
.order_by(DocumentAnnotation.page, DocumentAnnotation.created_at)
.all()
)
return {
"file_id": file_id,
"annotations": [_serialize_annotation(a) for a in annotations],
"total": len(annotations),
}
@router.post("/files/{file_id}/annotations", status_code=status.HTTP_201_CREATED)
@require_login
def create_annotation(
request: Request,
file_id: int,
db: DbSession,
page: int = Body(..., embed=True),
x: float = Body(..., embed=True),
y: float = Body(..., embed=True),
content: str = Body(..., embed=True),
width: float = Body(0, embed=True),
height: float = Body(0, embed=True),
annotation_type: str = Body("note", embed=True),
color: str | None = Body(None, embed=True),
):
"""Create a new annotation on a PDF page.
Path Parameters:
file_id: The ID of the document.
Request body (JSON):
page: Page number (1-based, required).
x: Horizontal position on the page (required).
y: Vertical position on the page (required).
content: Annotation text (required, max 5 000 characters).
width: Width of the annotation bounding box (default 0).
height: Height of the annotation bounding box (default 0).
annotation_type: One of ``note``, ``highlight``, ``underline``,
``strikethrough`` (default ``note``).
color: Optional CSS colour string (e.g. ``#ff0000``).
Returns:
The created annotation object.
"""
user_id = get_current_user_id(request)
owner_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not is_admin and not has_file_role(file_record, owner_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if not isinstance(content, str) or not content.strip():
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="content is required and must be non-empty",
)
content = content.strip()
if len(content) > MAX_ANNOTATION_CONTENT_LENGTH:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"content must be at most {MAX_ANNOTATION_CONTENT_LENGTH} characters",
)
if page < 1:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="page must be >= 1",
)
if annotation_type not in ALLOWED_ANNOTATION_TYPES:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"annotation_type must be one of: {', '.join(sorted(ALLOWED_ANNOTATION_TYPES))}",
)
annotation = DocumentAnnotation(
file_id=file_id,
user_id=user_id,
page=page,
x=x,
y=y,
width=width,
height=height,
content=content,
annotation_type=annotation_type,
color=color,
)
try:
db.add(annotation)
db.commit()
db.refresh(annotation)
except Exception:
db.rollback()
logger.exception("Failed to create annotation on file_id=%s", file_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to create annotation",
)
logger.info("Annotation created: id=%s, file_id=%s, user=%s", annotation.id, file_id, user_id)
return _serialize_annotation(annotation)
@router.put("/files/{file_id}/annotations/{annotation_id}")
@require_login
def update_annotation(
request: Request,
file_id: int,
annotation_id: int,
db: DbSession,
content: str | None = Body(None, embed=True),
x: float | None = Body(None, embed=True),
y: float | None = Body(None, embed=True),
width: float | None = Body(None, embed=True),
height: float | None = Body(None, embed=True),
annotation_type: str | None = Body(None, embed=True),
color: str | None = Body(None, embed=True),
):
"""Update an existing annotation.
Only the annotation author may update the annotation.
Path Parameters:
file_id: The ID of the document.
annotation_id: The ID of the annotation to update.
Request body (JSON):
Any subset of ``content``, ``x``, ``y``, ``width``, ``height``,
``annotation_type``, and ``color``.
Returns:
The updated annotation object.
"""
user_id = get_current_user_id(request)
annotation = (
db.query(DocumentAnnotation)
.filter(DocumentAnnotation.id == annotation_id, DocumentAnnotation.file_id == file_id)
.first()
)
if not annotation:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Annotation not found")
if annotation.user_id != user_id:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only edit your own annotations")
if content is not None:
content = content.strip() if isinstance(content, str) else ""
if not content:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="content must be non-empty",
)
if len(content) > MAX_ANNOTATION_CONTENT_LENGTH:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"content must be at most {MAX_ANNOTATION_CONTENT_LENGTH} characters",
)
annotation.content = content
if x is not None:
annotation.x = x
if y is not None:
annotation.y = y
if width is not None:
annotation.width = width
if height is not None:
annotation.height = height
if annotation_type is not None:
if annotation_type not in ALLOWED_ANNOTATION_TYPES:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"annotation_type must be one of: {', '.join(sorted(ALLOWED_ANNOTATION_TYPES))}",
)
annotation.annotation_type = annotation_type
if color is not None:
annotation.color = color
try:
db.commit()
db.refresh(annotation)
except Exception:
db.rollback()
logger.exception("Failed to update annotation id=%s", annotation_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to update annotation",
)
logger.info("Annotation updated: id=%s, user=%s", annotation_id, user_id)
return _serialize_annotation(annotation)
@router.delete("/files/{file_id}/annotations/{annotation_id}", status_code=status.HTTP_204_NO_CONTENT)
@require_login
def delete_annotation(request: Request, file_id: int, annotation_id: int, db: DbSession):
"""Delete an annotation.
Only the annotation author may delete the annotation.
Path Parameters:
file_id: The ID of the document.
annotation_id: The ID of the annotation to delete.
"""
user_id = get_current_user_id(request)
annotation = (
db.query(DocumentAnnotation)
.filter(DocumentAnnotation.id == annotation_id, DocumentAnnotation.file_id == file_id)
.first()
)
if not annotation:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Annotation not found")
if annotation.user_id != user_id:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only delete your own annotations")
try:
db.delete(annotation)
db.commit()
except Exception:
db.rollback()
logger.exception("Failed to delete annotation id=%s", annotation_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to delete annotation",
)
logger.info("Annotation deleted: id=%s, user=%s", annotation_id, user_id)
# ---------------------------------------------------------------------------
# Mentionable users endpoint
# ---------------------------------------------------------------------------
@router.get("/users/mentionable")
@require_login
def list_mentionable_users(request: Request, db: DbSession):
"""List users that can be @mentioned in comments.
Returns all user profiles that are not blocked, sorted by
``display_name``.
Returns:
A list of ``{user_id, display_name}`` objects.
"""
profiles = db.query(UserProfile).filter(UserProfile.is_blocked.is_(False)).order_by(UserProfile.display_name).all()
return [
{
"user_id": p.user_id,
"display_name": p.display_name or p.user_id,
}
for p in profiles
]
-61
View File
@@ -21,67 +21,6 @@ _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
-174
View File
@@ -5,10 +5,8 @@ Dropbox API endpoints
import logging
import os
from typing import Annotated, Optional
from urllib.parse import quote
import httpx
import requests
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
from sqlalchemy.orm import Session
@@ -25,93 +23,6 @@ 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(
@@ -299,91 +210,6 @@ 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(
+14 -64
View File
@@ -20,7 +20,6 @@ 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
@@ -30,7 +29,7 @@ from app.utils.file_queries import apply_status_filter
from app.utils.file_status import get_files_processing_status
from app.utils.filename_utils import sanitize_filename
from app.utils.input_validation import validate_search_query, validate_sort_field, validate_sort_order
from app.utils.user_scope import apply_owner_filter, get_current_owner_id, get_file_role
from app.utils.user_scope import apply_owner_filter, get_current_owner_id
# Set up logging
logger = logging.getLogger(__name__)
@@ -300,7 +299,6 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
"""
Delete a file record from the database.
This only removes the database entry, not the actual file.
Only the file owner (or an admin) may delete a document.
"""
# Check if file deletion is allowed
if not settings.allow_file_delete:
@@ -315,18 +313,6 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
if not file_record:
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
# Enforce owner-only deletion in multi-user mode
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
if not is_admin:
owner_id = get_current_owner_id(request)
role = get_file_role(file_record, owner_id, db)
if role != "owner":
raise HTTPException(
status_code=403,
detail="Only the file owner can delete this document",
)
# Log the deletion
logger.info(f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}")
@@ -353,7 +339,6 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
"""
Delete multiple file records from the database.
This only removes the database entries, not the actual files.
Only the file owner (or an admin) may delete each document.
"""
# Check if file deletion is allowed
if not settings.allow_file_delete:
@@ -368,18 +353,6 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
if not file_records:
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
# Enforce owner-only deletion in multi-user mode
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
if not is_admin:
owner_id = get_current_owner_id(request)
non_owner_ids = [f.id for f in file_records if get_file_role(f, owner_id, db) != "owner"]
if non_owner_ids:
raise HTTPException(
status_code=403,
detail=f"You can only delete files you own. Not owner of file IDs: {non_owner_ids}",
)
deleted_count = len(file_records)
deleted_ids = [f.id for f in file_records]
@@ -1294,12 +1267,7 @@ 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.
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).
"""
"""Check for an exact duplicate of the uploaded file and return a warning if found."""
if not settings.enable_deduplication:
return None
@@ -1318,8 +1286,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 is an exact duplicate of an already-processed document. "
"It has not been queued for processing again."
"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."
),
}
except Exception as e:
@@ -1330,12 +1298,7 @@ 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(...),
_rate_ok: None = Depends(require_upload_rate_limit),
):
async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...)):
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
workdir = settings.workdir
@@ -1421,25 +1384,6 @@ async def ui_upload(
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()
@@ -1503,8 +1447,6 @@ async def ui_upload(
".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)
@@ -1518,12 +1460,20 @@ async def ui_upload(
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)
return {
# 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 = {
"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
# ---------------------------------------------------------------------------
+6
View File
@@ -17,6 +17,7 @@ from sqlalchemy.orm import Session
from app.database import get_db
from app.models import UserImapAccount
from app.utils.encryption import decrypt_value, encrypt_value
from app.utils.network import is_private_ip
from app.utils.subscription import get_tier, get_user_tier_id
from app.utils.user_scope import get_current_owner_id
@@ -187,6 +188,11 @@ def _test_imap_connection(host: str, port: int, username: str, password: str, us
Returns a dict with ``{"success": bool, "message": str}``.
"""
# Security: Prevent SSRF by blocking connections to internal IPs
if is_private_ip(host):
logger.warning("SSRF blocked: Attempt to connect to private IP %s", host)
return {"success": False, "message": "Connection error: Invalid hostname or IP address"}
try:
if use_ssl:
mail = imaplib.IMAP4_SSL(host, port)
+11 -65
View File
@@ -32,21 +32,6 @@ 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"])
@@ -565,50 +550,9 @@ 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 httpx
import urllib.request
cfg = config or {}
creds = credentials or {}
@@ -635,21 +579,23 @@ def _test_webdav_connection(config: dict[str, Any] | None, credentials: dict[str
return {"success": False, "message": "URLs pointing to internal or private networks are not allowed"}
try:
auth = (username, password) if username and password else None
headers = {"Depth": "0"}
import base64
# Use httpx for secure connection testing, avoiding urllib vulnerabilities
resp = httpx.request("PROPFIND", url, auth=auth, headers=headers, timeout=10.0, follow_redirects=False)
if resp.status_code < 400:
return {"success": True, "message": "WebDAV connection successful"}
return {"success": False, "message": f"WebDAV returned HTTP {resp.status_code}"}
req = urllib.request.Request(url, method="PROPFIND") # noqa: S310
if username and password:
token = base64.b64encode(f"{username}:{password}".encode()).decode()
req.add_header("Authorization", f"Basic {token}")
req.add_header("Depth", "0")
with urllib.request.urlopen(req, timeout=10) as resp: # noqa: S310
if resp.status < 400:
return {"success": True, "message": "WebDAV connection successful"}
return {"success": False, "message": f"WebDAV returned HTTP {resp.status}"}
except Exception as exc: # noqa: BLE001
logger.warning("WebDAV connection error for %s: %s", hostname, exc)
return {"success": False, "message": "WebDAV connection failed — check URL and credentials"}
_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,
+8 -23
View File
@@ -120,7 +120,6 @@ class WhoAmIResponse(BaseModel):
email: str | None
avatar_url: str | None
is_admin: bool
preferred_language: str | None
# ---------------------------------------------------------------------------
@@ -274,44 +273,31 @@ async def list_devices(
return [_device_to_response(d) for d in devices]
@router.delete("/devices/{device_id}", status_code=status.HTTP_200_OK)
@router.delete("/devices/{device_id}", status_code=status.HTTP_204_NO_CONTENT)
@require_login
async def deactivate_device(
request: Request,
device_id: int,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, str]:
"""Deactivate or permanently delete a push-notification device registration.
) -> None:
"""Deactivate a push-notification device registration.
* **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.
The device record is kept for audit purposes but will no longer receive
push notifications.
"""
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")
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.
device.is_active = False
try:
db.delete(device)
db.commit()
except Exception:
db.rollback()
raise
logger.info("Mobile device permanently deleted: id=%s owner=%s", device_id, owner_id)
return {"detail": "Device deleted"}
logger.info("Mobile device deactivated: id=%s owner=%s", device_id, owner_id)
@router.get("/whoami", response_model=WhoAmIResponse)
@@ -358,5 +344,4 @@ async def whoami(
"email": email,
"avatar_url": avatar_url,
"is_admin": is_admin,
"preferred_language": profile.preferred_language if profile else None,
}
-98
View File
@@ -7,7 +7,6 @@ from datetime import datetime, timedelta
from typing import Annotated, Optional
import httpx
import requests
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
from sqlalchemy.orm import Session
@@ -57,7 +56,6 @@ 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),
}
@@ -185,102 +183,6 @@ 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:
+2 -8
View File
@@ -117,14 +117,8 @@ PIPELINE_STEP_TYPES: dict[str, dict[str, Any]] = {
},
"classify": {
"label": "Document Classification",
"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.).",
},
},
"description": "Classify the document type using AI without full metadata extraction.",
"config_schema": {},
},
}
-16
View File
@@ -31,7 +31,6 @@ from pydantic import BaseModel, Field
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.audit_log import get_client_ip
from app.utils.session_manager import (
@@ -152,11 +151,6 @@ async def create_challenge(
displayed to the user. The mobile app scans this QR code and
calls the ``/claim`` endpoint.
"""
if not settings.qr_login_enabled:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
)
ip = get_client_ip(request)
challenge = create_qr_challenge(db, owner_id, ip_address=ip)
@@ -193,11 +187,6 @@ async def poll_challenge_status(
The web UI calls this endpoint every few seconds to check if the
mobile app has scanned the QR code and claimed the challenge.
"""
if not settings.qr_login_enabled:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
)
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")
@@ -217,11 +206,6 @@ async def claim_challenge(
serves as proof that the user authorized this login from their web
session.
"""
if not settings.qr_login_enabled:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
)
ip = get_client_ip(request)
result = claim_qr_challenge(db, body.challenge_token, device_name=body.device_name, ip_address=ip)
-62
View File
@@ -55,12 +55,6 @@ class SettingUpdate(BaseModel):
value: Optional[str] = Field(None, description="Setting value (None to delete)")
class SettingValueUpdate(BaseModel):
"""Model for updating a setting value by key (key is provided in the URL path)."""
value: Optional[str] = Field(None, description="Setting value (None to delete)")
class SettingResponse(BaseModel):
"""Model for setting response"""
@@ -329,62 +323,6 @@ async def update_setting(
)
@router.put("/{key}")
async def put_setting(
key: str,
body: SettingValueUpdate,
request: Request,
db: DbSession,
admin: AdminUser,
):
"""
Update a specific setting by key (RESTful PUT).
Accepts a body with only ``value``; the key is taken from the URL path.
This is the endpoint used by the admin Connections wizard.
Admin only.
"""
validate_setting_key(key)
try:
if body.value is not None:
is_valid, error_message = validate_setting_value(key, body.value)
if not is_valid:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error_message)
user = request.session.get("user", {}) if hasattr(request, "session") else {}
changed_by = (
user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "admin"
)
success = save_setting_to_db(db, key, body.value, changed_by=changed_by)
if not success:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to save setting to database",
)
notify_settings_updated()
metadata = get_setting_metadata(key)
restart_required = metadata.get("restart_required", False)
return {
"success": True,
"message": f"Setting '{key}' updated successfully",
"restart_required": restart_required,
"key": key,
"value": body.value,
}
except HTTPException:
raise
except Exception as e:
logger.error(f"Error updating setting {key}: {e}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Failed to update setting: {key}",
)
@router.delete("/{key}")
async def delete_setting(key: str, request: Request, db: DbSession, admin: AdminUser):
"""
-355
View File
@@ -1,355 +0,0 @@
"""File-sharing API endpoints.
Provides CRUD operations for ``FileShare`` records, which grant named
users ``viewer`` or ``editor`` access to a document owned by someone
else. Only the file owner may create, update, or revoke shares.
"""
import logging
from typing import Annotated, Any
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session
from app.auth import require_login
from app.database import get_db
from app.models import FILE_SHARE_ROLE_VIEWER, FILE_SHARE_ROLES, FileRecord, FileShare, UserProfile
from app.utils.user_scope import get_current_owner_id, get_file_role
logger = logging.getLogger(__name__)
router = APIRouter(tags=["sharing"])
DbSession = Annotated[Session, Depends(get_db)]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _serialize_share(share: FileShare) -> dict[str, Any]:
"""Serialize a ``FileShare`` to a JSON-friendly dict."""
return {
"id": share.id,
"file_id": share.file_id,
"owner_id": share.owner_id,
"shared_with_user_id": share.shared_with_user_id,
"role": share.role,
"created_at": share.created_at.isoformat() if share.created_at else None,
"updated_at": share.updated_at.isoformat() if share.updated_at else None,
}
def _require_owner(file_record: FileRecord, user_id: str | None, db: Session) -> None:
"""Raise 403 unless the calling user is the file owner."""
if get_file_role(file_record, user_id, db) != "owner":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only the file owner can manage shares",
)
# ---------------------------------------------------------------------------
# List shares
# ---------------------------------------------------------------------------
@router.get("/files/{file_id}/shares")
@require_login
def list_shares(request: Request, file_id: int, db: DbSession):
"""List all shares for a document.
Only the file owner (or an admin) may call this endpoint.
Path Parameters:
file_id: The ID of the document.
Returns:
A list of share objects.
"""
user_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
role = get_file_role(file_record, user_id, db)
if role is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
if role != "owner" and not is_admin:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only the file owner can view shares",
)
shares = db.query(FileShare).filter(FileShare.file_id == file_id).all()
return [_serialize_share(s) for s in shares]
# ---------------------------------------------------------------------------
# Create share
# ---------------------------------------------------------------------------
@router.post("/files/{file_id}/shares", status_code=status.HTTP_201_CREATED)
@require_login
def create_share(
request: Request,
file_id: int,
db: DbSession,
shared_with_user_id: str = Body(..., embed=True),
role: str = Body(FILE_SHARE_ROLE_VIEWER, embed=True),
):
"""Share a document with another user.
Only the file owner may share the document. Sharing with a user
that already has access updates their role instead of creating a
duplicate record.
Path Parameters:
file_id: The ID of the document to share.
Request body (JSON):
shared_with_user_id: The stable user identifier of the recipient.
role: ``"viewer"`` (default) or ``"editor"``.
Returns:
The created or updated share object.
"""
owner_id = get_current_owner_id(request)
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
_require_owner(file_record, owner_id, db)
if role not in FILE_SHARE_ROLES:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"role must be one of: {', '.join(FILE_SHARE_ROLES)}",
)
if not shared_with_user_id or not shared_with_user_id.strip():
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="shared_with_user_id must be a non-empty string",
)
shared_with_user_id = shared_with_user_id.strip()
# Cannot share with yourself
if shared_with_user_id == owner_id:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="You cannot share a file with yourself",
)
try:
existing = (
db.query(FileShare)
.filter(FileShare.file_id == file_id, FileShare.shared_with_user_id == shared_with_user_id)
.first()
)
if existing:
# Update role if different
if existing.role != role:
existing.role = role
db.commit()
db.refresh(existing)
logger.info(
"Share updated: file_id=%s, shared_with=%s, role=%s, by owner=%s",
file_id,
shared_with_user_id,
role,
owner_id,
)
return _serialize_share(existing)
share = FileShare(
file_id=file_id,
owner_id=owner_id,
shared_with_user_id=shared_with_user_id,
role=role,
)
db.add(share)
db.commit()
db.refresh(share)
except HTTPException:
raise
except Exception:
db.rollback()
logger.exception("Failed to create share: file_id=%s, shared_with=%s", file_id, shared_with_user_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to create share",
)
logger.info(
"Share created: id=%s, file_id=%s, shared_with=%s, role=%s, by owner=%s",
share.id,
file_id,
shared_with_user_id,
role,
owner_id,
)
return _serialize_share(share)
# ---------------------------------------------------------------------------
# Update share role
# ---------------------------------------------------------------------------
@router.put("/files/{file_id}/shares/{share_id}")
@require_login
def update_share(
request: Request,
file_id: int,
share_id: int,
db: DbSession,
role: str = Body(..., embed=True),
):
"""Update the role of an existing share.
Only the file owner may change the role of a share.
Path Parameters:
file_id: The ID of the document.
share_id: The ID of the share record to update.
Request body (JSON):
role: New role — ``"viewer"`` or ``"editor"``.
Returns:
The updated share object.
"""
owner_id = get_current_owner_id(request)
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
_require_owner(file_record, owner_id, db)
if role not in FILE_SHARE_ROLES:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"role must be one of: {', '.join(FILE_SHARE_ROLES)}",
)
share = db.query(FileShare).filter(FileShare.id == share_id, FileShare.file_id == file_id).first()
if not share:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Share not found")
try:
share.role = role
db.commit()
db.refresh(share)
except Exception:
db.rollback()
logger.exception("Failed to update share: share_id=%s", share_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to update share",
)
logger.info("Share updated: id=%s, file_id=%s, new_role=%s, by owner=%s", share_id, file_id, role, owner_id)
return _serialize_share(share)
# ---------------------------------------------------------------------------
# Revoke share
# ---------------------------------------------------------------------------
@router.delete("/files/{file_id}/shares/{share_id}", status_code=status.HTTP_200_OK)
@require_login
def revoke_share(request: Request, file_id: int, share_id: int, db: DbSession):
"""Revoke a share, removing the user's access.
Only the file owner may revoke shares.
Path Parameters:
file_id: The ID of the document.
share_id: The ID of the share record to delete.
Returns:
A success message.
"""
owner_id = get_current_owner_id(request)
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
_require_owner(file_record, owner_id, db)
share = db.query(FileShare).filter(FileShare.id == share_id, FileShare.file_id == file_id).first()
if not share:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Share not found")
try:
db.delete(share)
db.commit()
except Exception:
db.rollback()
logger.exception("Failed to revoke share: share_id=%s", share_id)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to revoke share",
)
logger.info("Share revoked: id=%s, file_id=%s, by owner=%s", share_id, file_id, owner_id)
return {"status": "success", "message": "Share revoked successfully"}
# ---------------------------------------------------------------------------
# List users that the file is already shared with (for the share-picker UI)
# ---------------------------------------------------------------------------
@router.get("/files/{file_id}/shared-with")
@require_login
def list_shared_with(request: Request, file_id: int, db: DbSession):
"""Return the list of users a document is shared with and their roles.
Accessible to any user that has at least viewer access to the file,
so that editors/viewers can see who else has access.
Path Parameters:
file_id: The ID of the document.
Returns:
A list of ``{share_id, user_id, display_name, role}`` objects.
"""
user_id = get_current_owner_id(request)
user = request.session.get("user")
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
role = get_file_role(file_record, user_id, db)
if role is None and not is_admin:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
shares = db.query(FileShare).filter(FileShare.file_id == file_id).all()
results = []
for s in shares:
profile = db.query(UserProfile).filter(UserProfile.user_id == s.shared_with_user_id).first()
results.append(
{
"share_id": s.id,
"user_id": s.shared_with_user_id,
"display_name": (profile.display_name if profile and profile.display_name else s.shared_with_user_id),
"role": s.role,
}
)
return results
+2 -7
View File
@@ -11,12 +11,11 @@ from typing import Optional
import aiofiles
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi import APIRouter, 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
@@ -108,11 +107,7 @@ 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,
_rate_ok: None = Depends(require_upload_rate_limit),
):
async def process_url(request: Request, url_request: URLUploadRequest):
"""
Download a file from a URL and enqueue it for processing.
+69 -284
View File
@@ -45,264 +45,78 @@ OAUTH_PROVIDER_NAME = "Single Sign-On"
# Social login providers that are enabled and registered
SOCIAL_PROVIDERS: dict[str, dict[str, str]] = {}
if AUTH_ENABLED and settings.authentik_client_id and settings.authentik_client_secret:
oauth.register(
name="authentik",
client_id=settings.authentik_client_id,
client_secret=settings.authentik_client_secret,
server_metadata_url=settings.authentik_config_url,
client_kwargs={"scope": "openid profile email"},
)
OAUTH_CONFIGURED = True
OAUTH_PROVIDER_NAME = settings.oauth_provider_name or "Authentik SSO"
# ---------------------------------------------------------------------------
# Helpers for dynamic (re-)registration of OAuth providers
# ---------------------------------------------------------------------------
def _register_oauth_client(name: str, **kwargs: object) -> None:
"""Register (or re-register) an authlib OAuth client, clearing any cached instance.
authlib caches the constructed client object in ``oauth._clients`` after the
first ``register()`` call. Subsequent ``register()`` calls overwrite the
registry entry but the stale cached client is still returned by
``create_client()`` / ``__getattr__``. Popping the name from ``_clients``
before re-registering ensures the new credentials are picked up immediately.
Args:
name: Provider name (e.g. ``"google"``, ``"github"``).
**kwargs: Keyword arguments forwarded verbatim to ``oauth.register()``.
"""
oauth._clients.pop(name, None)
oauth.register(name, **kwargs)
def _dropbox_userinfo_compliance_fix(client, user_cls, token, data):
"""Normalize Dropbox userinfo response for authlib compatibility.
Dropbox's /2/users/get_current_account returns a non-standard response
format. This compliance fix normalizes the response data — the HTTP
method (POST) is handled by authlib's compliance infrastructure.
Args:
client: The OAuth client instance (required by authlib compliance fix interface).
user_cls: The user class (required by authlib compliance fix interface).
token: The OAuth token dict.
data: The raw userinfo response dict from Dropbox.
Returns:
The normalized userinfo dict with ``sub`` and ``name`` fields.
"""
# Dropbox returns account_id instead of sub
if "account_id" in data and "sub" not in data:
data["sub"] = data["account_id"]
# Normalize name field
name_info = data.get("name", {})
if isinstance(name_info, dict) and "display_name" in name_info:
data["name"] = name_info["display_name"]
return data
def _setup_social_providers() -> None:
"""Register all configured OAuth / social-login providers from current settings.
This function is **idempotent**: it clears ``SOCIAL_PROVIDERS``,
``OAUTH_CONFIGURED``, and ``OAUTH_PROVIDER_NAME`` before rebuilding them,
and calls :func:`_register_oauth_client` (which also clears the authlib
client cache) so that credential changes in the database are reflected
without an application restart.
Can safely be called multiple times, e.g. after a settings reload.
"""
global OAUTH_CONFIGURED, OAUTH_PROVIDER_NAME
SOCIAL_PROVIDERS.clear()
OAUTH_CONFIGURED = False
OAUTH_PROVIDER_NAME = "Single Sign-On"
if not AUTH_ENABLED:
return
# --- Authentik / OIDC ---
if settings.authentik_client_id and settings.authentik_client_secret:
_register_oauth_client(
"authentik",
client_id=settings.authentik_client_id,
client_secret=settings.authentik_client_secret,
server_metadata_url=settings.authentik_config_url,
# --- Social Login Providers ---------------------------------------------------
if AUTH_ENABLED and settings.social_auth_google_enabled:
if settings.social_auth_google_client_id and settings.social_auth_google_client_secret:
oauth.register(
name="google",
client_id=settings.social_auth_google_client_id,
client_secret=settings.social_auth_google_client_secret,
server_metadata_url="https://accounts.google.com/.well-known/openid-configuration",
client_kwargs={"scope": "openid profile email"},
)
OAUTH_CONFIGURED = True
OAUTH_PROVIDER_NAME = settings.oauth_provider_name or "Authentik SSO"
SOCIAL_PROVIDERS["google"] = {"name": "Google", "icon": "fab fa-google", "color": "red"}
logger.info("Social login provider registered: Google")
else:
logger.warning("SOCIAL_AUTH_GOOGLE_ENABLED=true but client ID/secret not configured")
# --- Social Login Providers ---
if AUTH_ENABLED and settings.social_auth_microsoft_enabled:
if settings.social_auth_microsoft_client_id and settings.social_auth_microsoft_client_secret:
tenant = settings.social_auth_microsoft_tenant or "common"
oauth.register(
name="microsoft",
client_id=settings.social_auth_microsoft_client_id,
client_secret=settings.social_auth_microsoft_client_secret,
server_metadata_url=f"https://login.microsoftonline.com/{tenant}/v2.0/.well-known/openid-configuration",
client_kwargs={"scope": "openid profile email"},
)
SOCIAL_PROVIDERS["microsoft"] = {"name": "Microsoft", "icon": "fab fa-microsoft", "color": "blue"}
logger.info("Social login provider registered: Microsoft (tenant=%s)", tenant)
else:
logger.warning("SOCIAL_AUTH_MICROSOFT_ENABLED=true but client ID/secret not configured")
# Google
if settings.social_auth_google_enabled:
_google_client_id = settings.social_auth_google_client_id
_google_client_secret = settings.social_auth_google_client_secret
if settings.social_auth_google_use_global_credentials and not (_google_client_id and _google_client_secret):
_google_client_id = settings.google_drive_client_id
_google_client_secret = settings.google_drive_client_secret
if AUTH_ENABLED and settings.social_auth_apple_enabled:
if settings.social_auth_apple_client_id and settings.social_auth_apple_team_id:
oauth.register(
name="apple",
client_id=settings.social_auth_apple_client_id,
server_metadata_url="https://appleid.apple.com/.well-known/openid-configuration",
client_kwargs={
"scope": "openid name email",
"response_mode": "form_post",
},
)
SOCIAL_PROVIDERS["apple"] = {"name": "Apple", "icon": "fab fa-apple", "color": "gray"}
logger.info("Social login provider registered: Apple")
else:
logger.warning("SOCIAL_AUTH_APPLE_ENABLED=true but client ID/team ID not configured")
if _google_client_id and _google_client_secret:
_register_oauth_client(
"google",
client_id=_google_client_id,
client_secret=_google_client_secret,
server_metadata_url="https://accounts.google.com/.well-known/openid-configuration",
client_kwargs={"scope": "openid profile email"},
)
SOCIAL_PROVIDERS["google"] = {"name": "Google", "icon": "fab fa-google", "color": "red"}
logger.info("Social login provider registered: Google")
else:
logger.warning("SOCIAL_AUTH_GOOGLE_ENABLED=true but client ID/secret not configured")
# Microsoft
if settings.social_auth_microsoft_enabled:
_microsoft_client_id = settings.social_auth_microsoft_client_id
_microsoft_client_secret = settings.social_auth_microsoft_client_secret
if settings.social_auth_microsoft_use_global_credentials and not (
_microsoft_client_id and _microsoft_client_secret
):
_microsoft_client_id = settings.onedrive_client_id
_microsoft_client_secret = settings.onedrive_client_secret
if _microsoft_client_id and _microsoft_client_secret:
tenant = settings.social_auth_microsoft_tenant or "common"
_register_oauth_client(
"microsoft",
client_id=_microsoft_client_id,
client_secret=_microsoft_client_secret,
server_metadata_url=f"https://login.microsoftonline.com/{tenant}/v2.0/.well-known/openid-configuration",
client_kwargs={"scope": "openid profile email"},
)
SOCIAL_PROVIDERS["microsoft"] = {"name": "Microsoft", "icon": "fab fa-microsoft", "color": "blue"}
logger.info("Social login provider registered: Microsoft (tenant=%s)", tenant)
else:
logger.warning("SOCIAL_AUTH_MICROSOFT_ENABLED=true but client ID/secret not configured")
# Apple
if settings.social_auth_apple_enabled:
if settings.social_auth_apple_client_id and settings.social_auth_apple_team_id:
_register_oauth_client(
"apple",
client_id=settings.social_auth_apple_client_id,
server_metadata_url="https://appleid.apple.com/.well-known/openid-configuration",
client_kwargs={
"scope": "openid name email",
"response_mode": "form_post",
},
)
SOCIAL_PROVIDERS["apple"] = {"name": "Apple", "icon": "fab fa-apple", "color": "gray"}
logger.info("Social login provider registered: Apple")
else:
logger.warning("SOCIAL_AUTH_APPLE_ENABLED=true but client ID/team ID not configured")
# Dropbox
if settings.social_auth_dropbox_enabled:
_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 and _dropbox_client_secret):
_dropbox_client_id = settings.dropbox_app_key
_dropbox_client_secret = settings.dropbox_app_secret
if _dropbox_client_id and _dropbox_client_secret:
_register_oauth_client(
"dropbox",
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",
userinfo_compliance_fix=_dropbox_userinfo_compliance_fix,
client_kwargs={
"token_endpoint_auth_method": "client_secret_post",
"token_access_type": "offline",
},
)
SOCIAL_PROVIDERS["dropbox"] = {"name": "Dropbox", "icon": "fab fa-dropbox", "color": "blue"}
logger.info("Social login provider registered: Dropbox")
else:
logger.warning("SOCIAL_AUTH_DROPBOX_ENABLED=true but client ID/secret not configured")
# GitHub
if settings.social_auth_github_enabled:
if settings.social_auth_github_client_id and settings.social_auth_github_client_secret:
_register_oauth_client(
"github",
client_id=settings.social_auth_github_client_id,
client_secret=settings.social_auth_github_client_secret,
authorize_url="https://github.com/login/oauth/authorize",
access_token_url="https://github.com/login/oauth/access_token",
userinfo_endpoint="https://api.github.com/user",
client_kwargs={"scope": "read:user user:email"},
)
SOCIAL_PROVIDERS["github"] = {"name": "GitHub", "icon": "fab fa-github", "color": "gray"}
logger.info("Social login provider registered: GitHub")
else:
logger.warning("SOCIAL_AUTH_GITHUB_ENABLED=true but client ID/secret not configured")
# Keycloak
if settings.social_auth_keycloak_enabled:
_kc_server = settings.social_auth_keycloak_server_url
_kc_realm = settings.social_auth_keycloak_realm
if (
settings.social_auth_keycloak_client_id
and settings.social_auth_keycloak_client_secret
and _kc_server
and _kc_realm
):
_kc_base = f"{_kc_server.rstrip('/')}/realms/{_kc_realm}"
_register_oauth_client(
"keycloak",
client_id=settings.social_auth_keycloak_client_id,
client_secret=settings.social_auth_keycloak_client_secret,
server_metadata_url=f"{_kc_base}/.well-known/openid-configuration",
client_kwargs={"scope": "openid profile email"},
)
SOCIAL_PROVIDERS["keycloak"] = {"name": "Keycloak", "icon": "fas fa-key", "color": "gray"}
logger.info("Social login provider registered: Keycloak (realm=%s)", _kc_realm)
else:
logger.warning("SOCIAL_AUTH_KEYCLOAK_ENABLED=true but required settings not configured")
# Generic OAuth2
if settings.social_auth_generic_oauth2_enabled:
if (
settings.social_auth_generic_oauth2_client_id
and settings.social_auth_generic_oauth2_client_secret
and settings.social_auth_generic_oauth2_authorize_url
and settings.social_auth_generic_oauth2_token_url
):
_register_oauth_client(
"generic_oauth2",
client_id=settings.social_auth_generic_oauth2_client_id,
client_secret=settings.social_auth_generic_oauth2_client_secret,
authorize_url=settings.social_auth_generic_oauth2_authorize_url,
access_token_url=settings.social_auth_generic_oauth2_token_url,
userinfo_endpoint=settings.social_auth_generic_oauth2_userinfo_url,
client_kwargs={"scope": settings.social_auth_generic_oauth2_scope},
)
_generic_name = settings.social_auth_generic_oauth2_name or "OAuth2"
SOCIAL_PROVIDERS["generic_oauth2"] = {
"name": _generic_name,
"icon": "fas fa-sign-in-alt",
"color": "indigo",
}
logger.info("Social login provider registered: Generic OAuth2 (%s)", _generic_name)
else:
logger.warning("SOCIAL_AUTH_GENERIC_OAUTH2_ENABLED=true but required settings not configured")
def refresh_social_providers() -> None:
"""Re-register all OAuth providers from the *current* settings object.
Call this after loading or reloading settings from the database so that
providers configured (or updated) through the admin UI take effect
immediately — **no application restart required**.
This function is safe to call multiple times and is idempotent.
"""
logger.info("Refreshing social login provider registrations from current settings")
_setup_social_providers()
# Perform the initial registration from environment / default settings at
# import time. The lifespan hook and settings_sync will call
# refresh_social_providers() again after DB settings are loaded so that
# any providers configured only in the database are also active.
_setup_social_providers()
if AUTH_ENABLED and settings.social_auth_dropbox_enabled:
if settings.social_auth_dropbox_client_id and settings.social_auth_dropbox_client_secret:
oauth.register(
name="dropbox",
client_id=settings.social_auth_dropbox_client_id,
client_secret=settings.social_auth_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",
client_kwargs={"token_endpoint_auth_method": "client_secret_post"},
)
SOCIAL_PROVIDERS["dropbox"] = {"name": "Dropbox", "icon": "fab fa-dropbox", "color": "blue"}
logger.info("Social login provider registered: Dropbox")
else:
logger.warning("SOCIAL_AUTH_DROPBOX_ENABLED=true but client ID/secret not configured")
router = APIRouter()
@@ -372,16 +186,6 @@ 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,
@@ -527,21 +331,13 @@ async def login(request: Request):
get_client_ip(request),
)
error = request.query_params.get("error")
message = request.query_params.get("message")
show_oauth = OAUTH_CONFIGURED
# SSO Auto Login: redirect directly to SSO provider if configured
if show_oauth and settings.sso_auto_login is True and not error and not message:
return RedirectResponse(url="/oauth-login", status_code=status.HTTP_302_FOUND)
return templates.TemplateResponse(
"login.html",
{
"request": request,
"error": error,
"message": message,
"show_oauth": show_oauth,
"error": request.query_params.get("error"),
"message": request.query_params.get("message"),
"show_oauth": OAUTH_CONFIGURED,
"oauth_provider_name": OAUTH_PROVIDER_NAME,
"social_providers": SOCIAL_PROVIDERS,
"app_version": settings.version,
@@ -627,17 +423,6 @@ def _normalize_social_userinfo(provider: str, token: dict, raw_userinfo: dict |
"picture": userinfo.get("profile_photo_url", ""),
}
if provider == "github":
# GitHub returns login, id, name, email, avatar_url
email = userinfo.get("email", "")
return {
"sub": str(userinfo.get("id", "")),
"email": email,
"name": userinfo.get("name", "") or userinfo.get("login", ""),
"preferred_username": userinfo.get("login", email),
"picture": userinfo.get("avatar_url", ""),
}
# Standard OIDC providers (Google, Microsoft, Apple)
return {
"sub": userinfo.get("sub", ""),
-2
View File
@@ -10,7 +10,6 @@ from app import tasks # noqa: F401 - Imports app/tasks.py so Celery can registe
# Import the shared Celery instance
from app.celery_app import celery
from app.config import settings
from app.tasks.automation_tasks import deliver_automation_hook_task # noqa: F401
from app.tasks.backup_tasks import cleanup_old_backups, create_backup # noqa: F401
from app.tasks.batch_tasks import ( # noqa: F401
backfill_missing_metadata,
@@ -23,7 +22,6 @@ 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
+30 -161
View File
@@ -13,24 +13,6 @@ 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
@@ -120,16 +102,6 @@ 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(
@@ -193,16 +165,6 @@ 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
@@ -244,10 +206,6 @@ class Settings(BaseSettings):
"Useful for admin-configured non-standard durations."
),
)
qr_login_enabled: bool = Field(
default=True,
description="Enable QR code-based login for mobile device authentication (default: True).",
)
qr_login_challenge_ttl_seconds: int = Field(
default=120,
description="Time-to-live in seconds for QR login challenges (default: 2 minutes).",
@@ -309,55 +267,12 @@ class Settings(BaseSettings):
authentik_client_secret: Optional[str] = None
authentik_config_url: Optional[str] = None
oauth_provider_name: Optional[str] = None # Name to display for the OAuth provider
sso_auto_login: bool = Field(
default=False,
description=(
"Automatically redirect to SSO login when authentication is required. "
"When enabled, users are sent directly to the SSO provider instead of "
"seeing the login page. Only effective when OIDC is configured."
),
)
# Keycloak SSO
social_auth_keycloak_enabled: bool = False
social_auth_keycloak_client_id: Optional[str] = None
social_auth_keycloak_client_secret: Optional[str] = None
social_auth_keycloak_server_url: Optional[str] = None
social_auth_keycloak_realm: Optional[str] = None
# Generic OAuth2 SSO
social_auth_generic_oauth2_enabled: bool = False
social_auth_generic_oauth2_client_id: Optional[str] = None
social_auth_generic_oauth2_client_secret: Optional[str] = None
social_auth_generic_oauth2_authorize_url: Optional[str] = None
social_auth_generic_oauth2_token_url: Optional[str] = None
social_auth_generic_oauth2_userinfo_url: Optional[str] = None
social_auth_generic_oauth2_scope: str = "openid profile email"
social_auth_generic_oauth2_name: str = "OAuth2"
# SAML2 SSO
social_auth_saml2_enabled: bool = False
social_auth_saml2_entity_id: Optional[str] = None
social_auth_saml2_sso_url: Optional[str] = None
social_auth_saml2_certificate: Optional[str] = None
social_auth_saml2_name: str = "SAML2"
# Social Login Providers
# Google OAuth2
social_auth_google_enabled: bool = False
social_auth_google_client_id: Optional[str] = None
social_auth_google_client_secret: Optional[str] = None
social_auth_google_use_global_credentials: bool = Field(
default=False,
description=(
"When True, Google social login uses the global GOOGLE_DRIVE_CLIENT_ID / "
"GOOGLE_DRIVE_CLIENT_SECRET credentials (the Google Drive OAuth integration credentials) "
"instead of requiring separate SOCIAL_AUTH_GOOGLE_CLIENT_ID / "
"SOCIAL_AUTH_GOOGLE_CLIENT_SECRET values. "
"Requires SOCIAL_AUTH_GOOGLE_ENABLED=True and the global Google Drive OAuth credentials to be set. "
"Default: False."
),
)
# Microsoft OAuth2 (Azure AD / Microsoft Entra ID)
social_auth_microsoft_enabled: bool = False
@@ -372,17 +287,6 @@ class Settings(BaseSettings):
"Default: common."
),
)
social_auth_microsoft_use_global_credentials: bool = Field(
default=False,
description=(
"When True, Microsoft social login uses the global ONEDRIVE_CLIENT_ID / "
"ONEDRIVE_CLIENT_SECRET credentials (the OneDrive integration credentials) "
"instead of requiring separate SOCIAL_AUTH_MICROSOFT_CLIENT_ID / "
"SOCIAL_AUTH_MICROSOFT_CLIENT_SECRET values. "
"Requires SOCIAL_AUTH_MICROSOFT_ENABLED=True and the global OneDrive credentials to be set. "
"Default: False."
),
)
# Apple Sign-In
social_auth_apple_enabled: bool = False
@@ -395,21 +299,6 @@ 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."
),
)
# GitHub OAuth2
social_auth_github_enabled: bool = False
social_auth_github_client_id: Optional[str] = None
social_auth_github_client_secret: Optional[str] = None
# Local user signup
allow_local_signup: bool = Field(
@@ -907,11 +796,6 @@ class Settings(BaseSettings):
),
)
# Telegram Bot
telegram_bot_token: Optional[str] = None
telegram_chat_id: Optional[str] = None
telegram_enabled: bool = False
# Notification settings
notification_urls: Union[List[str], str] = Field(
default_factory=list,
@@ -946,12 +830,6 @@ class Settings(BaseSettings):
description="Enable webhook delivery for document events",
)
# Automation hooks (Zapier / Make.com)
automation_hooks_enabled: bool = Field(
default=True,
description="Enable Zapier / Make.com automation hook subscriptions and delivery",
)
# ── Backup / restore settings ──────────────────────────────────────────────
backup_enabled: bool = Field(
default=True,
@@ -1237,18 +1115,43 @@ class Settings(BaseSettings):
),
)
# Per-user upload rate limiting (health-aware, Redis-backed sliding window)
# Database Connection Pool Configuration
# Controls SQLAlchemy QueuePool behaviour for PostgreSQL/MySQL.
# SQLite uses NullPool and ignores these settings.
db_pool_size: int = Field(
default=5,
description="Number of persistent connections kept in the pool. Ignored for SQLite.",
)
db_max_overflow: int = Field(
default=10,
description=("Maximum number of connections that can be opened beyond db_pool_size. Ignored for SQLite."),
)
db_pool_timeout: int = Field(
default=30,
description="Seconds to wait for a connection from the pool before raising an error. Ignored for SQLite.",
)
db_pool_recycle: int = Field(
default=1800,
description=(
"Seconds after which a connection is recycled to prevent stale connections. "
"Ignored for SQLite. Default: 1800 (30 minutes)."
),
)
# Per-user upload rate limiting (health-aware limiter)
# Controls how many uploads a single user may submit within a 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."
"Maximum number of uploads allowed per user within the upload_rate_limit_window. "
"The limiter may dynamically reduce this value when Redis queue depth or CPU load is high."
),
)
upload_rate_limit_window: int = Field(
default=60,
description="Sliding window size in seconds for per-user upload rate limiting (default: 60).",
description=(
"Sliding window in seconds over which upload_rate_limit_per_user is enforced. Default: 60 seconds."
),
)
# Rate Limiting Configuration (see SECURITY_AUDIT.md and docs/API.md)
@@ -1377,40 +1280,6 @@ class Settings(BaseSettings):
),
)
# ---------------------------------------------------------------------------
# Observability Sentry Browser JavaScript SDK (client-side)
# ---------------------------------------------------------------------------
# The same SENTRY_DSN is reused for the browser SDK. The DSN is a *public*
# key in Sentry's model and is intentionally embedded in client-side code.
# All three settings below default to 0.0 / disabled so that operators opt-in
# to the level of browser monitoring they want.
# ---------------------------------------------------------------------------
sentry_js_traces_sample_rate: float = Field(
default=0.0,
description=(
"Fraction of browser page-loads captured for client-side performance tracing "
"(0.0 1.0). 0.0 disables browser tracing; 1.0 captures every navigation. "
"Only active when SENTRY_DSN is set."
),
)
sentry_js_replay_session_sample_rate: float = Field(
default=0.0,
description=(
"Fraction of sessions recorded by Sentry Session Replay (0.0 1.0). "
"0.0 disables session recording; 1.0 records every session. "
"Only active when SENTRY_DSN is set."
),
)
sentry_js_replay_on_error_sample_rate: float = Field(
default=0.1,
description=(
"Fraction of sessions with an error that will be recorded by Sentry Session "
"Replay (0.0 1.0). Defaults to 0.1 (10 %) so that errors are captured "
"with replay context even when session-level recording is disabled. "
"Only active when SENTRY_DSN is set."
),
)
@model_validator(mode="before")
@classmethod
def strip_outer_quotes(cls, data: Any) -> Any:
+13 -28
View File
@@ -18,37 +18,22 @@ logger = logging.getLogger(__name__)
Base = declarative_base()
# ---------------------------------------------------------------------------
# Engine construction
# ---------------------------------------------------------------------------
# Parse the DATABASE_URL
DB_URL = settings.database_url
_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
_db_url = make_url(DB_URL)
if _db_url.get_backend_name() == "sqlite":
# SQLite does not benefit from connection pooling; NullPool avoids contention.
engine = create_engine(DB_URL, connect_args={"check_same_thread": False}, 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,
}
# PostgreSQL / MySQL / other: use a configurable QueuePool.
engine = create_engine(
DB_URL,
poolclass=QueuePool,
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)
-12
View File
@@ -189,18 +189,6 @@ async def lifespan(app: FastAPI):
finally:
db.close()
# Re-register OAuth / social-login providers now that DB settings are
# loaded. auth.py runs its initial registration at import time (before
# the lifespan runs), so providers that are only configured in the
# database would not be registered yet. Calling refresh here ensures
# they are active immediately on startup without any manual restart.
try:
from app.auth import refresh_social_providers
refresh_social_providers()
except Exception as e:
logging.warning(f"Could not refresh social login providers on startup: {e}")
# Initialize Sentry after DB settings are loaded so that values configured
# via the database UI (e.g. SENTRY_DSN) are respected in addition to env vars.
init_sentry()
-290
View File
@@ -1,290 +0,0 @@
"""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,
)
-164
View File
@@ -211,28 +211,6 @@ class WebhookConfig(Base):
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
class AutomationHook(Base):
"""Zapier / Make.com compatible webhook subscription for automation triggers.
External automation platforms subscribe to DocuElevate events via the REST
hooks protocol. When an event fires, DocuElevate POSTs a Zapier-compatible
flat JSON payload to ``target_url``. The ``hook_type`` field records which
platform created the subscription (informational only).
"""
__tablename__ = "automation_hooks"
id = Column(Integer, primary_key=True, index=True)
target_url = Column(String, nullable=False) # URL to POST events to
secret = Column(String, nullable=True) # Optional HMAC-SHA256 signing secret
events = Column(Text, nullable=False) # JSON list of subscribed event names
is_active = Column(Boolean, default=True, nullable=False)
hook_type = Column(String(50), nullable=False, default="generic") # zapier | make | generic
description = Column(String, nullable=True) # Optional human-readable label
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
class LocalUser(Base):
"""A locally-registered user authenticated by email and bcrypt password.
@@ -808,9 +786,6 @@ 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.
@@ -962,51 +937,6 @@ 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.
@@ -1192,97 +1122,3 @@ class PipelineRoutingRule(Base):
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
class DocumentComment(Base):
"""Threaded comment on a document.
Supports threaded replies via ``parent_id`` and @mentions via the
``mentions`` column (comma-separated user identifiers).
"""
__tablename__ = "document_comments"
id = Column(Integer, primary_key=True, index=True)
file_id = Column(Integer, ForeignKey(_FILES_ID_FK), nullable=False, index=True)
user_id = Column(String, nullable=False, index=True)
parent_id = Column(Integer, ForeignKey("document_comments.id"), nullable=True, index=True)
body = Column(Text, nullable=False)
mentions = Column(Text, nullable=True)
is_resolved = Column(Boolean, nullable=False, default=False, server_default="0")
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
class DocumentAnnotation(Base):
"""Text annotation on a specific page and position of a PDF document.
Stores the bounding-box coordinates (``x``, ``y``, ``width``,
``height``) relative to the page dimensions so that the annotation
can be rendered on top of the PDF viewer.
"""
__tablename__ = "document_annotations"
id = Column(Integer, primary_key=True, index=True)
file_id = Column(Integer, ForeignKey(_FILES_ID_FK), nullable=False, index=True)
user_id = Column(String, nullable=False, index=True)
page = Column(Integer, nullable=False)
x = Column(Float, nullable=False)
y = Column(Float, nullable=False)
width = Column(Float, nullable=False, default=0)
height = Column(Float, nullable=False, default=0)
content = Column(Text, nullable=False)
annotation_type = Column(String(50), nullable=False, default="note", server_default="note")
color = Column(String(20), nullable=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
# ---------------------------------------------------------------------------
# File sharing
# ---------------------------------------------------------------------------
# Valid roles for FileShare.role
FILE_SHARE_ROLE_VIEWER = "viewer"
FILE_SHARE_ROLE_EDITOR = "editor"
FILE_SHARE_ROLES = (FILE_SHARE_ROLE_VIEWER, FILE_SHARE_ROLE_EDITOR)
class FileShare(Base):
"""Grants a named user access to a ``FileRecord`` owned by someone else.
The ``owner_id`` column records who created the share (must be the file
owner). ``shared_with_user_id`` is the recipient's stable user
identifier (the same kind of string used in ``FileRecord.owner_id``).
Roles
-----
``viewer`` — can read the file, comments, and annotations; may add
comments/annotations; cannot delete or share.
``editor`` — all viewer rights plus the ability to edit document
metadata; cannot delete or re-share.
Only the file owner may create, update, or revoke shares.
"""
__tablename__ = "file_shares"
id = Column(Integer, primary_key=True, index=True)
# The document being shared.
file_id = Column(Integer, ForeignKey(_FILES_ID_FK), nullable=False, index=True)
# The user who granted the share (must match FileRecord.owner_id).
owner_id = Column(String, nullable=False, index=True)
# The user receiving the share.
shared_with_user_id = Column(String, nullable=False, index=True)
# "viewer" or "editor"
role = Column(String(20), nullable=False, default=FILE_SHARE_ROLE_VIEWER)
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("file_id", "shared_with_user_id", name="uq_file_share_file_user"),)
-44
View File
@@ -1,44 +0,0 @@
"""Celery task for asynchronous automation hook delivery with retry and backoff.
Uses :class:`~app.tasks.retry_config.BaseTaskWithRetry` so failed deliveries
are automatically retried with exponential backoff (default: 60 s, 300 s,
900 s) and ±20 % jitter.
"""
import logging
from typing import Any
from app.celery_app import celery
from app.tasks.retry_config import BaseTaskWithRetry
from app.utils.webhook import deliver_webhook
logger = logging.getLogger(__name__)
@celery.task(base=BaseTaskWithRetry, bind=True, name="automation.deliver_hook")
def deliver_automation_hook_task(self, url: str, payload: dict[str, Any], secret: str | None = None) -> dict[str, Any]:
"""Deliver an automation hook payload to *url* with automatic retries.
Args:
url: Target webhook URL (provided by Zapier / Make.com).
payload: The flat Zapier-compatible payload.
secret: Optional shared secret for HMAC-SHA256 signing.
Returns:
A dict with ``status`` and ``url`` on success.
Raises:
RuntimeError: Re-raised to trigger Celery retry on delivery failure.
"""
logger.info(
"Delivering automation hook to %s (attempt %d/%d)",
url,
self.request.retries + 1,
self.max_retries + 1,
)
success = deliver_webhook(url, payload, secret)
if success:
return {"status": "delivered", "url": url}
raise RuntimeError(f"Automation hook delivery to {url} failed")
-174
View File
@@ -1,174 +0,0 @@
"""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
+1 -1
View File
@@ -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", ".heic", ".heif"}
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".tif", ".webp", ".svg"}
HTML_EXTENSIONS = {".html", ".htm"}
-1
View File
@@ -78,7 +78,6 @@ def _convert_pdf_to_pdfa(input_path: str, output_path: str, pdfa_format: str = "
output_type,
"--quiet",
"--invalidate-digital-signatures",
"--",
input_path,
output_path,
]
+6
View File
@@ -18,6 +18,7 @@ from app.utils.allowed_types import (
DEFAULT_CATEGORIES,
get_allowed_types_for_categories,
)
from app.utils.network import is_private_ip
# Database session for per-user IMAP accounts (imported lazily to avoid circular imports)
_db_session_factory = None
@@ -405,6 +406,11 @@ def pull_inbox(
)
processed_emails = load_processed_emails()
# Security: Prevent SSRF by blocking connections to internal IPs
if is_private_ip(host):
logger.warning("SSRF blocked: Attempt to pull mailbox from private IP %s", host)
return
try:
mail = imaplib.IMAP4_SSL(host, port) if use_ssl else imaplib.IMAP4(host, port)
mail.login(username, password)
+1 -2
View File
@@ -555,9 +555,8 @@ def _upload_rclone(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], t
dest = dest.replace("//", "/")
try:
# SECURITY: Separate options from positional arguments using -- to prevent command injection
result = subprocess.run( # nosec B603 # noqa: S603 S607
["rclone", "copyto", f"--config={conf_path}", "--", file_path, dest], # noqa: S603 S607
["rclone", "copyto", f"--config={conf_path}", file_path, dest], # noqa: S603 S607
capture_output=True,
text=True,
timeout=300,
+1 -9
View File
@@ -68,8 +68,6 @@ IMAGE_MIME_TYPES: set[str] = {
"image/tiff",
"image/webp",
"image/svg+xml",
"image/heic",
"image/heif",
}
# ---------------------------------------------------------------------------
@@ -126,8 +124,6 @@ ALLOWED_EXTENSIONS: set[str] = {
".tif",
".webp",
".svg",
".heic",
".heif",
# Web
".html",
".htm",
@@ -238,7 +234,7 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
},
"images": {
"label": "Images",
"description": "Image files (.jpg, .png, .gif, .bmp, .tiff, .webp, .svg, .heic, .heif)",
"description": "Image files (.jpg, .png, .gif, .bmp, .tiff, .webp, .svg)",
"mime_types": frozenset(
{
"image/jpeg",
@@ -249,8 +245,6 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
"image/tiff",
"image/webp",
"image/svg+xml",
"image/heic",
"image/heif",
}
),
"extensions": frozenset(
@@ -264,8 +258,6 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
".tif",
".webp",
".svg",
".heic",
".heif",
}
),
},
-188
View File
@@ -1,188 +0,0 @@
"""Automation hook utilities for Zapier / Make.com integration.
Provides helpers to build Zapier-compatible flat payloads, query active
automation hook subscriptions, and fan-out event delivery to all matching
hooks via Celery tasks.
The payload format is intentionally *flat* (no nested ``data`` key) so that
Zapier and Make.com can map fields without JSONPath expressions. An ``id``
field is included for Zapier deduplication.
"""
import json
import logging
import time
import uuid
from typing import Any
from app.config import settings
from app.database import SessionLocal
from app.models import AutomationHook
from app.utils.webhook import VALID_EVENTS
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Payload helpers
# ---------------------------------------------------------------------------
def build_zapier_payload(event: str, data: dict[str, Any]) -> dict[str, Any]:
"""Build a flat, Zapier-compatible webhook payload.
Zapier works best with flat JSON objects that include an ``id`` field
for deduplication. This function merges event metadata into the
top-level object alongside the event-specific *data*.
Args:
event: The event name (e.g. ``document.processed``).
data: Event-specific key/value pairs.
Returns:
A flat dictionary suitable for Zapier / Make.com consumption.
"""
return {
"id": f"evt_{uuid.uuid4().hex[:16]}",
"event": event,
"timestamp": time.time(),
**data,
}
# ---------------------------------------------------------------------------
# Sample payloads (used by the /triggers/sample endpoint)
# ---------------------------------------------------------------------------
#: Example payloads that Zapier uses for field-mapping during Zap creation.
SAMPLE_PAYLOADS: dict[str, dict[str, Any]] = {
"document.uploaded": {
"id": "evt_sample0001",
"event": "document.uploaded",
"timestamp": 1710000000.0,
"document_id": 42,
"filename": "invoice_2024.pdf",
"content_type": "application/pdf",
"size_bytes": 204800,
"owner_id": "user@example.com",
},
"document.processed": {
"id": "evt_sample0002",
"event": "document.processed",
"timestamp": 1710000060.0,
"document_id": 42,
"filename": "invoice_2024.pdf",
"status": "processed",
"title": "Invoice #1234",
"owner_id": "user@example.com",
},
"document.failed": {
"id": "evt_sample0003",
"event": "document.failed",
"timestamp": 1710000120.0,
"document_id": 42,
"filename": "corrupt.pdf",
"status": "failed",
"error": "Unable to extract text from document",
"owner_id": "user@example.com",
},
"user.signup": {
"id": "evt_sample0004",
"event": "user.signup",
"timestamp": 1710000180.0,
"user_id": "newuser@example.com",
"display_name": "Jane Doe",
},
"user.plan_changed": {
"id": "evt_sample0005",
"event": "user.plan_changed",
"timestamp": 1710000240.0,
"user_id": "user@example.com",
"old_tier": "free",
"new_tier": "pro",
},
"user.payment_issue": {
"id": "evt_sample0006",
"event": "user.payment_issue",
"timestamp": 1710000300.0,
"user_id": "user@example.com",
"issue": "Credit card declined",
},
}
# ---------------------------------------------------------------------------
# Database queries
# ---------------------------------------------------------------------------
def get_active_hooks_for_event(event: str) -> list[dict[str, Any]]:
"""Return all active automation hooks subscribed to *event*.
Args:
event: The event name to filter on.
Returns:
A list of dicts with ``id``, ``target_url``, ``secret``, and
``events`` keys.
"""
db = SessionLocal()
try:
hooks = db.query(AutomationHook).filter(AutomationHook.is_active.is_(True)).all()
result: list[dict[str, Any]] = []
for hook in hooks:
try:
subscribed = json.loads(hook.events)
except (json.JSONDecodeError, TypeError):
subscribed = []
if event in subscribed:
result.append(
{
"id": hook.id,
"target_url": hook.target_url,
"secret": hook.secret,
"events": subscribed,
}
)
return result
finally:
db.close()
# ---------------------------------------------------------------------------
# Dispatch
# ---------------------------------------------------------------------------
def dispatch_automation_hooks(event: str, data: dict[str, Any]) -> None:
"""Fan-out an event to all matching active automation hooks.
Builds a Zapier-compatible flat payload and queues a Celery task for
each matching hook so delivery is asynchronous with automatic retries.
Args:
event: Event name (must be in :data:`VALID_EVENTS`).
data: Event-specific payload data.
"""
if not settings.automation_hooks_enabled:
return
if event not in VALID_EVENTS:
logger.warning("Ignoring unknown automation hook event: %s", event)
return
hooks = get_active_hooks_for_event(event)
if not hooks:
logger.debug("No active automation hooks for event %s", event)
return
payload = build_zapier_payload(event, data)
from app.tasks.automation_tasks import deliver_automation_hook_task
for hook in hooks:
try:
deliver_automation_hook_task.delay(hook["target_url"], payload, hook["secret"])
logger.debug("Queued automation hook delivery to %s for event %s", hook["target_url"], event)
except Exception as exc:
logger.error("Failed to queue automation hook to %s: %s", hook["target_url"], exc)
-378
View File
@@ -1,378 +0,0 @@
"""
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),
)
+4 -3
View File
@@ -15,7 +15,7 @@ import logging
import re
from typing import Any
from sqlalchemy import MetaData, create_engine, inspect, text
from sqlalchemy import MetaData, create_engine, func, inspect, select, table
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import make_url
from sqlalchemy.orm import sessionmaker
@@ -89,8 +89,9 @@ def preview_migration(source_url: str) -> dict[str, Any]:
logger.warning(f"Skipping table with invalid name format: {table_name}")
continue
# table_name is safe — sourced from inspect().get_table_names(), not user input
quoted_table = conn.dialect.identifier_preparer.quote(table_name)
row = conn.execute(text(f"SELECT COUNT(*) FROM {quoted_table}")).fetchone() # noqa: S608
t = table(table_name)
query = select(func.count()).select_from(t)
row = conn.execute(query).fetchone()
count = row[0] if row else 0
result.append({"name": table_name, "row_count": count})
total += count
+12 -1
View File
@@ -7,8 +7,19 @@ def hash_file(filepath: str | Path, chunk_size: int = 65536) -> str:
Returns the SHA-256 hash of the file at 'filepath'.
Reads the file in chunks to handle large files efficiently.
"""
from app.config import settings
filepath_obj = Path(filepath).resolve()
workdir_obj = Path(settings.workdir).resolve()
# Security check: Ensure the resolved path is strictly within the allowed workdir
try:
filepath_obj.relative_to(workdir_obj)
except ValueError:
raise FileNotFoundError(f"Access denied: path traversal attempt or file outside workdir '{filepath}'")
sha256 = hashlib.sha256()
with open(filepath, "rb") as f:
with open(filepath_obj, "rb") as f:
while True:
data = f.read(chunk_size)
if not data:
+45 -377
View File
@@ -39,50 +39,6 @@ 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",
@@ -99,18 +55,6 @@ 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",
@@ -206,14 +150,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"qr_login_enabled": {
"category": "Authentication",
"description": "Enable QR code-based login for mobile device authentication.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
"qr_login_challenge_ttl_seconds": {
"category": "Authentication",
"description": "Time-to-live in seconds for QR login challenges (default 120).",
@@ -270,17 +206,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"sso_auto_login": {
"category": "Authentication",
"description": (
"Automatically redirect to SSO login when authentication is required. "
"Skips the login page and sends users directly to the configured SSO provider."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Social Login Providers
"social_auth_google_enabled": {
"category": "Social Login",
@@ -311,20 +236,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"social_auth_google_use_global_credentials": {
"category": "Social Login",
"description": (
"When True, Google social login uses the global GOOGLE_DRIVE_CLIENT_ID / "
"GOOGLE_DRIVE_CLIENT_SECRET credentials (the Google Drive OAuth integration) "
"instead of requiring separate SOCIAL_AUTH_GOOGLE_CLIENT_ID / "
"SOCIAL_AUTH_GOOGLE_CLIENT_SECRET values. "
"Requires SOCIAL_AUTH_GOOGLE_ENABLED=True and global Google Drive OAuth credentials to be set."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_microsoft_enabled": {
"category": "Social Login",
"description": (
@@ -367,20 +278,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"social_auth_microsoft_use_global_credentials": {
"category": "Social Login",
"description": (
"When True, Microsoft social login uses the global ONEDRIVE_CLIENT_ID / "
"ONEDRIVE_CLIENT_SECRET credentials (the OneDrive integration credentials) "
"instead of requiring separate SOCIAL_AUTH_MICROSOFT_CLIENT_ID / "
"SOCIAL_AUTH_MICROSOFT_CLIENT_SECRET values. "
"Requires SOCIAL_AUTH_MICROSOFT_ENABLED=True and global OneDrive credentials to be set."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_apple_enabled": {
"category": "Social Login",
"description": (
@@ -429,19 +326,6 @@ 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": (
@@ -469,182 +353,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"social_auth_github_enabled": {
"category": "Social Login",
"description": (
"Enable GitHub Sign-In. Requires SOCIAL_AUTH_GITHUB_CLIENT_ID and "
"SOCIAL_AUTH_GITHUB_CLIENT_SECRET from GitHub Developer Settings."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
"help_link": "https://github.com/settings/developers",
"help_link_label": "GitHub Developer Settings",
},
"social_auth_github_client_id": {
"category": "Social Login",
"description": "GitHub OAuth2 client ID from GitHub Developer Settings.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_github_client_secret": {
"category": "Social Login",
"description": "GitHub OAuth2 client secret from GitHub Developer Settings.",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": True,
},
# Keycloak SSO
"social_auth_keycloak_enabled": {
"category": "Social Login",
"description": "Enable Keycloak SSO. Requires server URL, realm, client ID, and client secret.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_keycloak_client_id": {
"category": "Social Login",
"description": "Keycloak OAuth2 client ID.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_keycloak_client_secret": {
"category": "Social Login",
"description": "Keycloak OAuth2 client secret.",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": True,
},
"social_auth_keycloak_server_url": {
"category": "Social Login",
"description": "Keycloak server base URL (e.g. https://keycloak.example.com).",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_keycloak_realm": {
"category": "Social Login",
"description": "Keycloak realm name.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
# Generic OAuth2 SSO
"social_auth_generic_oauth2_enabled": {
"category": "Social Login",
"description": "Enable a generic OAuth2 SSO provider.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_client_id": {
"category": "Social Login",
"description": "Generic OAuth2 client ID.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_client_secret": {
"category": "Social Login",
"description": "Generic OAuth2 client secret.",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_authorize_url": {
"category": "Social Login",
"description": "Generic OAuth2 authorization URL.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_token_url": {
"category": "Social Login",
"description": "Generic OAuth2 token endpoint URL.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_userinfo_url": {
"category": "Social Login",
"description": "Generic OAuth2 userinfo endpoint URL.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_scope": {
"category": "Social Login",
"description": "Space-separated list of OAuth2 scopes to request (default: openid profile email).",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_generic_oauth2_name": {
"category": "Social Login",
"description": "Display name for the generic OAuth2 provider button on the login page.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
# SAML2 SSO
"social_auth_saml2_enabled": {
"category": "Social Login",
"description": "Enable SAML2 SSO authentication.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_saml2_entity_id": {
"category": "Social Login",
"description": "SAML2 Identity Provider Entity ID.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_saml2_sso_url": {
"category": "Social Login",
"description": "SAML2 Identity Provider SSO URL.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"social_auth_saml2_certificate": {
"category": "Social Login",
"description": "SAML2 Identity Provider X.509 certificate (PEM format).",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": True,
},
"social_auth_saml2_name": {
"category": "Social Login",
"description": "Display name for the SAML2 provider.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
# AI Services
"openai_api_key": {
"category": "AI Services",
@@ -1011,18 +719,6 @@ 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",
@@ -2175,30 +1871,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": False,
},
"telegram_enabled": {
"category": "Notifications",
"description": "Enable Telegram bot notifications.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
"telegram_bot_token": {
"category": "Notifications",
"description": "Telegram Bot API token from @BotFather.",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"telegram_chat_id": {
"category": "Notifications",
"description": "Telegram chat ID to send notifications to.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Notifications Settings
"notification_urls": {
"category": "Notifications",
@@ -2297,18 +1969,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": False,
},
"automation_hooks_enabled": {
"category": "Feature Flags",
"description": (
"Enable Zapier / Make.com automation hook subscriptions and delivery. "
"When enabled, external automation platforms can subscribe to DocuElevate events "
"via the REST hooks protocol. Default: True."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
"compliance_enabled": {
"category": "Feature Flags",
"description": (
@@ -2897,6 +2557,51 @@ SETTING_METADATA = {
"required": False,
"restart_required": False,
},
# Database Connection Pool
"db_pool_size": {
"category": "Core",
"description": (
"Number of persistent connections kept in the SQLAlchemy QueuePool. "
"Has no effect for SQLite databases. Default: 5."
),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"db_max_overflow": {
"category": "Core",
"description": (
"Maximum extra connections that can be opened beyond db_pool_size. "
"Has no effect for SQLite databases. Default: 10."
),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"db_pool_timeout": {
"category": "Core",
"description": (
"Seconds to wait for a connection from the pool before raising an error. "
"Has no effect for SQLite databases. Default: 30."
),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"db_pool_recycle": {
"category": "Core",
"description": (
"Seconds after which idle connections are recycled to prevent stale connections. "
"Has no effect for SQLite databases. Default: 1800 (30 minutes)."
),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
# Per-user upload rate limiting
"upload_rate_limit_per_user": {
"category": "Security",
@@ -3232,43 +2937,6 @@ SETTING_METADATA = {
"required": False,
"restart_required": True,
},
"sentry_js_traces_sample_rate": {
"category": "Observability",
"description": (
"Fraction of browser page-loads captured for client-side Sentry performance tracing (0.01.0). "
"0.0 (default) disables browser tracing; 1.0 captures every navigation. "
"Only active when SENTRY_DSN is set."
),
"type": "float",
"sensitive": False,
"required": False,
"restart_required": True,
},
"sentry_js_replay_session_sample_rate": {
"category": "Observability",
"description": (
"Fraction of sessions recorded by Sentry Session Replay (0.01.0). "
"0.0 (default) disables session recording; 1.0 records every session. "
"Only active when SENTRY_DSN is set."
),
"type": "float",
"sensitive": False,
"required": False,
"restart_required": True,
},
"sentry_js_replay_on_error_sample_rate": {
"category": "Observability",
"description": (
"Fraction of error sessions recorded by Sentry Session Replay (0.01.0). "
"Defaults to 0.1 (10%) so that errors are captured with replay context "
"even when session-level recording is disabled. "
"Only active when SENTRY_DSN is set."
),
"type": "float",
"sensitive": False,
"required": False,
"restart_required": True,
},
}
-10
View File
@@ -71,16 +71,6 @@ def notify_settings_updated() -> None:
except Exception as exc:
logger.warning(f"Could not reload in-process settings: {exc}")
# Re-register OAuth / social-login providers so that any provider whose
# credentials were just saved (or updated) in the database is active
# immediately on the login page — no restart required.
try:
from app.auth import refresh_social_providers
refresh_social_providers()
except Exception as exc:
logger.warning(f"Could not refresh social login providers after settings update: {exc}")
# Re-check OCR language availability in the background whenever settings
# are updated. This ensures that if a user changes tesseract_language or
# easyocr_languages via the UI, the new language data is downloaded without
+5 -91
View File
@@ -11,21 +11,14 @@ import logging
from fastapi import Request
from sqlalchemy import or_
from sqlalchemy.orm import Query, Session
from sqlalchemy.orm import Query
from sqlalchemy.sql import false
from app.config import settings
from app.models import FILE_SHARE_ROLE_EDITOR, FILE_SHARE_ROLE_VIEWER, FileRecord, FileShare
from app.models import FileRecord
logger = logging.getLogger(__name__)
# Role hierarchy: higher index = more rights
_ROLE_RANK: dict[str, int] = {
FILE_SHARE_ROLE_VIEWER: 1,
FILE_SHARE_ROLE_EDITOR: 2,
"owner": 3,
}
def _owner_id_from_user(user: dict) -> str | None:
"""Extract the owner identifier from a user dict.
@@ -100,9 +93,8 @@ def apply_owner_filter(query: Query, request: Request) -> Query:
"""Conditionally filter a ``FileRecord`` query by the current user.
When multi-user mode is enabled, only files whose ``owner_id``
matches the authenticated user are returned, **plus** any files that
have been explicitly shared with the user via ``FileShare``. Admin
users bypass the filter and see all documents.
matches the authenticated user are returned. Admin users bypass
the filter and see all documents.
When ``unowned_docs_visible_to_all`` is ``True`` (default), documents
with ``owner_id IS NULL`` (unclaimed) are also included for every
@@ -130,89 +122,11 @@ def apply_owner_filter(query: Query, request: Request) -> Query:
# No authenticated user — return empty result set
return query.filter(false())
# Build filter: user's own documents + documents shared with them
# Build filter: user's own documents
conditions = [FileRecord.owner_id == owner_id]
# Include files explicitly shared with this user
from sqlalchemy import select as sa_select
conditions.append(FileRecord.id.in_(sa_select(FileShare.file_id).where(FileShare.shared_with_user_id == owner_id)))
# Optionally include unclaimed (owner_id IS NULL) documents
if settings.unowned_docs_visible_to_all:
conditions.append(FileRecord.owner_id.is_(None))
return query.filter(or_(*conditions))
def get_file_role(file_record: FileRecord, user_id: str | None, db: Session) -> str | None:
"""Return the effective role a user has on a ``FileRecord``.
Roles (in descending order of privilege):
``"owner"`` — the user's ``owner_id`` matches ``file_record.owner_id``,
or multi-user mode is disabled (everyone is effectively an
owner in single-user mode).
``"editor"`` — the user has an explicit ``FileShare`` with role=editor.
``"viewer"`` — the user has an explicit ``FileShare`` with role=viewer,
or the file is unclaimed (``owner_id IS NULL``) and
``unowned_docs_visible_to_all`` is True.
``None`` — no access.
Args:
file_record: The ``FileRecord`` to check.
user_id: The stable identifier of the requesting user.
db: An active SQLAlchemy session.
Returns:
One of ``"owner"``, ``"editor"``, ``"viewer"``, or ``None``.
"""
if not settings.multi_user_enabled:
# Single-user mode: full access for everyone
return "owner"
if user_id is None:
return None
# Owner always has full access
if file_record.owner_id == user_id:
return "owner"
# Unclaimed document — limited access when setting allows it
if file_record.owner_id is None and settings.unowned_docs_visible_to_all:
return FILE_SHARE_ROLE_VIEWER
# Check for an explicit share
share = (
db.query(FileShare)
.filter(FileShare.file_id == file_record.id, FileShare.shared_with_user_id == user_id)
.first()
)
if share:
return share.role
return None
def has_file_role(
file_record: FileRecord,
user_id: str | None,
db: Session,
minimum_role: str = FILE_SHARE_ROLE_VIEWER,
) -> bool:
"""Return ``True`` if the user's effective role meets the minimum required.
Args:
file_record: The document to check.
user_id: Requesting user's stable identifier.
db: Active SQLAlchemy session.
minimum_role: The minimum role required (``"viewer"``, ``"editor"``,
or ``"owner"``).
Returns:
``True`` when the user's role rank is >= the minimum rank.
"""
role = get_file_role(file_record, user_id, db)
if role is None:
return False
return _ROLE_RANK.get(role, 0) >= _ROLE_RANK.get(minimum_role, 0)
+10 -19
View File
@@ -145,8 +145,6 @@ def dispatch_webhook_event(event: str, data: dict[str, Any]) -> None:
It delegates to :func:`deliver_webhook_task` (Celery) for each matching
webhook so delivery happens asynchronously with automatic retries.
Also dispatches to automation hooks (Zapier / Make.com) if enabled.
Args:
event: Event name (must be in :data:`VALID_EVENTS`).
data: Event-specific payload data.
@@ -158,23 +156,16 @@ def dispatch_webhook_event(event: str, data: dict[str, Any]) -> None:
webhooks = get_active_webhooks_for_event(event)
if not webhooks:
logger.debug("No active webhooks for event %s", event)
else:
payload = build_payload(event, data)
return
# Import here to avoid circular dependency with celery_app
from app.tasks.webhook_tasks import deliver_webhook_task
payload = build_payload(event, data)
for wh in webhooks:
try:
deliver_webhook_task.delay(wh["url"], payload, wh["secret"])
logger.debug("Queued webhook delivery to %s for event %s", wh["url"], event)
except Exception as exc:
logger.error("Failed to queue webhook to %s: %s", wh["url"], exc)
# Import here to avoid circular dependency with celery_app
from app.tasks.webhook_tasks import deliver_webhook_task
# Also fan-out to Zapier / Make.com automation hooks
try:
from app.utils.automation_hooks import dispatch_automation_hooks
dispatch_automation_hooks(event, data)
except Exception as exc:
logger.error("Failed to dispatch automation hooks for event %s: %s", event, exc)
for wh in webhooks:
try:
deliver_webhook_task.delay(wh["url"], payload, wh["secret"])
logger.debug("Queued webhook delivery to %s for event %s", wh["url"], event)
except Exception as exc:
logger.error("Failed to queue webhook to %s: %s", wh["url"], exc)
+29 -20
View File
@@ -96,21 +96,6 @@ def _inject_global_context(ctx: dict) -> None:
)
ctx.setdefault("enable_factory_reset", getattr(settings, "enable_factory_reset", False))
# Sentry Browser SDK config (injected into every page so the JS SDK can initialise)
# Normalize empty-string DSN to None so the {% if sentry_dsn %} template guard works correctly.
_raw_dsn = getattr(settings, "sentry_dsn", None)
ctx.setdefault("sentry_dsn", _raw_dsn if _raw_dsn else None)
ctx.setdefault("sentry_environment", getattr(settings, "sentry_environment", "production"))
ctx.setdefault("sentry_js_traces_sample_rate", getattr(settings, "sentry_js_traces_sample_rate", 0.0))
ctx.setdefault(
"sentry_js_replay_session_sample_rate",
getattr(settings, "sentry_js_replay_session_sample_rate", 0.0),
)
ctx.setdefault(
"sentry_js_replay_on_error_sample_rate",
getattr(settings, "sentry_js_replay_on_error_sample_rate", 0.1),
)
req = ctx.get("request")
if req is not None:
# CSRF token
@@ -162,12 +147,36 @@ def _inject_global_context(ctx: dict) -> None:
def template_response_with_version(*args, **kwargs):
"""Wrapper for TemplateResponse to include version and CSRF token in all templates"""
# If context dict is provided, add version to it
if len(args) >= 2 and isinstance(args[1], dict):
_inject_global_context(args[1])
elif "context" in kwargs and isinstance(kwargs["context"], dict):
"""Wrapper for TemplateResponse to include version and CSRF token in all templates.
Handles both old-style and new-style Starlette TemplateResponse calls:
- Old-style (Starlette <1.0): TemplateResponse(name, {"request": req, ...}, ...)
- New-style (Starlette 1.0+): TemplateResponse(request, name, context={...}, ...)
"""
if len(args) >= 1 and isinstance(args[0], str):
# Old-style call: first positional arg is the template name (string).
# Convert to new-style: (request, name, context=..., ...)
name = args[0]
if len(args) >= 2 and isinstance(args[1], dict):
context = args[1]
# Old-style may have status_code as 3rd positional arg
if len(args) >= 3 and "status_code" not in kwargs:
kwargs["status_code"] = args[2]
else:
context = kwargs.pop("context", {})
request_obj = context.pop("request", None)
if request_obj is not None:
context["request"] = request_obj
_inject_global_context(context)
if request_obj is not None:
return original_template_response(request_obj, name, context=context, **kwargs)
return original_template_response(name, context=context, **kwargs)
# New-style call: (request, name, context=..., ...)
if "context" in kwargs and isinstance(kwargs["context"], dict):
_inject_global_context(kwargs["context"])
elif len(args) >= 3 and isinstance(args[2], dict):
_inject_global_context(args[2])
return original_template_response(*args, **kwargs)
+1 -27
View File
@@ -14,19 +14,6 @@ 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(
@@ -43,8 +30,6 @@ 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 = (
@@ -61,12 +46,6 @@ 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",
{
@@ -77,12 +56,9 @@ async def dropbox_setup_page(
"integration_name": integration.name,
"integration_type": integration.integration_type,
"folder_path": folder_path,
# 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_key_value": "",
"app_secret_value": "",
"refresh_token_value": "",
"global_creds_available": global_creds_available,
"callback_url": callback_url,
},
)
@@ -102,7 +78,6 @@ async def dropbox_setup_page(
"integration_id": integration_id,
"integration_name": None,
"integration_type": None,
"callback_url": callback_url,
},
)
@@ -133,6 +108,5 @@ 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),
},
)
+7 -198
View File
@@ -19,43 +19,6 @@ router = APIRouter()
_FILE_NOT_FOUND = "File not found"
def _resolve_owner_context(request: Request, file_record, db: Session) -> dict:
"""Return owner display info and the current user's effective role.
Returns a dict with:
- ``current_user_role``: one of "owner" / "editor" / "viewer" / None
- ``owner_display``: human-readable owner string (display_name or user_id)
- ``multi_user_enabled``: whether multi-user mode is active
"""
from app.config import settings
from app.models import UserProfile
from app.utils.user_scope import get_current_owner_id, get_file_role
multi_user_enabled = settings.multi_user_enabled
current_owner_id = get_current_owner_id(request)
user_session = request.session.get("user")
is_admin = isinstance(user_session, dict) and bool(user_session.get("is_admin"))
if is_admin:
current_user_role: str | None = "owner"
else:
current_user_role = get_file_role(file_record, current_owner_id, db)
# Build a human-readable owner label
if file_record.owner_id:
profile = db.query(UserProfile).filter(UserProfile.user_id == file_record.owner_id).first()
owner_display: str | None = profile.display_name if profile and profile.display_name else file_record.owner_id
else:
owner_display = None # No owner (unowned)
return {
"current_user_role": current_user_role,
"owner_display": owner_display,
"multi_user_enabled": multi_user_enabled,
}
@router.get("/files")
@require_login
def files_page(
@@ -243,93 +206,10 @@ def files_page(
@router.get("/files/{file_id}")
@require_login
def file_summary_page(request: Request, file_id: int, db: Session = Depends(get_db)):
"""
Return the file summary page — a concise overview with links to detail, processing, and annotations views.
"""
try:
import json
import os
from app.models import FileRecord
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
return templates.TemplateResponse(
"file_summary.html",
{"request": request, "file": None, "error": f"File with ID {file_id} not found"},
)
from app.config import settings
workdir = os.path.realpath(settings.workdir)
def _safe_exists(path: str | None) -> bool:
"""Return True only when *path* exists and resides within workdir."""
if not path:
return False
resolved = os.path.realpath(path)
try:
common = os.path.commonpath([resolved, workdir])
except ValueError:
return False
return common == workdir and os.path.exists(resolved)
original_file_exists = _safe_exists(file_record.original_file_path)
processed_file_exists = _safe_exists(file_record.processed_file_path)
# Load AI metadata — JSON sidecar file first, then DB column
gpt_metadata = None
if file_record.processed_file_path:
metadata_path = os.path.splitext(os.path.realpath(file_record.processed_file_path))[0] + ".json"
if _safe_exists(metadata_path):
try:
with open(metadata_path, "r", encoding="utf-8") as f:
gpt_metadata = json.load(f)
except Exception as e:
logger.warning(f"Failed to load metadata sidecar for file {file_id}: {e}")
if gpt_metadata is None and file_record.ai_metadata:
try:
gpt_metadata = json.loads(file_record.ai_metadata)
except Exception as e:
logger.warning(f"Failed to parse ai_metadata for file {file_id}: {e}")
# Quick processing status
try:
from app.utils.step_manager import get_step_summary as _get_step_summary
step_summary = _get_step_summary(db, file_id)
except Exception:
step_summary = None
pipeline_info = _resolve_pipeline(db, file_record)
owner_ctx = _resolve_owner_context(request, file_record, db)
return templates.TemplateResponse(
"file_summary.html",
{
"request": request,
"file": file_record,
"gpt_metadata": gpt_metadata,
"original_file_exists": original_file_exists,
"processed_file_exists": processed_file_exists,
"step_summary": step_summary,
"pipeline_info": pipeline_info,
**owner_ctx,
},
)
except Exception as e:
logger.error(f"Error retrieving file summary {file_id}: {str(e)}")
return templates.TemplateResponse("file_summary.html", {"request": request, "file": None, "error": str(e)})
@router.get("/files/{file_id}/detail")
@require_login
def file_view_page(request: Request, file_id: int, db: Session = Depends(get_db)):
"""
Return the document detail page — document-centric view with metadata, preview, and extracted text.
Return the document view page — document-centric view with metadata, preview, and extracted text.
Process-oriented details are available via /files/{file_id}/detail.
"""
try:
import json
@@ -392,7 +272,6 @@ def file_view_page(request: Request, file_id: int, db: Session = Depends(get_db)
# Resolve the pipeline assigned to this file (explicit or system default)
pipeline_info = _resolve_pipeline(db, file_record)
owner_ctx = _resolve_owner_context(request, file_record, db)
return templates.TemplateResponse(
"file_view.html",
@@ -404,7 +283,6 @@ def file_view_page(request: Request, file_id: int, db: Session = Depends(get_db)
"processed_file_exists": processed_file_exists,
"step_summary": step_summary,
"pipeline_info": pipeline_info,
**owner_ctx,
},
)
except Exception as e:
@@ -412,11 +290,11 @@ def file_view_page(request: Request, file_id: int, db: Session = Depends(get_db)
return templates.TemplateResponse("file_view.html", {"request": request, "file": None, "error": str(e)})
@router.get("/files/{file_id}/process")
@router.get("/files/{file_id}/detail")
@require_login
def file_detail_page(request: Request, file_id: int, db: Session = Depends(get_db)):
"""
Return the file processing page showing processing history and pipeline information.
Return the file detail page showing processing history and file information
"""
try:
import json
@@ -497,77 +375,6 @@ def file_detail_page(request: Request, file_id: int, db: Session = Depends(get_d
return templates.TemplateResponse("file_detail.html", {"request": request, "file": None, "error": str(e)})
@router.get("/files/{file_id}/annotations")
@require_login
def file_annotations_page(request: Request, file_id: int, db: Session = Depends(get_db)):
"""
Return the comments & annotations page for a file.
"""
try:
import os
from app.models import FileRecord
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not file_record:
return templates.TemplateResponse(
"file_annotations.html",
{"request": request, "file": None, "error": f"File with ID {file_id} not found"},
)
from app.config import settings
workdir = os.path.realpath(settings.workdir)
def _safe_exists(path: str | None) -> bool:
"""Return True only when *path* exists and resides within workdir."""
if not path:
return False
resolved = os.path.realpath(path)
try:
common = os.path.commonpath([resolved, workdir])
except ValueError:
return False
return common == workdir and os.path.exists(resolved)
original_file_exists = _safe_exists(file_record.original_file_path)
processed_file_exists = _safe_exists(file_record.processed_file_path)
# Determine whether the file is a PDF (for EmbedPDF viewer)
mime = file_record.mime_type or ""
is_pdf = mime == "application/pdf" or (file_record.original_filename or "").lower().endswith(".pdf")
# Determine the current user's role on this file (and owner display info)
owner_ctx = _resolve_owner_context(request, file_record, db)
return templates.TemplateResponse(
"file_annotations.html",
{
"request": request,
"file": file_record,
"original_file_exists": original_file_exists,
"processed_file_exists": processed_file_exists,
"is_pdf": is_pdf,
**owner_ctx,
},
)
except Exception as e:
logger.error(f"Error retrieving annotations for file {file_id}: {str(e)}")
return templates.TemplateResponse("file_annotations.html", {"request": request, "file": None, "error": str(e)})
@router.get("/files/{file_id}/comments")
@require_login
def file_comments_redirect(request: Request, file_id: int):
"""
Redirect /files/{file_id}/comments to /files/{file_id}/annotations.
"""
from starlette.responses import RedirectResponse
return RedirectResponse(url=f"/files/{file_id}/annotations", status_code=302)
# ---------------------------------------------------------------------------
# Pipeline ↔ Celery-log stage mapping
# ---------------------------------------------------------------------------
@@ -589,7 +396,9 @@ _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": ["classify_document"],
# "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": [],
}
# These internal bookkeeping stages are always shown in the flow regardless of
+4 -11
View File
@@ -45,9 +45,6 @@ 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",
{
@@ -61,13 +58,10 @@ async def google_drive_setup_page(
"use_oauth": True,
"oauth_configured": bool(integration.credentials),
"sa_configured": False,
"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 ""
),
"client_id": False,
"client_id_value": "",
"client_secret": False,
"client_secret_value": "",
"refresh_token": False,
"refresh_token_value": "",
"has_credentials_json": False,
@@ -96,7 +90,6 @@ 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),
+5 -10
View File
@@ -44,9 +44,6 @@ 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",
{
@@ -57,12 +54,11 @@ async def onedrive_setup_page(
"integration_name": integration.name,
"integration_type": integration.integration_type,
"folder_path": folder_path,
"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",
"client_id": False,
"client_id_value": "",
"client_secret": False,
"client_secret_value": "",
"tenant_id": "common",
"refresh_token": False,
"refresh_token_value": "",
},
@@ -79,7 +75,6 @@ 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),
-340
View File
@@ -195,346 +195,6 @@ async def credentials_page(request: Request, db: Session = Depends(get_db)):
)
@router.get("/admin/connections")
@require_login
@require_admin_access
async def connections_page(request: Request, db: Session = Depends(get_db)):
"""
Connections management page - admin only.
Allows administrators to configure external authentication providers,
SSO settings, and service integrations through a wizard-like interface.
"""
try:
db_settings = get_all_settings_from_db(db)
def _get_effective(key: str):
"""Return DB value if present, else fall back to settings attr."""
if key in db_settings and db_settings[key] is not None:
return db_settings[key]
return getattr(settings, key, None)
def _is_truthy(val) -> bool:
if isinstance(val, bool):
return val
if isinstance(val, str):
return val.lower() in ("true", "1", "yes")
return bool(val)
# Build service status list
services = []
# --- SSO (Authentik / OIDC) ---
_oidc_linked = bool(_get_effective("authentik_client_id") and _get_effective("authentik_client_secret"))
services.append(
{
"key": "oidc",
"name": _get_effective("oauth_provider_name") or "Single Sign-On",
"icon": "fas fa-lock",
"type": "SSO",
"linked": _oidc_linked,
"description": "OpenID Connect SSO provider",
"settings_keys": [
"authentik_client_id",
"authentik_client_secret",
"authentik_config_url",
"oauth_provider_name",
],
}
)
# --- Google ---
_google_id = _get_effective("social_auth_google_client_id")
_google_secret = _get_effective("social_auth_google_client_secret")
if _is_truthy(_get_effective("social_auth_google_use_global_credentials")) and not (
_google_id and _google_secret
):
_google_id = _google_id or _get_effective("google_drive_client_id")
_google_secret = _google_secret or _get_effective("google_drive_client_secret")
_google_linked = bool(
_is_truthy(_get_effective("social_auth_google_enabled")) and _google_id and _google_secret
)
services.append(
{
"key": "google",
"name": "Google",
"icon": "fab fa-google",
"type": "Sign-in authentication",
"linked": _google_linked,
"description": "Sign-in authentication",
"settings_keys": [
"social_auth_google_enabled",
"social_auth_google_client_id",
"social_auth_google_client_secret",
"social_auth_google_use_global_credentials",
],
}
)
# --- GitHub ---
_github_linked = bool(
_is_truthy(_get_effective("social_auth_github_enabled"))
and _get_effective("social_auth_github_client_id")
and _get_effective("social_auth_github_client_secret")
)
services.append(
{
"key": "github",
"name": "GitHub",
"icon": "fab fa-github",
"type": "Sign-in authentication",
"linked": _github_linked,
"description": "Sign-in authentication",
"settings_keys": [
"social_auth_github_enabled",
"social_auth_github_client_id",
"social_auth_github_client_secret",
],
}
)
# --- Microsoft ---
_ms_id = _get_effective("social_auth_microsoft_client_id")
_ms_secret = _get_effective("social_auth_microsoft_client_secret")
if _is_truthy(_get_effective("social_auth_microsoft_use_global_credentials")) and not (_ms_id and _ms_secret):
_ms_id = _ms_id or _get_effective("onedrive_client_id")
_ms_secret = _ms_secret or _get_effective("onedrive_client_secret")
_microsoft_linked = bool(_is_truthy(_get_effective("social_auth_microsoft_enabled")) and _ms_id and _ms_secret)
services.append(
{
"key": "microsoft",
"name": "Microsoft",
"icon": "fab fa-microsoft",
"type": "Sign-in authentication",
"linked": _microsoft_linked,
"description": "Sign-in authentication",
"settings_keys": [
"social_auth_microsoft_enabled",
"social_auth_microsoft_client_id",
"social_auth_microsoft_client_secret",
"social_auth_microsoft_tenant",
"social_auth_microsoft_use_global_credentials",
],
}
)
# --- Apple ---
_apple_linked = bool(
_is_truthy(_get_effective("social_auth_apple_enabled"))
and _get_effective("social_auth_apple_client_id")
and _get_effective("social_auth_apple_team_id")
)
services.append(
{
"key": "apple",
"name": "Apple",
"icon": "fab fa-apple",
"type": "Sign-in authentication",
"linked": _apple_linked,
"description": "Sign-in authentication",
"settings_keys": [
"social_auth_apple_enabled",
"social_auth_apple_client_id",
"social_auth_apple_team_id",
"social_auth_apple_key_id",
"social_auth_apple_private_key",
],
}
)
# --- Dropbox ---
_dbx_id = _get_effective("social_auth_dropbox_client_id")
_dbx_secret = _get_effective("social_auth_dropbox_client_secret")
if _is_truthy(_get_effective("social_auth_dropbox_use_global_credentials")) and not (_dbx_id and _dbx_secret):
_dbx_id = _dbx_id or _get_effective("dropbox_app_key")
_dbx_secret = _dbx_secret or _get_effective("dropbox_app_secret")
_dropbox_linked = bool(_is_truthy(_get_effective("social_auth_dropbox_enabled")) and _dbx_id and _dbx_secret)
services.append(
{
"key": "dropbox",
"name": "Dropbox",
"icon": "fab fa-dropbox",
"type": "Sign-in authentication",
"linked": _dropbox_linked,
"description": "Sign-in authentication",
"settings_keys": [
"social_auth_dropbox_enabled",
"social_auth_dropbox_client_id",
"social_auth_dropbox_client_secret",
"social_auth_dropbox_use_global_credentials",
],
}
)
# --- Keycloak ---
_keycloak_linked = bool(
_is_truthy(_get_effective("social_auth_keycloak_enabled"))
and _get_effective("social_auth_keycloak_client_id")
and _get_effective("social_auth_keycloak_client_secret")
and _get_effective("social_auth_keycloak_server_url")
and _get_effective("social_auth_keycloak_realm")
)
services.append(
{
"key": "keycloak",
"name": "Keycloak",
"icon": "fas fa-key",
"type": "SSO",
"linked": _keycloak_linked,
"description": "SSO",
"settings_keys": [
"social_auth_keycloak_enabled",
"social_auth_keycloak_client_id",
"social_auth_keycloak_client_secret",
"social_auth_keycloak_server_url",
"social_auth_keycloak_realm",
],
}
)
# --- Generic OAuth2 ---
_generic_oauth2_linked = bool(
_is_truthy(_get_effective("social_auth_generic_oauth2_enabled"))
and _get_effective("social_auth_generic_oauth2_client_id")
and _get_effective("social_auth_generic_oauth2_client_secret")
and _get_effective("social_auth_generic_oauth2_authorize_url")
and _get_effective("social_auth_generic_oauth2_token_url")
)
services.append(
{
"key": "generic_oauth2",
"name": "Generic OAuth2",
"icon": "fas fa-sign-in-alt",
"type": "SSO",
"linked": _generic_oauth2_linked,
"description": "SSO",
"settings_keys": [
"social_auth_generic_oauth2_enabled",
"social_auth_generic_oauth2_client_id",
"social_auth_generic_oauth2_client_secret",
"social_auth_generic_oauth2_authorize_url",
"social_auth_generic_oauth2_token_url",
"social_auth_generic_oauth2_userinfo_url",
"social_auth_generic_oauth2_scope",
"social_auth_generic_oauth2_name",
],
}
)
# --- SAML2 ---
_saml2_configured = bool(
_is_truthy(_get_effective("social_auth_saml2_enabled"))
and _get_effective("social_auth_saml2_sso_url")
and _get_effective("social_auth_saml2_entity_id")
)
services.append(
{
"key": "saml2",
"name": settings.social_auth_saml2_name or "SAML2",
"icon": "fas fa-id-badge",
"type": "SSO (SAML)",
"linked": _saml2_configured,
"description": "SSO (SAML)",
"settings_keys": [
"social_auth_saml2_enabled",
"social_auth_saml2_entity_id",
"social_auth_saml2_sso_url",
"social_auth_saml2_certificate",
"social_auth_saml2_name",
],
}
)
# --- SMTP Mail ---
_smtp_configured = bool(_get_effective("email_host") and _get_effective("email_username"))
services.append(
{
"key": "smtp",
"name": "SMTP Mail",
"icon": "fas fa-envelope",
"type": "Email Notifications",
"linked": _smtp_configured,
"description": "Email Notifications",
"settings_keys": [
"email_host",
"email_port",
"email_username",
"email_password",
"email_use_tls",
"email_sender",
],
}
)
# --- Telegram Bot ---
_telegram_configured = bool(
_is_truthy(_get_effective("telegram_enabled")) and _get_effective("telegram_bot_token")
)
services.append(
{
"key": "telegram",
"name": "Telegram Bot",
"icon": "fab fa-telegram",
"type": "Notifications",
"linked": _telegram_configured,
"description": "Configure Telegram bot connectivity, access controls, and feedback behavior.",
"settings_keys": [
"telegram_enabled",
"telegram_bot_token",
"telegram_chat_id",
],
}
)
# Get setting details for the modal forms
service_settings = {}
for svc in services:
svc_settings = []
for skey in svc["settings_keys"]:
meta = get_setting_metadata(skey)
# Get current effective value
val = _get_effective(skey)
display_val = val
if meta.get("sensitive") and val:
display_val = mask_sensitive_value(val)
svc_settings.append(
{
"key": skey,
"value": val,
"display_value": display_val if display_val is not None else "",
"metadata": meta,
}
)
service_settings[svc["key"]] = svc_settings
# Feature toggles
sso_auto_login = _is_truthy(_get_effective("sso_auto_login"))
qr_login_enabled = _is_truthy(_get_effective("qr_login_enabled"))
frontend_url_configured = bool(_get_effective("public_base_url"))
return templates.TemplateResponse(
"admin_connections.html",
{
"request": request,
"services": services,
"service_settings": service_settings,
"sso_auto_login": sso_auto_login,
"oauth_configured": _oidc_linked,
"qr_login_enabled": qr_login_enabled,
"frontend_url_configured": frontend_url_configured,
"app_version": settings.version,
},
)
except HTTPException:
raise
except Exception as e:
logger.error(f"Error loading connections page: {e}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to load connections page",
)
@router.get("/admin/settings/audit-log")
@require_login
@require_admin_access