Compare commits

..

1 Commits

Author SHA1 Message Date
google-labs-jules[bot] 10df0b21bf 🧪 Extract and complete tests for saved searches API
Extracted existing `TestSavedSearchesCRUD` from `tests/test_api_advanced_filters.py` into a dedicated `tests/test_api_saved_searches.py` file to better organize testing logic and reflect the application's file structure.

Significantly improved code coverage of `app/api/saved_searches.py` from 0% (missing configuration imports during tests) to 100% by testing previously untested edge cases including:
- Reaching the maximum saved search limit per user.
- Database commit errors (`HTTP_500_INTERNAL_SERVER_ERROR`) during create, update, and delete actions.
- Validation failures for `filters` field checking for non-dict types (`status.HTTP_422_UNPROCESSABLE_ENTITY`).
- Conflicting names during updates where an existing saved search matches the new name.
- Proper fallback logic across authentication methods for `_get_user_id`.

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
2026-03-23 16:31:34 +00:00
377 changed files with 2834 additions and 142388 deletions
-91
View File
@@ -1,91 +0,0 @@
# =============================================================================
# Docker build context exclusions
# Reducing the build context speeds up builds and prevents unnecessary cache
# invalidation when unrelated files change.
# =============================================================================
# ── Version control ──────────────────────────────────────────────────────────
.git
# ── GitHub / CI tooling ──────────────────────────────────────────────────────
.github
# ── IDE / local dev ──────────────────────────────────────────────────────────
.vscode
.jules
# ── Pre-commit / linting config (not needed at runtime) ──────────────────────
.pre-commit-config.yaml
pyproject.toml
codecov.yml
crowdin.yml
# ── Test suite ───────────────────────────────────────────────────────────────
tests/
requirements-dev.txt
coverage.json
COVERAGE_REPORT.md
.coverage
htmlcov/
.pytest_cache/
junit.xml
coverage.xml
# ── Mobile app / browser extension / legacy placeholder ─────────────────────
# backend/ is an empty placeholder directory not part of the Python application
mobile/
browser-extension/
backend/
# ── Helm charts ──────────────────────────────────────────────────────────────
helm/
# ── Scripts (run before Docker build, output files are COPYd separately) ─────
scripts/
# ── Benchmark and one-off utility scripts ────────────────────────────────────
benchmark_*.py
fix_test*.py
run_fast_tests.sh
# ── Root-level Markdown files (docs/ is kept for docs-builder stage) ─────────
# Note: *.md only matches files at the root level, not inside subdirectories
*.md
# ── Python bytecode / compiled artifacts ─────────────────────────────────────
__pycache__/
*.pyc
*.pyo
*.pyd
*.so
*.egg
*.egg-info/
# ── Virtual environments ──────────────────────────────────────────────────────
.venv/
venv/
env/
# ── Environment / secret files ───────────────────────────────────────────────
.env
.env.local
.env.*.local
# ── Runtime state files ───────────────────────────────────────────────────────
*.log
celerybeat-schedule
celerybeat.pid
# ── Build artifacts ───────────────────────────────────────────────────────────
build/
dist/
.cache/
.mypy_cache/
.ruff_cache/
site/
docs_build/
# ── Editor temp files ─────────────────────────────────────────────────────────
*.swp
*.swo
*~
-48
View File
@@ -7,28 +7,6 @@ GOTENBERG_URL=http://gotenberg:3000
ALLOW_FILE_DELETE=true # Allow deletion of file records ALLOW_FILE_DELETE=true # Allow deletion of file records
COMPLIANCE_ENABLED=true # Enable compliance templates dashboard (GDPR, HIPAA, SOC 2) COMPLIANCE_ENABLED=true # Enable compliance templates dashboard (GDPR, HIPAA, SOC 2)
# **System Reset / Factory Reset**
# FACTORY_RESET_ON_STARTUP=false # Wipe all user data on every startup (demo/testing only)
# ENABLE_FACTORY_RESET=false # Show the System Reset page in admin UI
# **Logging**
# LOG_LEVEL controls the Python root-logger level.
# Accepted values: DEBUG, INFO, WARNING, ERROR, CRITICAL (default: INFO).
# When DEBUG=true and LOG_LEVEL is not set, the level is automatically lowered to DEBUG.
# LOG_LEVEL=INFO
# DEBUG=false
# Log output format: "text" (human-readable, default) or "json" (structured JSON lines).
# Use "json" when shipping logs to Grafana Loki, Splunk, ELK, Datadog, or any SIEM.
# LOG_FORMAT=text
# Forward application logs to a syslog receiver (in addition to stdout).
# Useful for traditional (non-container) deployments and centralised SIEM ingestion.
# LOG_SYSLOG_ENABLED=false
# LOG_SYSLOG_HOST=localhost
# LOG_SYSLOG_PORT=514
# LOG_SYSLOG_PROTOCOL=udp # udp | tcp
# **UI / Appearance** # **UI / Appearance**
# Default colour scheme: system (follow OS), light, or dark # Default colour scheme: system (follow OS), light, or dark
# Individual users can always override with the navbar dark-mode toggle. # Individual users can always override with the navbar dark-mode toggle.
@@ -164,16 +142,6 @@ AUTH_ENABLED=true
# Generate a secure random string, for example: # Generate a secure random string, for example:
# python -c "import secrets; print(secrets.token_hex(32))" # python -c "import secrets; print(secrets.token_hex(32))"
SESSION_SECRET=b39fd43f68d0491ca942f28a16e484b1e763fe9accf4445ca2669a5f3b179eb4 SESSION_SECRET=b39fd43f68d0491ca942f28a16e484b1e763fe9accf4445ca2669a5f3b179eb4
# Session lifetime in days (default: 30). Common values: 30, 60, 90.
# Determines how long a user stays logged in before needing to re-authenticate.
# SESSION_LIFETIME_DAYS=30
# Override with a custom value (takes precedence over SESSION_LIFETIME_DAYS):
# SESSION_LIFETIME_CUSTOM_DAYS=
# Time-to-live in seconds for QR code login challenges (default: 120 = 2 minutes).
# QR_LOGIN_CHALLENGE_TTL_SECONDS=120
ADMIN_USERNAME=admin ADMIN_USERNAME=admin
ADMIN_PASSWORD=your_secure_password ADMIN_PASSWORD=your_secure_password
ADMIN_GROUP_NAME=admin ADMIN_GROUP_NAME=admin
@@ -272,13 +240,6 @@ OPENAI_MODEL=gpt-4o-mini
# AZURE_OPENAI_API_VERSION=2024-02-01 # AZURE_OPENAI_API_VERSION=2024-02-01
# AI_MODEL=gpt-4o # deployment name in Azure # AI_MODEL=gpt-4o # deployment name in Azure
# **Document Translation**
# After processing, documents whose detected language differs from the default
# target language are automatically translated. Only the original and this
# default-language version are persisted; other translations are on-the-fly.
# Users can override this in their profile settings.
# DEFAULT_DOCUMENT_LANGUAGE=en
# Azure Document Intelligence (OCR separate from AI provider above) # Azure Document Intelligence (OCR separate from AI provider above)
# **Email Settings (shared SMTP password reset, verification, and system notifications)** # **Email Settings (shared SMTP password reset, verification, and system notifications)**
EMAIL_HOST=smtp.example.com EMAIL_HOST=smtp.example.com
@@ -447,15 +408,6 @@ ONEDRIVE_TENANT_ID=common
ONEDRIVE_REFRESH_TOKEN=your-refresh-token ONEDRIVE_REFRESH_TOKEN=your-refresh-token
ONEDRIVE_FOLDER_PATH=Documents/Uploads ONEDRIVE_FOLDER_PATH=Documents/Uploads
# SharePoint
SHAREPOINT_CLIENT_ID=your-client-id
SHAREPOINT_CLIENT_SECRET=your-client-secret
SHAREPOINT_TENANT_ID=common
SHAREPOINT_REFRESH_TOKEN=your-refresh-token
SHAREPOINT_SITE_URL=https://tenant.sharepoint.com/sites/sitename
SHAREPOINT_DOCUMENT_LIBRARY=Documents
SHAREPOINT_FOLDER_PATH=Uploads
# WebDAV # WebDAV
# WEBDAV_ENABLED=true # Set to false to disable WebDAV uploads without removing credentials # WEBDAV_ENABLED=true # Set to false to disable WebDAV uploads without removing credentials
WEBDAV_URL=https://webdav.example.com/path WEBDAV_URL=https://webdav.example.com/path
+2 -14
View File
@@ -44,18 +44,6 @@ jobs:
- run: ruff check app/ tests/ - run: ruff check app/ tests/
- run: ruff format --check app/ tests/ - run: ruff format --check app/ tests/
migration-chain:
name: Alembic Migration Chain Check
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.11"
- name: Validate migration chain
run: python scripts/check_alembic_migrations.py
html-lint: html-lint:
name: HTML Accessibility Lint name: HTML Accessibility Lint
runs-on: ubuntu-latest runs-on: ubuntu-latest
@@ -100,7 +88,7 @@ jobs:
python-version: "3.11" python-version: "3.11"
cache: 'pip' cache: 'pip'
- run: pip install pip-audit>=2.7.0 - run: pip install pip-audit>=2.7.0
- run: pip-audit -r requirements.txt --desc on --ignore-vuln CVE-2026-4539 - run: pip-audit -r requirements.txt --desc on
run-tests: run-tests:
name: Execute All Tests (Quick + Integration) name: Execute All Tests (Quick + Integration)
@@ -150,7 +138,7 @@ jobs:
build: build:
name: Build & Push Docker Image name: Build & Push Docker Image
runs-on: ubuntu-latest runs-on: ubuntu-latest
needs: [run-tests, mypy, dependency-scan, html-lint, migration-chain] needs: [run-tests, mypy, dependency-scan, html-lint]
if: github.event_name == 'push' if: github.event_name == 'push'
steps: steps:
- name: Checkout Code - name: Checkout Code
-2
View File
@@ -200,5 +200,3 @@ cython_debug/
# Build metadata files - generated at build time # Build metadata files - generated at build time
GIT_SHA GIT_SHA
RUNTIME_INFO RUNTIME_INFO
node_modules
frontend/node_modules
+4 -4
View File
@@ -1,4 +1,4 @@
## 2026-06-01 - [Fix XSS in status_dashboard.html] ## 2024-05-24 - SSRF in WebDAV connection test
**Vulnerability:** A Cross-Site Scripting (XSS) vulnerability existed in `frontend/templates/status_dashboard.html` where untrusted configuration settings (`value`), external service messages (`data.message`), and token expirations (`data.token_info.expires_in_human`) were injected directly into the DOM via `.innerHTML` without sanitization. **Vulnerability:** The `_test_webdav_connection` function had a custom SSRF check that failed to resolve DNS names, allowing attackers to bypass the check by providing a domain that resolves to an internal IP (e.g., `127.0.0.1`).
**Learning:** Even internal or admin-focused dashboards can be vulnerable if they display external or user-configurable data without escaping. Constructing HTML strings dynamically from unvalidated sources is a common vector for DOM-based XSS. **Learning:** DNS resolution is required for robust SSRF protection when validating URLs provided by users.
**Prevention:** Always use a sanitization function like `escapeHtml` to escape dangerous characters (`<`, `>`, `&`, `"`, `'`) before assigning dynamic content to `.innerHTML`, or prefer `.textContent` when only plaintext is intended. **Prevention:** Use a centralized `is_private_ip` function (now in `app/utils/network.py`) that resolves the hostname to its IPs and checks if any are private.
-10
View File
@@ -48,16 +48,6 @@ repos:
.env.demo .env.demo
)$ )$
# Alembic migration chain validation
- repo: local
hooks:
- id: check-alembic-migrations
name: Check Alembic migration chain
entry: python scripts/check_alembic_migrations.py
language: python
pass_filenames: false
files: ^migrations/versions/.*\.py$
# Conventional commits validation # Conventional commits validation
- repo: https://github.com/compilerla/conventional-pre-commit - repo: https://github.com/compilerla/conventional-pre-commit
rev: v3.0.0 rev: v3.0.0
+1 -1
View File
@@ -1 +1 @@
2026-06-01T03:41:15Z 2026-03-15T21:39:26Z
-1768
View File
File diff suppressed because it is too large Load Diff
+18 -53
View File
@@ -1,48 +1,14 @@
# syntax=docker/dockerfile:1 # Use multi-stage build for a smaller final image
FROM python:3.14.1 AS builder
# ── Stage 1: Python dependency builder ────────────────────────────────────── WORKDIR /app
# Use the same slim variant as the runtime to keep Python versions in sync.
# build-essential + libffi-dev cover the few packages (e.g. cryptography) that
# need a C compiler; they are discarded after this stage.
FROM python:3.14.3-slim AS builder
WORKDIR /build # Copy requirements first for better layer caching
COPY requirements.txt /app/
RUN pip install --no-cache-dir -r requirements.txt
RUN apt-get update && apt-get install -y --no-install-recommends \ # ── Documentation build stage ───────────────────────────────────────────────
build-essential \ FROM python:3.14.1-slim AS docs-builder
libffi-dev \
&& rm -rf /var/lib/apt/lists/*
# Create an isolated virtual environment so only installed packages are copied
# to the runtime image (no pip, setuptools, or other builder artefacts).
RUN python -m venv /opt/venv
ENV PATH="/opt/venv/bin:$PATH" \
PYTHONDONTWRITEBYTECODE=1 \
PIP_NO_CACHE_DIR=1
COPY requirements.txt /build/
RUN pip install --no-cache-dir -r requirements.txt \
# Remove bytecode and cache to keep the venv lean
&& find /opt/venv -type f -name "*.pyc" -delete \
&& find /opt/venv -type d -name "__pycache__" -exec rm -rf {} + 2>/dev/null || true
# ── Stage 2: Frontend asset builder ─────────────────────────────────────────
# Compiles Tailwind CSS (a devDependency) into the minified styles.css.
# npm ci installs ALL deps (including devDependencies) so the tailwindcss CLI
# is available; using --omit=dev would cause 'tailwindcss: not found'.
FROM node:20-slim AS frontend-builder
WORKDIR /frontend
COPY frontend/package.json frontend/package-lock.json ./
RUN npm ci
COPY frontend/ ./
RUN npm run build
# ── Stage 4: Documentation builder ──────────────────────────────────────────
FROM python:3.14.3-slim AS docs-builder
WORKDIR /docs WORKDIR /docs
@@ -57,13 +23,14 @@ COPY mkdocs.yml /docs/mkdocs.yml
# Build the static documentation site # Build the static documentation site
RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build
# ── Stage 5: Runtime image ─────────────────────────────────────────────────── # Second stage for the actual runtime
FROM python:3.14.3-slim FROM python:3.14.3-slim
WORKDIR /app WORKDIR /app
# Copy only the pre-built virtual environment from the builder # Copy installed packages from builder stage
COPY --from=builder /opt/venv /opt/venv COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
COPY --from=builder /usr/local/bin /usr/local/bin
# Install system-level OCR tools required for local OCR workflows: # Install system-level OCR tools required for local OCR workflows:
# tesseract-ocr OCR engine used by pytesseract and ocrmypdf # tesseract-ocr OCR engine used by pytesseract and ocrmypdf
@@ -95,17 +62,15 @@ COPY ./RUNTIME_INFO /app/RUNTIME_INFO
# Copy the pre-built MkDocs documentation site (served at /help) # Copy the pre-built MkDocs documentation site (served at /help)
COPY --from=docs-builder /docs/docs_build /app/docs_build COPY --from=docs-builder /docs/docs_build /app/docs_build
# Copy the compiled Tailwind CSS (built in the frontend-builder stage) # Create runtime_info directory
COPY --from=frontend-builder /frontend/static/styles.css /app/frontend/static/styles.css RUN mkdir -p /app/runtime_info
# Create necessary runtime directories in a single layer # Create necessary directories
RUN mkdir -p /app/runtime_info /workdir RUN mkdir -p /workdir
# Set environment variables # Set environment variables
ENV PATH="/opt/venv/bin:$PATH" \ ENV PYTHONPATH=/app
PYTHONPATH=/app \ ENV PYTHONUNBUFFERED=1
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1
# Expose the port the app runs on # Expose the port the app runs on
EXPOSE 8000 EXPOSE 8000
+13 -35
View File
@@ -1,31 +1,13 @@
# syntax=docker/dockerfile:1
# Local development Dockerfile (avoids CI-only build metadata files) # Local development Dockerfile (avoids CI-only build metadata files)
FROM python:3.14.1 AS builder
# ── Stage 1: Python dependency builder ────────────────────────────────────── WORKDIR /app
FROM python:3.14.3-slim AS builder
WORKDIR /build COPY requirements.txt /app/
RUN pip install --no-cache-dir -r requirements.txt
RUN apt-get update && apt-get install -y --no-install-recommends \ # ── Documentation build stage ───────────────────────────────────────────────
build-essential \ FROM python:3.14.1-slim AS docs-builder
libffi-dev \
&& rm -rf /var/lib/apt/lists/*
# Create an isolated virtual environment
RUN python -m venv /opt/venv
ENV PATH="/opt/venv/bin:$PATH" \
PYTHONDONTWRITEBYTECODE=1 \
PIP_NO_CACHE_DIR=1
COPY requirements.txt /build/
RUN pip install --no-cache-dir -r requirements.txt \
&& find /opt/venv -type f -name "*.pyc" -delete \
&& find /opt/venv -type d -name "__pycache__" -exec rm -rf {} + 2>/dev/null || true
# ── Stage 2: Documentation builder ──────────────────────────────────────────
FROM python:3.14.3-slim AS docs-builder
WORKDIR /docs WORKDIR /docs
@@ -37,25 +19,23 @@ COPY mkdocs.yml /docs/mkdocs.yml
RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build
# ── Stage 3: Runtime image ─────────────────────────────────────────────────── FROM python:3.14.1-slim
FROM python:3.14.3-slim
WORKDIR /app WORKDIR /app
COPY --from=builder /opt/venv /opt/venv COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
COPY --from=builder /usr/local/bin /usr/local/bin
# Install system-level OCR tools required for local OCR workflows: # Install system-level OCR tools required for local OCR workflows:
# tesseract-ocr OCR engine used by pytesseract and ocrmypdf # tesseract-ocr OCR engine used by pytesseract and ocrmypdf
# ghostscript required by ocrmypdf for PDF/PS operations # ghostscript required by ocrmypdf for PDF/PS operations
# poppler-utils provides pdfinfo/pdftoppm used by pdf2image # poppler-utils provides pdfinfo/pdftoppm used by pdf2image
# unpaper optional deskewing pre-processor used by ocrmypdf # unpaper optional deskewing pre-processor used by ocrmypdf
# wget used by ocr_language_manager to download tessdata files
RUN apt-get update && apt-get install -y --no-install-recommends \ RUN apt-get update && apt-get install -y --no-install-recommends \
tesseract-ocr \ tesseract-ocr \
ghostscript \ ghostscript \
poppler-utils \ poppler-utils \
unpaper \ unpaper \
wget \
&& apt-get clean && rm -rf /var/lib/apt/lists/* && apt-get clean && rm -rf /var/lib/apt/lists/*
COPY ./app /app/app COPY ./app /app/app
@@ -73,13 +53,11 @@ COPY --from=docs-builder /docs/docs_build /app/docs_build
RUN echo "local" > /app/GIT_SHA \ RUN echo "local" > /app/GIT_SHA \
&& echo "local" > /app/RUNTIME_INFO && echo "local" > /app/RUNTIME_INFO
# Create necessary runtime directories in a single layer RUN mkdir -p /app/runtime_info
RUN mkdir -p /app/runtime_info /workdir RUN mkdir -p /workdir
ENV PATH="/opt/venv/bin:$PATH" \ ENV PYTHONPATH=/app
PYTHONPATH=/app \ ENV PYTHONUNBUFFERED=1
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1
EXPOSE 8000 EXPOSE 8000
+1 -1
View File
@@ -1 +1 @@
425805a 237af31
+69 -115
View File
@@ -1,6 +1,6 @@
# DocuElevate Milestones # DocuElevate Milestones
**Last Updated:** 2026-05-23 **Last Updated:** 2026-02-08
This document outlines the release milestones, versioning strategy, and detailed feature breakdown for DocuElevate. This document outlines the release milestones, versioning strategy, and detailed feature breakdown for DocuElevate.
@@ -19,15 +19,17 @@ DocuElevate follows [Semantic Versioning 2.0.0](https://semver.org/):
--- ---
## Current State (Continuous Releases) ## Current Release: v0.5.0 (February 2026)
DocuElevate ships continuously via automated semantic versioning. Use **GitHub Releases** for the latest build artifacts and **GitHub Milestones** (below) for roadmap tracking. ### Status: Stable
- Production-ready document processing
### Last Shipped Milestone: v0.5.0 (Released February 8, 2026) - Multi-provider storage support
- Database-backed settings management with encryption - **Database-backed settings management with encryption**
- Setup wizard for first-time configuration - **Setup wizard for first-time configuration**
- Admin UI for runtime configuration - **Admin UI for runtime configuration**
- Release automation via semantic-release - **Automated semantic versioning and releases**
- OAuth2 authentication with admin group support
- Basic web UI and REST API
### Important Note on Versioning ### Important Note on Versioning
As of February 2026, DocuElevate uses **automated semantic versioning**: As of February 2026, DocuElevate uses **automated semantic versioning**:
@@ -128,78 +130,87 @@ As of February 2026, DocuElevate uses **automated semantic versioning**:
## Upcoming Milestones ## Upcoming Milestones
### v0.6.0 - Clarity: Enhanced Search & UI (Target: July 31, 2026) ### v0.6.0 - Enhanced Search & UI Improvements (April 2026)
**Target Date:** July 31, 2026 **Target Date:** April 1, 2026
**Status:** 📋 Planned **Status:** 📋 Planned
**Theme:** Search, Discovery, Modern UX **Theme:** User Experience, Search, Performance
**Epic:** #863
#### Goals #### Goals
- Hybrid discovery: keyword + semantic search, fast filtering, saved searches - Implement full-text search across documents
- Preview-first UX (open, skim, and act quickly) - Responsive mobile interface
- Modern UX polish (accessibility, responsiveness, performance) - Dark mode support
- Document preview in browser
- Performance optimizations
- Improved error handling and user feedback
#### Deliverables #### Deliverables
- Semantic search foundation (vectorization + ranking signals) - Full-text search API and UI
- Saved searches / smart views - Advanced filtering capabilities
- In-browser preview + “quick actions” (tag, route, export) - Responsive CSS framework integration
- Bulk operations and pagination improvements - Dark mode toggle
- UX polish (dark mode/accessibility where applicable) - In-browser document viewer
- Loading states and progress indicators
- Performance benchmarks
- Mobile-optimized interface
#### Breaking Changes #### Breaking Changes
- Potential pagination/search response changes (must be versioned and documented) - API response format changes for search endpoints (documented)
#### Migration Path #### Migration Path
- Version endpoints where needed and keep previous versions working for at least 2 minor milestones - Search endpoint changes will be versioned (/api/v1/search → /api/v2/search)
- Old endpoints deprecated but functional for 2 releases
--- ---
### v0.7.0 - Conductor: Workflow Automation & Integrations (Target: September 30, 2026) ### v0.4.5 - Workflow Automation (June 2026)
**Target Date:** September 30, 2026 **Target Date:** June 1, 2026
**Status:** 📋 Planned **Status:** 📋 Planned
**Theme:** Automation, Integration, Webhooks **Theme:** Automation, Integration, Webhooks
**Epic:** #864
#### Goals #### Goals
- First-class workflow model (steps, state, retries) that matches what the system actually executes - Custom processing pipelines
- Workflow-aware UI status, retries, and observability - Conditional routing based on document type
- Webhooks + event-driven automation foundations - Webhook support for external integrations
- Rule-based classification
- Scheduled batch processing
#### Deliverables #### Deliverables
- Workflow object model and storage - Pipeline configuration UI
- Workflow-aware file detail view + status dashboard - Webhook management interface
- Scheduling primitives (recurring jobs / delayed runs) - Rule engine for document routing
- Webhook system (outbound events + inbound triggers) - Batch processing scheduler
- Integration templates and documentation - Integration examples and templates
- Webhook payload documentation
--- ---
### v0.8.0 - Signal: AI Quality, RAG, and Multi-language (Target: November 30, 2026) ### v0.7.0 - Advanced AI & Multi-language (August 2026)
**Target Date:** November 30, 2026 **Target Date:** August 1, 2026
**Status:** 📋 Planned **Status:** 📋 Planned
**Theme:** AI Quality, Retrieval, Internationalization **Theme:** AI Enhancement, Internationalization
**Epic:** #865
#### Goals #### Goals
- “Chat with Library” foundations (retrieval + UI) - Custom AI model support
- Local AI options for privacy-sensitive setups - Multi-language OCR
- Measurable AI quality (confidence + human review loop) - Document similarity detection
- Expand multilingual capability across OCR + UI - Duplicate detection
- UI internationalization (i18n)
- API localization
#### Deliverables #### Deliverables
- Vector DB integration and embeddings pipeline - Custom model integration API
- Chat UI foundations and retrieval API - Multi-language OCR configuration
- Confidence scoring + human review/edit loop for extracted fields - Similarity algorithm implementation
- Multi-language OCR configuration improvements - Duplicate detection service
- Expanded i18n coverage + localized docs - Translation framework (10+ languages)
- Localized documentation
--- ---
### v1.0.0 - Summit: Enterprise Edition (Target: March 31, 2027) ### v1.0.0 - Enterprise Edition (November 2026)
**Target Date:** March 31, 2027 **Target Date:** November 1, 2026
**Status:** 📋 Planned **Status:** 📋 Planned
**Theme:** Enterprise Features, Scalability, Multi-tenancy **Theme:** Enterprise Features, Scalability, Multi-tenancy
**Epic:** #866
This is our first major release, marking production-ready enterprise capabilities. This is our first major release, marking production-ready enterprise capabilities.
@@ -233,60 +244,6 @@ This is our first major release, marking production-ready enterprise capabilitie
--- ---
### v2.0.0 - Horizon: Platform Expansion (Target: September 30, 2027)
**Target Date:** September 30, 2027
**Status:** 📋 Planned
**Theme:** Ecosystem, Platform, Distribution
**Epic:** #867
#### Goals
- Make DocuElevate extensible by design (plugins + templates)
- Expand integrations and developer experience
- Harden multi-surface experiences (web, mobile, extension, CLI) as a cohesive product
#### Deliverables
- Plugin system foundations and public extension points
- Template library for pipelines/workflows + “starter kits”
- Integration hub patterns (webhooks, events, connectors)
- SDK + documentation for extensions
---
### v2.1.0 - Sentinel: Governance & Policy (Target: March 31, 2028)
**Target Date:** March 31, 2028
**Status:** 📋 Planned
**Theme:** Governance, Compliance, Policy-driven Automation
**Epic:** #868
#### Goals
- Make governance first-class (retention, legal hold, PII workflows)
- Provide tamper-evident auditing and admin controls
- Introduce policy-driven approvals for sensitive automation
#### Deliverables
- Retention policies + legal hold primitives
- PII detection + redaction workflows
- Tamper-evident audit trails + admin activity feed
- Policy-as-code concepts for workflows (with approval gates)
---
### v3.0.0 - Constellation: Integration Hub & Agent Platform (Target: September 30, 2028)
**Target Date:** September 30, 2028
**Status:** 📋 Planned
**Theme:** Ecosystem, Agents, Interoperability
**Epic:** #869
#### Goals
- Make DocuElevate the “system of record” for document intelligence in an organization
- Support external automation ecosystems (Zapier/Make/n8n) and agent runtimes
- Provide a clean interoperability layer for modern AI tools
#### Deliverables
- DocuElevate MCP server (search, retrieve, summarize, route) and documentation
- Connector marketplace concepts (curated + community)
- Event stream + webhooks at scale (delivery guarantees, retries, signing)
## Release Process ## Release Process
### Automated Semantic Versioning (v0.6.0+) ### Automated Semantic Versioning (v0.6.0+)
@@ -298,15 +255,15 @@ Starting with v0.6.0, releases are fully automated using `python-semantic-releas
4. **Automatic Updates**: 4. **Automatic Updates**:
- Updates `VERSION` file - Updates `VERSION` file
- Generates/updates `CHANGELOG.md` - Generates/updates `CHANGELOG.md`
- Creates Git tag (e.g., `v0.173.1`) - Creates Git tag (e.g., `v0.6.0`)
- Creates GitHub Release with notes - Creates GitHub Release with notes
- Triggers Docker image builds - Triggers Docker image builds
5. **No Manual Steps**: VERSION and CHANGELOG are never edited manually 5. **No Manual Steps**: VERSION and CHANGELOG are never edited manually
### Version Bump Rules ### Version Bump Rules
- `feat:` commits → Minor version (e.g., 0.173.1 → 0.174.0) - `feat:` commits → Minor version (0.5.0 → 0.6.0)
- `fix:`, `perf:` → Patch version (e.g., 0.173.1 → 0.173.2) - `fix:`, `perf:` → Patch version (0.5.0 → 0.5.1)
- `feat!:`, `BREAKING CHANGE:` → Major version (e.g., 0.173.1 → 1.0.0) - `feat!:`, `BREAKING CHANGE:` → Major version (0.5.0 → 1.0.0)
- Other types (docs, chore, etc.) → No version bump - Other types (docs, chore, etc.) → No version bump
### Pre-release Checklist (Automated) ### Pre-release Checklist (Automated)
@@ -340,13 +297,10 @@ Starting with v0.6.0, releases are fully automated using `python-semantic-releas
| v0.3.2 | 2026-02-06 | Security Updates | Released | | v0.3.2 | 2026-02-06 | Security Updates | Released |
| v0.3.3 | 2026-02-08 | Drag-and-Drop Upload | Released | | v0.3.3 | 2026-02-08 | Drag-and-Drop Upload | Released |
| v0.5.0 | 2026-02-08 | **Settings & Encryption** | **Released** | | v0.5.0 | 2026-02-08 | **Settings & Encryption** | **Released** |
| v0.6.0 | 2026-07-31 | **Clarity:** Search & UX | Planned | | v0.6.0 | 2026-04 | Search & UX | Planned |
| v0.7.0 | 2026-09-30 | **Conductor:** Workflows & Integrations | Planned | | v0.7.0 | 2026-08 | Advanced AI | Planned |
| v0.8.0 | 2026-11-30 | **Signal:** AI Quality, RAG, Multi-language | Planned | | v1.0.0 | 2026-11 | Enterprise | Planned |
| v1.0.0 | 2027-03-31 | **Summit:** Enterprise | Planned | | v2.0.0 | 2027-Q3 | Platform Expansion | Future |
| v2.0.0 | 2027-09-30 | **Horizon:** Platform Expansion | Future |
| v2.1.0 | 2028-03-31 | **Sentinel:** Governance & Policy | Future |
| v3.0.0 | 2028-09-30 | **Constellation:** Integration Hub & Agents | Future |
--- ---
+1 -4
View File
@@ -35,7 +35,7 @@ DocuElevate is an intelligent document processing system that automates the inge
- **AI-Powered Metadata Extraction** — pluggable AI providers including OpenAI, Anthropic Claude, Google Gemini, Ollama (local), OpenRouter, Portkey, and Azure OpenAI via LiteLLM - **AI-Powered Metadata Extraction** — pluggable AI providers including OpenAI, Anthropic Claude, Google Gemini, Ollama (local), OpenRouter, Portkey, and Azure OpenAI via LiteLLM
- **Multi-Engine OCR** — Azure Document Intelligence, Tesseract, EasyOCR, Mistral OCR, Google Cloud Document AI, and AWS Textract with configurable merge strategies - **Multi-Engine OCR** — Azure Document Intelligence, Tesseract, EasyOCR, Mistral OCR, Google Cloud Document AI, and AWS Textract with configurable merge strategies
- **13 Storage Destinations** — Dropbox, Google Drive, OneDrive, Amazon S3, Nextcloud, WebDAV, FTP, SFTP, iCloud Drive, Email (SMTP), Paperless-ngx, Evernote, and Rclone - **12 Storage Destinations** — Dropbox, Google Drive, OneDrive, Amazon S3, Nextcloud, WebDAV, FTP, SFTP, iCloud Drive, Email (SMTP), Paperless-ngx, and Rclone
- **Multi-Channel Ingestion** — web upload, browser extension, mobile app, CLI, REST API, IMAP email, and watched folders (local, cloud, FTP/SFTP) - **Multi-Channel Ingestion** — web upload, browser extension, mobile app, CLI, REST API, IMAP email, and watched folders (local, cloud, FTP/SFTP)
- **Processing Pipelines** — customizable multi-step workflows with conditional routing rules - **Processing Pipelines** — customizable multi-step workflows with conditional routing rules
- **Full-Text Search** — powered by Meilisearch for instant document discovery - **Full-Text Search** — powered by Meilisearch for instant document discovery
@@ -106,7 +106,6 @@ Processed documents are distributed to any combination of configured destination
| **iCloud Drive** | Apple cloud | | **iCloud Drive** | Apple cloud |
| **Email (SMTP)** | Send as attachment | | **Email (SMTP)** | Send as attachment |
| **Paperless-ngx** | Document management system | | **Paperless-ngx** | Document management system |
| **Evernote** | Notes with PDF attachments |
| **Rclone** | 70+ cloud providers via Rclone | | **Rclone** | 70+ cloud providers via Rclone |
## Features ## Features
@@ -246,7 +245,6 @@ See the [Kubernetes Deployment Guide](docs/KubernetesDeployment.md) for full det
| [Google Drive](docs/GoogleDriveSetup.md) | Google Drive service account / OAuth | | [Google Drive](docs/GoogleDriveSetup.md) | Google Drive service account / OAuth |
| [OneDrive](docs/OneDriveSetup.md) | Microsoft OneDrive setup | | [OneDrive](docs/OneDriveSetup.md) | Microsoft OneDrive setup |
| [Amazon S3](docs/AmazonS3Setup.md) | S3 bucket configuration | | [Amazon S3](docs/AmazonS3Setup.md) | S3 bucket configuration |
| [Evernote](docs/EvernoteSetup.md) | Evernote note creation |
| [Authentication](docs/AuthenticationSetup.md) | OAuth2, OIDC, and social login | | [Authentication](docs/AuthenticationSetup.md) | OAuth2, OIDC, and social login |
| [Notifications](docs/NotificationsSetup.md) | Notification backend setup | | [Notifications](docs/NotificationsSetup.md) | Notification backend setup |
@@ -329,7 +327,6 @@ The following is a summary of the licenses used by our direct dependencies:
| pypdf | BSD | | pypdf | BSD |
| Requests | Apache 2.0 | | Requests | Apache 2.0 |
| Dropbox SDK | MIT | | Dropbox SDK | MIT |
| Evernote SDK | BSD |
| Azure AI Document Intelligence | MIT | | Azure AI Document Intelligence | MIT |
| Authlib | BSD | | Authlib | BSD |
| Starlette | BSD | | Starlette | BSD |
+133 -103
View File
@@ -1,139 +1,169 @@
# DocuElevate Roadmap # DocuElevate Roadmap
**Last Updated:** 2026-05-23 **Last Updated:** 2026-02-08
**Version:** 2.0 **Version:** 1.0
## Vision ## Vision
DocuElevate aims to be the premier open-source intelligent document processing platform, providing seamless integration with cloud storage providers, advanced AI-powered metadata extraction, and enterprise-grade security and scalability. DocuElevate aims to be the premier open-source intelligent document processing platform, providing seamless integration with cloud storage providers, advanced AI-powered metadata extraction, and enterprise-grade security and scalability.
## How to Read This Roadmap
DocuElevate ships frequently (automated semantic versioning), so this roadmap is organized around **milestone outcomes** and **themes**, not exact build numbers.
- **P0** = required for the milestone to feel “done”
- **P1** = strongly desired; may slip if needed
- **P2** = nice-to-have / opportunistic
For the detailed milestone breakdown and target dates, see [MILESTONES.md](MILESTONES.md).
## Release Naming ## Release Naming
Each major milestone release carries a codename to anchor key project moments. These names appear in the status dashboard, build metadata, and changelog. For details, see [docs/ReleaseNaming.md](docs/ReleaseNaming.md). Each major milestone release carries a codename to anchor key project moments. These names appear in the status dashboard, build metadata, and changelog. For details, see [docs/ReleaseNaming.md](docs/ReleaseNaming.md).
| Milestone | Codename | Theme | | Version Range | Codename | Theme |
|----------|------------------|-------| |---------------|---------------|--------------------------------------------------|
| v0.6.0 | **Clarity** | Search, discovery, and modern UX | | 0.5.x | **Foundation** | Core platform, multi-provider storage, AI, UI |
| v0.7.0 | **Conductor** | Workflows, orchestration, and integrations | | 0.6.x | **Clarity** | Enhanced search, filtering, UI/UX improvements |
| v0.8.0 | **Signal** | AI quality, multilingual, and “Chat with Library” foundations | | 0.7.x | **Conductor** | Workflow automation, pipelines, rule-based logic |
| v1.0.0 | **Summit** | Enterprise readiness (multi-tenancy, RBAC, scaling) | | 1.0.x | **Summit** | Enterprise features, multi-tenancy, RBAC |
| v2.0.0 | **Horizon** | Platform expansion and ecosystem maturity | | 1.1.x | **Bridge** | Collaboration, sharing, analytics |
| v2.1.0+ | **Sentinel** | Governance, compliance, and policy-driven automation | | 2.0.x | **Horizon** | On-premise AI, platform expansion |
| v3.0.0 | **Constellation**| Integration hub, agents, and interoperability |
## Current Product Capabilities (Today) ## Current Status (v0.5.0 "Foundation")
### Core Features ✅ ### Core Features ✅
- Multi-channel ingestion (web upload, IMAP email, watched folders, mobile, CLI, API) - Multi-provider document storage (Dropbox, Google Drive, OneDrive, Nextcloud, S3, etc.)
- Multi-engine OCR + AI extraction with configurable providers - IMAP email integration for document ingestion
- Customizable processing pipelines and routing rules - OCR processing via Azure Document Intelligence
- Full-text search and document discovery - AI-powered metadata extraction via OpenAI
- Multi-destination distribution (cloud providers, DMS, protocols, email) - PDF conversion via Gotenberg
- Admin UI for configuration (database-backed settings, encryption, setup wizard) - Web UI for document upload and management
- Production hardening building blocks (CI/CD, security docs, deployment guides) - **Database-backed settings management with admin UI**
- **Fernet encryption for sensitive configuration**
- **Setup wizard for first-time installation**
- REST API with OpenAPI documentation
- Celery-based async task processing
- OAuth2 authentication via Authentik with admin group support
## Feature Landscape (Themes) ## Short-term Goals (Q1-Q2 2026) - v0.4.x to v0.5.x "Foundation"
### 1) Search & Discovery ### Quality & Stability 🎯
- **P0:** hybrid search (keyword + semantic), fast faceted filtering, saved searches - **Test Coverage** (High Priority)
- **P1:** “explain results” (why a document matched), query suggestions, pinned results - [ ] Achieve 80% code coverage for core modules
- **P2:** entity search (people/companies/amounts/dates) and graph-style exploration - [ ] Add integration tests for all storage providers
- [ ] Add end-to-end workflow tests
- [ ] Performance benchmarks and load testing
### 2) AI Quality & Trust - **Code Quality** (High Priority)
- **P0:** confidence scoring, human review/edit loop, extraction evaluation harness - [ ] Enable strict linting in CI/CD
- **P1:** per-document-type schemas/templates, active learning (feedback improves extraction) - [ ] Refactor large modules for better maintainability
- **P2:** multi-model routing (choose model by cost/latency/accuracy per step) - [ ] Add comprehensive type hints
- [ ] Improve error handling and user feedback
### 3) Workflow Automation & Orchestration - **Security** (Critical Priority)
- **P0:** first-class workflow model (steps, state, retries), workflow-aware UI status - [x] Fix known vulnerabilities in dependencies
- **P1:** visual workflow builder, scheduling, webhooks, and event-driven triggers - [ ] Implement rate limiting on API endpoints
- **P2:** agentic workflows (“autopilot” suggestions with approval gates) - [ ] Add CSRF protection
- [ ] Security audit by external party
- [ ] Implement API key rotation
- [ ] Add audit logging for sensitive operations
### 4) Integrations & Ecosystem (Including MCP) - **Release Automation** (Completed ✅)
- **P0:** stable webhooks + outbound actions (Slack/Teams, email, DMS), bi-directional sync where supported - [x] Implement semantic-release for automated versioning
- **P1:** “Integration Hub” (Zapier/Make/n8n style), connector templates, secrets handling patterns - [x] Add conventional commit validation
- **P2:** **MCP**: ship a DocuElevate MCP server (search, retrieve, summarize, route) + allow MCP tools as pipeline steps - [x] Automate CHANGELOG generation
- [x] Integrate Docker builds with releases
### 5) Governance, Compliance, and Security ### Features - v0.4.0
- **P0:** audit trails, tamper-evident logs, API key lifecycle/rotation, admin activity feed - **Enhanced Search & Filtering** → _preparing for v0.6.0 "Clarity"_
- **P1:** retention policies, legal hold, PII detection + redaction, data residency controls - [ ] Full-text search across documents
- **P2:** compliance packs (SOC2/GDPR/HIPAA), BYOK/KMS integration paths - [ ] Advanced filtering by metadata, tags, date ranges
- [ ] Saved search queries
- [ ] Bulk operations on search results
### 6) Enterprise & Scale - **Improved UI/UX**
- **P0:** multi-tenancy, RBAC, horizontal scaling reference architecture - [ ] Responsive mobile interface
- **P1:** SCIM provisioning, SAML/Okta/Azure AD hardening, quotas/billing at org level - [ ] Dark mode support
- **P2:** multi-region deployment patterns and disaster recovery playbooks - [ ] Document preview in browser
- [ ] Drag-and-drop file upload
- [ ] Progress indicators for long-running tasks
- [ ] Real-time notifications via WebSocket
## Release Plan (Extended) ### Features - v0.5.0 "Foundation"
- **Workflow Automation** → _evolving into v0.7.0 "Conductor"_
- [ ] Custom processing pipelines
- [ ] Conditional routing based on document type
- [ ] Scheduled batch processing
- [ ] Webhook support for external integrations
- [ ] Rule-based document classification
This plan extends the existing milestones with a clearer thematic arc and a forward-looking “beyond v2.0” horizon. Each milestone links to an epic issue that owns scope and sub-issues. - **Advanced AI Features**
- [ ] Custom AI models for specialized document types
- [ ] Multi-language OCR support
- [ ] Document similarity detection
- [ ] Automatic duplicate detection
- [ ] Intelligent document splitting
### v0.6.0 — Clarity (Search & UX) ## Medium-term Goals (Q3-Q4 2026) - v1.0.x "Summit"
- **Outcome:** users can reliably find, preview, and act on documents in seconds
- **P0:** semantic search + hybrid ranking, saved searches, fast filters, preview-first UX
- **P1:** bulk operations, query suggestions, accessibility/dark mode polish
- **Tracking:** GitHub milestone `v0.6.0 - Enhanced Search & UI` (epic #863)
### v0.7.0 — Conductor (Workflows & Integrations) ### Enterprise Features - v1.0.0 "Summit"
- **Outcome:** workflows are explicit, inspectable, and automatable end-to-end - **Multi-tenancy**
- **P0:** workflow object model + workflow-aware UI status, retries, pipeline definitions - [ ] Organization/team management
- **P1:** workflow builder, scheduling, inbound/outbound webhooks - [ ] Role-based access control (RBAC)
- **P2:** integration templates + “connector marketplace” concepts - [ ] Per-tenant configuration
- **Tracking:** GitHub milestone `v0.7.0 - Workflow Automation` (epic #864) - [ ] Resource quotas and limits
- [ ] Audit logs per organization
### v0.8.0 — Signal (AI Quality + “Chat with Library” Foundations) - **Scalability**
- **Outcome:** AI features are measurable, reviewable, and safe to trust - [ ] Horizontal scaling support
- **P0:** vector DB + embeddings pipeline, chat UI foundations, local AI options - [ ] Distributed task processing
- **P1:** confidence scoring and review loop, extraction evaluation harness - [ ] Caching layer (Redis/Memcached)
- **P2:** multilingual UX + localization expansion - [ ] Database connection pooling
- **Tracking:** GitHub milestone `v0.8.0 - Advanced AI & Multi-language` (epic #865) - [ ] Message queue optimization
### v1.0.0 — Summit (Enterprise Readiness) - **Advanced Integrations**
- **Outcome:** teams can run DocuElevate with strong isolation, access control, and scale - [ ] Microsoft SharePoint integration
- **P0:** multi-tenancy, RBAC, audit logging, scaling guidance - [ ] Slack/Teams bot integration
- **P1:** SSO hardening (SAML/LDAP), org-level quotas and billing hooks - [ ] Zapier/Make.com integration
- **P2:** enterprise admin experience (policies, approvals, reporting) - [ ] Custom webhook receivers
- **Tracking:** GitHub milestone `v1.0.0 - Enterprise Edition` (epic #866) - [ ] GraphQL API
### v2.0.0 — Horizon (Platform Expansion) ### Features - v1.1.0 "Bridge"
- **Outcome:** DocuElevate becomes an extensible platform with a thriving ecosystem - **Collaboration**
- **P0:** plugin system foundations, SDK + templates, deeper integrations - [ ] Document sharing with expiring links
- **P1:** marketplace patterns, app distribution, mobile/extension maturity - [ ] Comments and annotations
- **P2:** multi-workspace experiences (personal + org) - [ ] Version history and rollback
- **Tracking:** GitHub milestone `v2.0.0 - Platform Expansion` (epic #867) - [ ] Real-time collaborative editing metadata
- [ ] Activity feed
### v2.1.0+ — Sentinel (Governance & Policy) - **Reporting & Analytics**
- **Outcome:** governance becomes a first-class layer (policy-driven automation) - [ ] Processing statistics dashboard
- **P0:** retention + legal hold, PII detection/redaction, tamper-evident audit trails - [ ] Storage usage analytics
- **P1:** BYOK/KMS integration patterns, advanced access policies, compliance reporting - [ ] AI confidence scores and accuracy tracking
- **P2:** “policy as code” for workflows + approvals (change management) - [ ] Cost analysis per provider
- **Tracking:** GitHub milestone `v2.1.0 - Governance & Policy (Sentinel)` (epic #868) - [ ] Export reports (PDF, CSV, Excel)
### v3.0.0 — Constellation (Integration Hub & Agent Platform) ## Long-term Goals (2027+) - v2.0+ "Horizon"
- **Outcome:** DocuElevate plugs into modern automation and AI ecosystems as a first-class system of record
- **P0:** MCP server, durable event stream + production-grade webhooks
- **P1:** connector templates + curated catalog, agent-friendly permissioning and auditing
- **P2:** bring-your-own-agent patterns (sandboxing, scoped credentials)
- **Tracking:** GitHub milestone `v3.0.0 - Integration Hub & Agent Platform (Constellation)` (epic #869)
## Research Bets (Optional / Experimental) ### Strategic Initiatives
- **On-Premise AI Models**
- [ ] Self-hosted OCR (Tesseract, EasyOCR)
- [ ] Local LLM integration (Ollama, LLaMA)
- [ ] GPU acceleration support
- [ ] Model fine-tuning interface
- [ ] Hybrid cloud/on-premise processing
These are longer-horizon bets that should only be productized if they prove real user value. - **Advanced Document Management**
- [ ] Document lifecycle management
- [ ] Retention policies and auto-deletion
- [ ] Compliance templates (GDPR, HIPAA, SOC2)
- [ ] Digital signature support
- [ ] Encryption at rest and in transit
- Knowledge graph over extracted entities (contracts ↔ vendors ↔ invoices) - **Platform Expansion**
- Auto-generated “case files” (collections) from intent (“tax 2025”, “project alpha”) - [ ] Desktop applications (Electron)
- Privacy-preserving learning (federated patterns) to improve extraction quality - [ ] Mobile apps (iOS/Android)
- Document provenance (signing, attestations) and tamper detection - [ ] Browser extensions
- [ ] Command-line interface (CLI)
- [ ] VS Code extension for developers
### Research & Innovation
- [ ] Machine learning for custom document types
- [ ] Blockchain for document provenance
- [ ] Federated learning for privacy-preserving AI
- [ ] Edge computing support
- [ ] Quantum-resistant encryption
## Community & Ecosystem ## Community & Ecosystem
+6 -6
View File
@@ -1,10 +1,10 @@
DocuElevate Build Information DocuElevate Build Information
============================== ==============================
Version: 0.173.4 Version: 0.145.2
Build Date: 2026-06-01T03:41:15Z Build Date: 2026-03-15T21:39:26Z
Git Commit: 425805ab23882944aee0cb02b5e497bc536549c0 Git Commit: 237af31f5fe598cfc3c08f2bbba79b3d0925787e
Git Short SHA: 425805a Git Short SHA: 237af31
Git Branch: main Git Branch: main
Commit Date: 2026-06-01T05:40:53+02:00 Commit Date: 2026-03-15T22:39:04+01:00
Build Timestamp: 2026-06-01T03:41:16Z Build Timestamp: 2026-03-15T21:39:26Z
============================== ==============================
+1 -1
View File
@@ -1 +1 @@
0.173.4 0.145.2
-16
View File
@@ -9,12 +9,9 @@ from fastapi import APIRouter
from app.api.admin_users import router as admin_users_router 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.api_tokens import router as api_tokens_router
from app.api.audit_logs import router as audit_logs_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.azure import router as azure_router
from app.api.backup import router as backup_router from app.api.backup import router as backup_router
from app.api.billing import router as billing_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.compliance import router as compliance_router
from app.api.database import router as database_router from app.api.database import router as database_router
from app.api.diagnostic import router as diagnostic_router from app.api.diagnostic import router as diagnostic_router
@@ -36,21 +33,16 @@ from app.api.pipelines import router as pipelines_router
from app.api.plans import router as plans_router from app.api.plans import router as plans_router
from app.api.process import router as process_router from app.api.process import router as process_router
from app.api.profile import router as profile_router from app.api.profile import router as profile_router
from app.api.qr_auth import router as qr_auth_router
from app.api.queue import router as queue_router from app.api.queue import router as queue_router
from app.api.routing_rules import router as routing_rules_router from app.api.routing_rules import router as routing_rules_router
from app.api.saved_searches import router as saved_searches_router from app.api.saved_searches import router as saved_searches_router
from app.api.scheduled_jobs import router as scheduled_jobs_router from app.api.scheduled_jobs import router as scheduled_jobs_router
from app.api.search import router as search_router from app.api.search import router as search_router
from app.api.sessions import router as sessions_router
from app.api.settings import router as settings_router from app.api.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 public_router as shared_links_public_router
from app.api.shared_links import router as shared_links_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.similarity import router as similarity_router
from app.api.subscriptions import router as subscriptions_router from app.api.subscriptions import router as subscriptions_router
from app.api.system_reset import router as system_reset_router
from app.api.translation import router as translation_router
from app.api.url_upload import router as url_upload_router from app.api.url_upload import router as url_upload_router
# Import all the individual routers # Import all the individual routers
@@ -103,12 +95,4 @@ router.include_router(scheduled_jobs_router)
router.include_router(audit_logs_router) router.include_router(audit_logs_router)
router.include_router(i18n_router) router.include_router(i18n_router)
router.include_router(mobile_router) router.include_router(mobile_router)
router.include_router(sessions_router)
router.include_router(qr_auth_router)
router.include_router(compliance_router) router.include_router(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)
+26 -123
View File
@@ -13,7 +13,7 @@ plaintext is returned exactly once at creation time.
import hashlib import hashlib
import logging import logging
import secrets import secrets
from datetime import datetime, timedelta, timezone from datetime import datetime, timezone
from typing import Annotated, Any from typing import Annotated, Any
from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
@@ -42,9 +42,6 @@ TOKEN_HASH_ITERATIONS = 100_000
#: PBKDF2 salt for API token hashing (not secret, but fixed for determinism). #: PBKDF2 salt for API token hashing (not secret, but fixed for determinism).
TOKEN_HASH_SALT = b"api-token-v1" TOKEN_HASH_SALT = b"api-token-v1"
#: Name prefix used for tokens created by the mobile app flow.
MOBILE_TOKEN_PREFIX = "Mobile App"
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Auth helper # Auth helper
@@ -94,21 +91,6 @@ def hash_token(token: str) -> str:
return dk.hex() return dk.hex()
def _token_to_dict(t: ApiToken) -> dict[str, Any]:
"""Convert an ``ApiToken`` ORM instance to a serialisable dict."""
return {
"id": t.id,
"name": t.name,
"token_prefix": t.token_prefix,
"is_active": t.is_active,
"last_used_at": t.last_used_at,
"last_used_ip": t.last_used_ip,
"created_at": t.created_at,
"revoked_at": t.revoked_at,
"expires_at": t.expires_at,
}
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Pydantic schemas # Pydantic schemas
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -118,12 +100,6 @@ class TokenCreate(BaseModel):
"""Schema for creating a new API token.""" """Schema for creating a new API token."""
name: str = Field(..., min_length=1, max_length=255, description="Human-readable label for the 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): class TokenResponse(BaseModel):
@@ -137,7 +113,6 @@ class TokenResponse(BaseModel):
last_used_ip: str | None last_used_ip: str | None
created_at: datetime | None created_at: datetime | None
revoked_at: datetime | None revoked_at: datetime | None
expires_at: datetime | None
model_config = {"from_attributes": True} model_config = {"from_attributes": True}
@@ -168,16 +143,11 @@ async def create_token(
token_hash_value = hash_token(plaintext) token_hash_value = hash_token(plaintext)
prefix = plaintext[:12] # "de_" prefix + 9 random chars = 12 chars total 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( db_token = ApiToken(
owner_id=owner_id, owner_id=owner_id,
name=body.name, name=body.name,
token_hash=token_hash_value, token_hash=token_hash_value,
token_prefix=prefix, token_prefix=prefix,
expires_at=expires_at,
) )
try: try:
db.add(db_token) db.add(db_token)
@@ -198,7 +168,6 @@ async def create_token(
"last_used_ip": db_token.last_used_ip, "last_used_ip": db_token.last_used_ip,
"created_at": db_token.created_at, "created_at": db_token.created_at,
"revoked_at": db_token.revoked_at, "revoked_at": db_token.revoked_at,
"expires_at": db_token.expires_at,
"token": plaintext, "token": plaintext,
} }
@@ -208,114 +177,48 @@ async def list_tokens(
owner_id: CurrentOwner, owner_id: CurrentOwner,
db: DbSession, db: DbSession,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""List non-mobile API tokens for the authenticated user. """List all API tokens for the authenticated user."""
tokens = db.query(ApiToken).filter(ApiToken.owner_id == owner_id).order_by(ApiToken.created_at.desc()).all()
Mobile tokens (whose names start with ``"Mobile App"``) are excluded return [
from this list; they are managed on the dedicated Devices page via {
``GET /api/api-tokens/mobile``. "id": t.id,
""" "name": t.name,
tokens = ( "token_prefix": t.token_prefix,
db.query(ApiToken) "is_active": t.is_active,
.filter( "last_used_at": t.last_used_at,
ApiToken.owner_id == owner_id, "last_used_ip": t.last_used_ip,
~ApiToken.name.startswith(MOBILE_TOKEN_PREFIX), "created_at": t.created_at,
) "revoked_at": t.revoked_at,
.order_by(ApiToken.created_at.desc()) }
.all() for t in tokens
) ]
return [_token_to_dict(t) for t in tokens]
@router.get("/mobile", response_model=list[TokenResponse])
async def list_mobile_tokens(
owner_id: CurrentOwner,
db: DbSession,
) -> list[dict[str, Any]]:
"""List mobile API tokens for the authenticated user.
Returns tokens whose names start with ``"Mobile App"`` — these are
created via the mobile SSO flow or QR code login.
"""
tokens = (
db.query(ApiToken)
.filter(
ApiToken.owner_id == owner_id,
ApiToken.name.startswith(MOBILE_TOKEN_PREFIX),
)
.order_by(ApiToken.created_at.desc())
.all()
)
return [_token_to_dict(t) for t in tokens]
@router.delete("/{token_id}", status_code=status.HTTP_200_OK) @router.delete("/{token_id}", status_code=status.HTTP_200_OK)
async def revoke_or_delete_token( async def revoke_token(
token_id: int, token_id: int,
owner_id: CurrentOwner, owner_id: CurrentOwner,
db: DbSession, db: DbSession,
) -> dict[str, str]: ) -> 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 The token row is kept for audit purposes but marked inactive with a
but marked inactive with a ``revoked_at`` timestamp. ``revoked_at`` timestamp.
* **Already-revoked token** hard-deleted: the row is permanently
removed from the database.
""" """
db_token = db.query(ApiToken).filter(ApiToken.id == token_id, ApiToken.owner_id == owner_id).first() db_token = db.query(ApiToken).filter(ApiToken.id == token_id, ApiToken.owner_id == owner_id).first()
if not db_token: if not db_token:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Token not found") raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Token not found")
if db_token.is_active: if not db_token.is_active:
# Soft-revoke the active token. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already revoked")
try:
db_token.is_active = False
db_token.revoked_at = datetime.now(timezone.utc)
db.commit()
except Exception:
db.rollback()
raise
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
return {"detail": "Token revoked"}
# Hard-delete an already-revoked token.
try: try:
db.delete(db_token) db_token.is_active = False
db_token.revoked_at = datetime.now(timezone.utc)
db.commit() db.commit()
except Exception: except Exception:
db.rollback() db.rollback()
raise raise
logger.info("API token permanently deleted: id=%s owner=%s", token_id, owner_id)
return {"detail": "Token deleted"}
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
@router.post("/{token_id}/reactivate", status_code=status.HTTP_200_OK, response_model=TokenResponse) return {"detail": "Token revoked"}
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)
-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,
}
+4 -3
View File
@@ -168,8 +168,9 @@ async def create_checkout_session(
checkout_session = client.checkout.sessions.create(params=session_params) checkout_session = client.checkout.sessions.create(params=session_params)
logger.info( logger.info(
"Created Stripe checkout session %s for plan %s", "Created Stripe checkout session %s for user %s plan %s",
checkout_session.id, checkout_session.id,
owner_id,
body.plan_id, body.plan_id,
) )
return {"checkout_url": checkout_session.url, "session_id": checkout_session.id} return {"checkout_url": checkout_session.url, "session_id": checkout_session.id}
@@ -212,7 +213,7 @@ async def create_portal_session(
} }
) )
logger.info("Created Stripe portal session for user") logger.info("Created Stripe portal session for user %s", owner_id)
return {"portal_url": portal.url} return {"portal_url": portal.url}
@@ -260,7 +261,7 @@ async def stripe_webhook(request: Request, db: Session = Depends(get_db)) -> dic
@require_login @require_login
async def billing_success(request: Request) -> Any: async def billing_success(request: Request) -> Any:
"""Show a success page after a completed Stripe Checkout.""" """Show a success page after a completed Stripe Checkout."""
return _templates.TemplateResponse(request, "billing_success.html") return _templates.TemplateResponse("billing_success.html", {"request": request})
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
-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() 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") @router.get("/diagnostic/health")
@require_login @require_login
+47 -235
View File
@@ -5,9 +5,7 @@ Dropbox API endpoints
import logging import logging
import os import os
from typing import Annotated, Optional from typing import Annotated, Optional
from urllib.parse import quote
import httpx
import requests import requests
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -25,104 +23,6 @@ logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
def _require_admin(request: Request) -> dict:
"""Ensure the caller is an admin. Raises 403 otherwise."""
user = request.session.get("user")
if not user or not user.get("is_admin"):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
return user
AdminUser = Annotated[dict, Depends(_require_admin)]
def _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") @router.post("/dropbox/exchange-token")
@require_login @require_login
async def exchange_dropbox_token( async def exchange_dropbox_token(
@@ -232,60 +132,57 @@ async def test_dropbox_token(request: Request):
"message": "Dropbox credentials are not fully configured", "message": "Dropbox credentials are not fully configured",
} }
async with httpx.AsyncClient() as client: # Check token validity by getting current account info
# Check token validity by getting current account info headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"}
headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"} response = requests.post(
response = await client.post( "https://api.dropboxapi.com/2/users/get_current_account",
headers=headers,
timeout=settings.http_request_timeout,
)
# If token is invalid, try refreshing it
if response.status_code == 401:
logger.info("Dropbox access token invalid or expired, trying to refresh")
# Get a new access token using the refresh token
refresh_url = "https://api.dropbox.com/oauth2/token"
refresh_data = {
"grant_type": "refresh_token",
"refresh_token": settings.dropbox_refresh_token,
"client_id": settings.dropbox_app_key,
"client_secret": settings.dropbox_app_secret,
}
refresh_response = requests.post(refresh_url, data=refresh_data, timeout=settings.http_request_timeout)
if refresh_response.status_code != 200:
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
return {
"status": "error",
"message": "Refresh token has expired or is invalid",
"needs_reauth": True,
}
token_info = refresh_response.json()
access_token = token_info.get("access_token")
# Try again with the new access token
headers = {"Authorization": f"Bearer {access_token}"}
response = requests.post(
"https://api.dropboxapi.com/2/users/get_current_account", "https://api.dropboxapi.com/2/users/get_current_account",
headers=headers, headers=headers,
timeout=settings.http_request_timeout, timeout=settings.http_request_timeout,
) )
# If token is invalid, try refreshing it if response.status_code != 200:
if response.status_code == 401: logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
logger.info("Dropbox access token invalid or expired, trying to refresh") return {
"status": "error",
"message": f"Token validation failed with status {response.status_code}: {response.text}",
}
# Get a new access token using the refresh token # Get account info
refresh_url = "https://api.dropbox.com/oauth2/token" account_info = response.json()
refresh_data = {
"grant_type": "refresh_token",
"refresh_token": settings.dropbox_refresh_token,
"client_id": settings.dropbox_app_key,
"client_secret": settings.dropbox_app_secret,
}
refresh_response = await client.post(
refresh_url, data=refresh_data, timeout=settings.http_request_timeout
)
if refresh_response.status_code != 200:
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
return {
"status": "error",
"message": "Refresh token has expired or is invalid",
"needs_reauth": True,
}
token_info = refresh_response.json()
access_token = token_info.get("access_token")
# Try again with the new access token
headers = {"Authorization": f"Bearer {access_token}"}
response = await client.post(
"https://api.dropboxapi.com/2/users/get_current_account",
headers=headers,
timeout=settings.http_request_timeout,
)
if response.status_code != 200:
logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
return {
"status": "error",
"message": f"Token validation failed with status {response.status_code}: {response.text}",
}
# Get account info
account_info = response.json()
account_email = account_info.get("email", "Unknown account") account_email = account_info.get("email", "Unknown account")
account_name = account_info.get("name", {}).get("display_name", "Unknown user") account_name = account_info.get("name", {}).get("display_name", "Unknown user")
@@ -310,100 +207,15 @@ async def test_dropbox_token(request: Request):
return {"status": "error", "message": f"Connection error: {str(e)}"} 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") @router.post("/dropbox/save-settings")
@require_login
async def save_dropbox_settings( async def save_dropbox_settings(
request: Request, request: Request,
refresh_token: Annotated[str, Form(...)], refresh_token: Annotated[str, Form(...)],
_admin: AdminUser,
db: Session = Depends(get_db),
app_key: Annotated[Optional[str], Form()] = None, app_key: Annotated[Optional[str], Form()] = None,
app_secret: Annotated[Optional[str], Form()] = None, app_secret: Annotated[Optional[str], Form()] = None,
folder_path: Annotated[Optional[str], Form()] = None, folder_path: Annotated[Optional[str], Form()] = None,
db: Session = Depends(get_db),
): ):
""" """
Save Dropbox settings to database (primary) and .env file (best-effort). Save Dropbox settings to database (primary) and .env file (best-effort).
+23 -28
View File
@@ -73,38 +73,33 @@ def list_duplicate_groups(
groups = [] groups = []
total_duplicate_files = 0 total_duplicate_files = 0
if dup_hashes: for filehash in dup_hashes:
# Fetch all matching files (both original and duplicates) in a single batch query # Find the original (non-duplicate) record with this hash
all_records = ( original = (
db.query(FileRecord).filter(FileRecord.filehash.in_(dup_hashes)).order_by(FileRecord.id.asc()).all() db.query(FileRecord)
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
.order_by(FileRecord.id.asc())
.first()
) )
# Group records by hash # Find all duplicate records for this hash
originals_by_hash = {} duplicates = (
duplicates_by_hash = {h: [] for h in dup_hashes} db.query(FileRecord)
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(True))
.order_by(FileRecord.id.asc())
.all()
)
for record in all_records: total_duplicate_files += len(duplicates)
h = record.filehash
if not record.is_duplicate:
# Store only the first original record per hash, matching the old .first() behaviour
if h not in originals_by_hash:
originals_by_hash[h] = record
else:
duplicates_by_hash[h].append(record)
total_duplicate_files += 1
for filehash in dup_hashes: groups.append(
original = originals_by_hash.get(filehash) {
duplicates = duplicates_by_hash.get(filehash, []) "filehash": filehash,
"original": _file_record_to_dict(original) if original else None,
groups.append( "duplicates": [_file_record_to_dict(d) for d in duplicates],
{ "duplicate_count": len(duplicates),
"filehash": filehash, }
"original": _file_record_to_dict(original) if original else None, )
"duplicates": [_file_record_to_dict(d) for d in duplicates],
"duplicate_count": len(duplicates),
}
)
total_pages = (total_groups + per_page - 1) // per_page if total_groups > 0 else 1 total_pages = (total_groups + per_page - 1) // per_page if total_groups > 0 else 1
+49 -155
View File
@@ -11,7 +11,6 @@ import zipfile
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Annotated, List, Optional from typing import Annotated, List, Optional
import aiofiles
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile, status from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile, status
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from sqlalchemy import asc, desc from sqlalchemy import asc, desc
@@ -20,7 +19,6 @@ from sqlalchemy.orm import Session
from app.auth import require_login from app.auth import require_login
from app.config import settings from app.config import settings
from app.database import get_db 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.models import FileProcessingStep, FileRecord, ProcessingLog
from app.tasks.convert_to_pdf import convert_to_pdf from app.tasks.convert_to_pdf import convert_to_pdf
from app.tasks.process_document import process_document from app.tasks.process_document import process_document
@@ -30,7 +28,7 @@ from app.utils.file_queries import apply_status_filter
from app.utils.file_status import get_files_processing_status from app.utils.file_status import get_files_processing_status
from app.utils.filename_utils import sanitize_filename 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.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 # Set up logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -300,7 +298,6 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
""" """
Delete a file record from the database. Delete a file record from the database.
This only removes the database entry, not the actual file. 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 # Check if file deletion is allowed
if not settings.allow_file_delete: if not settings.allow_file_delete:
@@ -315,18 +312,6 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found") 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 # Log the deletion
logger.info(f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}") logger.info(f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}")
@@ -353,7 +338,6 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
""" """
Delete multiple file records from the database. Delete multiple file records from the database.
This only removes the database entries, not the actual files. 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 # Check if file deletion is allowed
if not settings.allow_file_delete: if not settings.allow_file_delete:
@@ -361,25 +345,11 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
try: try:
# Find all file records # Find all file records
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)) file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
query = apply_owner_filter(query, request)
file_records = query.all()
if not file_records: if not file_records:
raise HTTPException(status_code=404, detail="No files found with the provided IDs") 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_count = len(file_records)
deleted_ids = [f.id for f in file_records] deleted_ids = [f.id for f in file_records]
@@ -414,9 +384,7 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
""" """
try: try:
# Find all file records # Find all file records
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)) file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
query = apply_owner_filter(query, request)
file_records = query.all()
if not file_records: if not file_records:
raise HTTPException(status_code=404, detail="No files found with the provided IDs") raise HTTPException(status_code=404, detail="No files found with the provided IDs")
@@ -488,9 +456,7 @@ def bulk_reprocess_files_cloud_ocr(request: Request, file_ids: List[int], db: Db
Useful for re-running OCR on files with poor text quality or missing OCR text. Useful for re-running OCR on files with poor text quality or missing OCR text.
""" """
try: try:
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)) file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
query = apply_owner_filter(query, request)
file_records = query.all()
if not file_records: if not file_records:
raise HTTPException(status_code=404, detail="No files found with the provided IDs") raise HTTPException(status_code=404, detail="No files found with the provided IDs")
@@ -570,9 +536,7 @@ def bulk_download_files(request: Request, file_ids: List[int], db: DbSession):
Files not found on disk are silently skipped. Files not found on disk are silently skipped.
""" """
try: try:
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)) file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
query = apply_owner_filter(query, request)
file_records = query.all()
if not file_records: if not file_records:
raise HTTPException(status_code=404, detail="No files found with the provided IDs") raise HTTPException(status_code=404, detail="No files found with the provided IDs")
@@ -654,9 +618,7 @@ def reprocess_single_file(request: Request, file_id: int, db: DbSession):
""" """
try: try:
# Find the file record # Find the file record
query = db.query(FileRecord).filter(FileRecord.id == file_id) file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
query = apply_owner_filter(query, request)
file_record = query.first()
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found") raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
@@ -712,9 +674,7 @@ def reprocess_with_cloud_ocr(request: Request, file_id: int, db: DbSession):
""" """
try: try:
# Find the file record # Find the file record
query = db.query(FileRecord).filter(FileRecord.id == file_id) file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
query = apply_owner_filter(query, request)
file_record = query.first()
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found") raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
@@ -977,9 +937,7 @@ def retry_subtask(
""" """
try: try:
# Find the file record # Find the file record
query = db.query(FileRecord).filter(FileRecord.id == file_id) file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
query = apply_owner_filter(query, request)
file_record = query.first()
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found") raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
@@ -1121,9 +1079,7 @@ def get_file_preview(
try: try:
# Find the file record # Find the file record
query = db.query(FileRecord).filter(FileRecord.id == file_id) file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
query = apply_owner_filter(query, request)
file_record = query.first()
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found") raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
@@ -1203,9 +1159,7 @@ def download_file(
try: try:
# Find the file record # Find the file record
query = db.query(FileRecord).filter(FileRecord.id == file_id) file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
query = apply_owner_filter(query, request)
file_record = query.first()
if not file_record: if not file_record:
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found") raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
@@ -1263,79 +1217,9 @@ def download_file(
raise HTTPException(status_code=500, detail=f"Error downloading file: {str(e)}") raise HTTPException(status_code=500, detail=f"Error downloading file: {str(e)}")
async def _save_upload_file_chunks(file: UploadFile, target_path: str, max_size: int) -> int:
"""Save an uploaded file in chunks and enforce the maximum size limit."""
try:
written_size = 0
with open(target_path, "wb") as f:
chunk_size = 65536 # 64 KB chunks
while True:
chunk = await file.read(chunk_size)
if not chunk:
break
written_size += len(chunk)
if written_size > max_size:
# Exceeded limit mid-stream; clean up and reject
f.close()
os.remove(target_path)
raise HTTPException(
status_code=413,
detail=f"File too large: exceeded {max_size} bytes during upload. "
f"See SECURITY_AUDIT.md for configuration details.",
)
f.write(chunk)
return written_size
except HTTPException:
raise
except Exception as e:
if os.path.exists(target_path):
os.remove(target_path)
raise HTTPException(status_code=500, detail=f"Failed to save file: {e}")
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).
"""
if not settings.enable_deduplication:
return None
try:
filehash = hash_file(target_path)
existing = (
db.query(FileRecord)
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
.order_by(FileRecord.id.asc())
.first()
)
if existing:
logger.info(f"Exact duplicate detected on upload: '{safe_filename}' matches file ID {existing.id}")
return {
"duplicate_type": "exact",
"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."
),
}
except Exception as e:
logger.warning(f"Duplicate check failed for uploaded file '{safe_filename}': {e}")
return None
@router.post("/ui-upload") @router.post("/ui-upload")
@require_login @require_login
async def ui_upload( async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...)):
request: Request,
db: DbSession,
file: UploadFile = File(...),
_rate_ok: None = Depends(require_upload_rate_limit),
):
"""Endpoint to accept a user-uploaded file and enqueue it for processing.""" """Endpoint to accept a user-uploaded file and enqueue it for processing."""
workdir = settings.workdir workdir = settings.workdir
@@ -1393,7 +1277,7 @@ async def ui_upload(
# enforcing the size limit during the read so memory usage stays bounded. # enforcing the size limit during the read so memory usage stays bounded.
try: try:
written_size = 0 written_size = 0
async with aiofiles.open(target_path, "wb") as f: with open(target_path, "wb") as f:
chunk_size = 65536 # 64 KB chunks chunk_size = 65536 # 64 KB chunks
while True: while True:
chunk = await file.read(chunk_size) chunk = await file.read(chunk_size)
@@ -1402,14 +1286,14 @@ async def ui_upload(
written_size += len(chunk) written_size += len(chunk)
if written_size > max_size: if written_size > max_size:
# Exceeded limit mid-stream; clean up and reject # Exceeded limit mid-stream; clean up and reject
await f.close() f.close()
os.remove(target_path) os.remove(target_path)
raise HTTPException( raise HTTPException(
status_code=413, status_code=413,
detail=f"File too large: exceeded {max_size} bytes during upload. " detail=f"File too large: exceeded {max_size} bytes during upload. "
f"See SECURITY_AUDIT.md for configuration details.", f"See SECURITY_AUDIT.md for configuration details.",
) )
await f.write(chunk) f.write(chunk)
except HTTPException: except HTTPException:
raise raise
except Exception as e: except Exception as e:
@@ -1421,25 +1305,6 @@ async def ui_upload(
logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'") logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'")
file_size = written_size 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 # Determine if the file is a PDF or needs conversion
mime_type, _ = mimetypes.guess_type(target_path) mime_type, _ = mimetypes.guess_type(target_path)
file_ext = os.path.splitext(target_path)[1].lower() file_ext = os.path.splitext(target_path)[1].lower()
@@ -1503,8 +1368,6 @@ async def ui_upload(
".tif", ".tif",
".webp", ".webp",
".svg", ".svg",
".heic",
".heif",
}: }:
# If it's an image, convert to PDF first # 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) task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
@@ -1518,12 +1381,42 @@ async def ui_upload(
logger.warning(f"Unsupported MIME type {mime_type} for {target_path}, attempting conversion") 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) 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 = None
if settings.enable_deduplication:
try:
filehash = hash_file(target_path)
existing = (
db.query(FileRecord)
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
.order_by(FileRecord.id.asc())
.first()
)
if existing:
exact_duplicate_warning = {
"duplicate_type": "exact",
"original_file_id": existing.id,
"original_filename": existing.original_filename,
"message": (
"This file appears to be an exact duplicate of an already-processed document. "
"It will still be queued but will be flagged as a duplicate."
),
}
logger.info(f"Exact duplicate detected on upload: '{safe_filename}' matches file ID {existing.id}")
except Exception as e:
logger.warning(f"Duplicate check failed for uploaded file '{safe_filename}': {e}")
response: dict = {
"task_id": task.id, "task_id": task.id,
"status": "queued", "status": "queued",
"original_filename": safe_filename, "original_filename": safe_filename,
"stored_filename": target_filename, "stored_filename": target_filename,
} }
if exact_duplicate_warning:
response["duplicate_warning"] = exact_duplicate_warning
return response
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -1565,7 +1458,7 @@ def claim_file(request: Request, file_id: int, db: DbSession):
logger.exception(f"Error claiming file {file_id}: {e}") logger.exception(f"Error claiming file {file_id}: {e}")
raise HTTPException(status_code=500, detail="Failed to claim document") raise HTTPException(status_code=500, detail="Failed to claim document")
logger.info("File %d claimed by user", file_id) logger.info(f"File {file_id} claimed by user '{owner_id}'")
return {"status": "success", "message": "Document claimed successfully", "file_id": file_id, "owner_id": owner_id} return {"status": "success", "message": "Document claimed successfully", "file_id": file_id, "owner_id": owner_id}
@@ -1605,7 +1498,7 @@ def bulk_claim_files(request: Request, file_ids: list[int], db: DbSession):
logger.exception(f"Error during bulk claim: {e}") logger.exception(f"Error during bulk claim: {e}")
raise HTTPException(status_code=500, detail="Failed to claim documents") raise HTTPException(status_code=500, detail="Failed to claim documents")
logger.info("Bulk claim: claimed=%s, skipped=%s", claimed, [s["file_id"] for s in skipped]) logger.info(f"Bulk claim by '{owner_id}': claimed={claimed}, skipped={[s['file_id'] for s in skipped]}")
return { return {
"status": "success", "status": "success",
"claimed_count": len(claimed), "claimed_count": len(claimed),
@@ -1658,7 +1551,8 @@ def assign_owner(request: Request, db: DbSession, owner_id: str = Query(...), fi
logger.exception(f"Error assigning owner: {e}") logger.exception(f"Error assigning owner: {e}")
raise HTTPException(status_code=500, detail="Failed to assign owner") raise HTTPException(status_code=500, detail="Failed to assign owner")
logger.info("Admin assigned owner to %d file(s)", updated) admin_name = get_current_owner_id(request) or "admin"
logger.info(f"Admin '{admin_name}' assigned owner_id='{owner_id}' to {updated} file(s)")
return { return {
"status": "success", "status": "success",
"message": f"Assigned owner to {updated} document(s)", "message": f"Assigned owner to {updated} document(s)",
+13 -26
View File
@@ -23,17 +23,6 @@ logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
def _require_admin(request: Request) -> dict:
"""Ensure the caller is an admin. Raises 403 otherwise."""
user = request.session.get("user")
if not user or not user.get("is_admin"):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
return user
AdminUser = Annotated[dict, Depends(_require_admin)]
@router.post("/google-drive/exchange-token") @router.post("/google-drive/exchange-token")
@require_login @require_login
async def exchange_google_drive_token( async def exchange_google_drive_token(
@@ -373,15 +362,15 @@ def format_time_remaining(time_delta):
@router.post("/google-drive/save-settings") @router.post("/google-drive/save-settings")
async def save_google_drive_settings( @require_login
async def save_dropbox_settings(
request: Request, request: Request,
refresh_token: Annotated[str, Form(...)], refresh_token: Annotated[str, Form(...)],
_admin: AdminUser,
db: Session = Depends(get_db),
client_id: Annotated[Optional[str], Form()] = None, client_id: Annotated[Optional[str], Form()] = None,
client_secret: Annotated[Optional[str], Form()] = None, client_secret: Annotated[Optional[str], Form()] = None,
folder_id: Annotated[Optional[str], Form()] = None, folder_id: Annotated[Optional[str], Form()] = None,
use_oauth: Annotated[str, Form()] = "true", use_oauth: Annotated[str, Form()] = "true",
db: Session = Depends(get_db),
): ):
""" """
Save Google Drive settings to the .env file (best-effort) and persist to database. Save Google Drive settings to the .env file (best-effort) and persist to database.
@@ -414,10 +403,9 @@ async def save_google_drive_settings(
if folder_id: if folder_id:
drive_settings["GOOGLE_DRIVE_FOLDER_ID"] = folder_id drive_settings["GOOGLE_DRIVE_FOLDER_ID"] = folder_id
# Best-effort .env file write — failures here are non-fatal # Try to update the .env file, but don't fail if it doesn't exist (for Docker containers)
env_file_written = False if os.path.exists(env_path):
try: try:
if os.path.exists(env_path):
logger.info(f"Updating Google Drive settings in {env_path}") logger.info(f"Updating Google Drive settings in {env_path}")
# Read the current .env file # Read the current .env file
@@ -450,13 +438,12 @@ async def save_google_drive_settings(
f.write("\n".join(new_env_lines) + "\n") f.write("\n".join(new_env_lines) + "\n")
logger.info("Successfully updated Google Drive settings in .env file") logger.info("Successfully updated Google Drive settings in .env file")
env_file_written = True except Exception as e:
else: logger.warning(f"Failed to update .env file: {str(e)}, but will continue with in-memory update")
logger.warning( else:
f".env file not found at {env_path}, skipping file update but continuing with in-memory update" logger.warning(
) f".env file not found at {env_path}, skipping file update but continuing with in-memory update"
except Exception as env_err: )
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
# Update the settings in memory (this always happens) # Update the settings in memory (this always happens)
if refresh_token: if refresh_token:
@@ -494,7 +481,7 @@ async def save_google_drive_settings(
return { return {
"status": "success", "status": "success",
"message": "Google Drive settings have been saved", "message": "Google Drive settings have been saved",
"in_memory_only": not env_file_written, "in_memory_only": not os.path.exists(env_path),
} }
except Exception as e: except Exception as e:
-6
View File
@@ -17,7 +17,6 @@ from sqlalchemy.orm import Session
from app.database import get_db from app.database import get_db
from app.models import UserImapAccount from app.models import UserImapAccount
from app.utils.encryption import decrypt_value, encrypt_value 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.subscription import get_tier, get_user_tier_id
from app.utils.user_scope import get_current_owner_id from app.utils.user_scope import get_current_owner_id
@@ -188,11 +187,6 @@ def _test_imap_connection(host: str, port: int, username: str, password: str, us
Returns a dict with ``{"success": bool, "message": str}``. 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: try:
if use_ssl: if use_ssl:
mail = imaplib.IMAP4_SSL(host, port) mail = imaplib.IMAP4_SSL(host, port)
+12 -83
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.subscription import get_tier, get_user_tier_id
from app.utils.user_scope import get_current_owner_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__) logger = logging.getLogger(__name__)
router = APIRouter(prefix="/integrations", tags=["integrations"]) router = APIRouter(prefix="/integrations", tags=["integrations"])
@@ -515,12 +500,6 @@ def _test_imap_connection(config: dict[str, Any] | None, credentials: dict[str,
if not host or not username or not password: if not host or not username or not password:
return {"success": False, "message": "Missing required fields: host, username, and password"} return {"success": False, "message": "Missing required fields: host, username, and password"}
from app.utils.network import is_private_ip
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: try:
if use_ssl: if use_ssl:
mail = imaplib.IMAP4_SSL(host, port) mail = imaplib.IMAP4_SSL(host, port)
@@ -549,28 +528,17 @@ def _test_s3_connection(config: dict[str, Any] | None, credentials: dict[str, An
creds = credentials or {} creds = credentials or {}
bucket = cfg.get("bucket", "") bucket = cfg.get("bucket", "")
region = cfg.get("region", "us-east-1") region = cfg.get("region", "us-east-1")
endpoint_url = cfg.get("endpoint_url")
if not bucket: if not bucket:
return {"success": False, "message": "Missing required field: bucket"} return {"success": False, "message": "Missing required field: bucket"}
if endpoint_url:
from urllib.parse import urlparse
from app.utils.network import is_private_ip
parsed_url = urlparse(endpoint_url)
if parsed_url.hostname and is_private_ip(parsed_url.hostname):
logger.warning("SSRF blocked: Attempt to connect to private IP via S3 endpoint %s", endpoint_url)
return {"success": False, "message": "Connection error: Invalid endpoint URL or private IP"}
try: try:
client = boto3.client( client = boto3.client(
"s3", "s3",
region_name=region, region_name=region,
aws_access_key_id=creds.get("access_key_id", ""), aws_access_key_id=creds.get("access_key_id", ""),
aws_secret_access_key=creds.get("secret_access_key", ""), aws_secret_access_key=creds.get("secret_access_key", ""),
endpoint_url=endpoint_url, endpoint_url=cfg.get("endpoint_url"),
) )
client.head_bucket(Bucket=bucket) client.head_bucket(Bucket=bucket)
return {"success": True, "message": f"S3 bucket '{bucket}' is accessible"} return {"success": True, "message": f"S3 bucket '{bucket}' is accessible"}
@@ -582,50 +550,9 @@ def _test_s3_connection(config: dict[str, Any] | None, credentials: dict[str, An
return {"success": False, "message": "S3 connection failed"} 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]: 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.""" """Test a WebDAV/Nextcloud connection by issuing an HTTP PROPFIND."""
import httpx import urllib.request
cfg = config or {} cfg = config or {}
creds = credentials or {} creds = credentials or {}
@@ -652,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"} return {"success": False, "message": "URLs pointing to internal or private networks are not allowed"}
try: try:
auth = (username, password) if username and password else None import base64
headers = {"Depth": "0"}
# Use httpx for secure connection testing, avoiding urllib vulnerabilities req = urllib.request.Request(url, method="PROPFIND") # noqa: S310
resp = httpx.request("PROPFIND", url, auth=auth, headers=headers, timeout=10.0, follow_redirects=False) if username and password:
if resp.status_code < 400: token = base64.b64encode(f"{username}:{password}".encode()).decode()
return {"success": True, "message": "WebDAV connection successful"} req.add_header("Authorization", f"Basic {token}")
return {"success": False, "message": f"WebDAV returned HTTP {resp.status_code}"} 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 except Exception as exc: # noqa: BLE001
logger.warning("WebDAV connection error for %s: %s", hostname, exc) logger.warning("WebDAV connection error for %s: %s", hostname, exc)
return {"success": False, "message": "WebDAV connection failed — check URL and credentials"} return {"success": False, "message": "WebDAV connection failed — check URL and credentials"}
_CONNECTION_TESTERS: dict[str, Any] = { _CONNECTION_TESTERS: dict[str, Any] = {
IntegrationType.DROPBOX: _test_dropbox_connection,
IntegrationType.IMAP: _test_imap_connection, IntegrationType.IMAP: _test_imap_connection,
IntegrationType.S3: _test_s3_connection, IntegrationType.S3: _test_s3_connection,
IntegrationType.WEBDAV: _test_webdav_connection, IntegrationType.WEBDAV: _test_webdav_connection,
+9 -9
View File
@@ -101,9 +101,9 @@ async def signup_page(request: Request) -> Any:
if not settings.allow_local_signup: if not settings.allow_local_signup:
return RedirectResponse(url="/login?error=Registration+is+not+enabled", status_code=302) return RedirectResponse(url="/login?error=Registration+is+not+enabled", status_code=302)
return templates.TemplateResponse( return templates.TemplateResponse(
request,
"signup.html", "signup.html",
context={ {
"request": request,
"csrf_token": getattr(request.state, "csrf_token", ""), "csrf_token": getattr(request.state, "csrf_token", ""),
"app_version": settings.version, "app_version": settings.version,
}, },
@@ -113,16 +113,16 @@ async def signup_page(request: Request) -> Any:
@router.get("/verify-email-sent", include_in_schema=False) @router.get("/verify-email-sent", include_in_schema=False)
async def verify_email_sent_page(request: Request) -> Any: async def verify_email_sent_page(request: Request) -> Any:
"""Render the verify-email-sent confirmation page.""" """Render the verify-email-sent confirmation page."""
return templates.TemplateResponse(request, "verify_email_sent.html") return templates.TemplateResponse("verify_email_sent.html", {"request": request})
@router.get("/forgot-username", include_in_schema=False) @router.get("/forgot-username", include_in_schema=False)
async def forgot_username_page(request: Request) -> Any: async def forgot_username_page(request: Request) -> Any:
"""Render the forgot-username page where users can request a username reminder email.""" """Render the forgot-username page where users can request a username reminder email."""
return templates.TemplateResponse( return templates.TemplateResponse(
request,
"forgot_username.html", "forgot_username.html",
context={ {
"request": request,
"csrf_token": getattr(request.state, "csrf_token", ""), "csrf_token": getattr(request.state, "csrf_token", ""),
"app_version": settings.version, "app_version": settings.version,
}, },
@@ -133,9 +133,9 @@ async def forgot_username_page(request: Request) -> Any:
async def forgot_password_page(request: Request) -> Any: async def forgot_password_page(request: Request) -> Any:
"""Render the forgot-password page where users can request a reset email.""" """Render the forgot-password page where users can request a reset email."""
return templates.TemplateResponse( return templates.TemplateResponse(
request,
"forgot_password.html", "forgot_password.html",
context={ {
"request": request,
"csrf_token": getattr(request.state, "csrf_token", ""), "csrf_token": getattr(request.state, "csrf_token", ""),
"app_version": settings.version, "app_version": settings.version,
}, },
@@ -147,9 +147,9 @@ async def reset_password_page(request: Request) -> Any:
"""Render the password reset form page.""" """Render the password reset form page."""
token = request.query_params.get("token", "") token = request.query_params.get("token", "")
return templates.TemplateResponse( return templates.TemplateResponse(
request,
"password_reset_form.html", "password_reset_form.html",
context={ {
"request": request,
"token": token, "token": token,
"csrf_token": getattr(request.state, "csrf_token", ""), "csrf_token": getattr(request.state, "csrf_token", ""),
"app_version": settings.version, "app_version": settings.version,
+8 -23
View File
@@ -120,7 +120,6 @@ class WhoAmIResponse(BaseModel):
email: str | None email: str | None
avatar_url: str | None avatar_url: str | None
is_admin: bool is_admin: bool
preferred_language: str | None
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -274,44 +273,31 @@ async def list_devices(
return [_device_to_response(d) for d in 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 @require_login
async def deactivate_device( async def deactivate_device(
request: Request, request: Request,
device_id: int, device_id: int,
owner_id: CurrentOwner, owner_id: CurrentOwner,
db: DbSession, db: DbSession,
) -> dict[str, str]: ) -> None:
"""Deactivate or permanently delete a push-notification device registration. """Deactivate a push-notification device registration.
* **Active device** soft-deactivated: the record is kept for audit The device record is kept for audit purposes but will no longer receive
purposes but will no longer receive push notifications. push notifications.
* **Already-inactive device** hard-deleted: the record is permanently
removed from the database.
""" """
device = db.get(MobileDevice, device_id) device = db.get(MobileDevice, device_id)
if not device or device.owner_id != owner_id: if not device or device.owner_id != owner_id:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Device not found") raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Device not found")
if device.is_active: device.is_active = False
device.is_active = False
try:
db.commit()
except Exception:
db.rollback()
raise
logger.info("Mobile device deactivated: id=%s owner=%s", device_id, owner_id)
return {"detail": "Device deactivated"}
# Hard-delete an already-inactive device.
try: try:
db.delete(device)
db.commit() db.commit()
except Exception: except Exception:
db.rollback() db.rollback()
raise 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) @router.get("/whoami", response_model=WhoAmIResponse)
@@ -358,5 +344,4 @@ async def whoami(
"email": email, "email": email,
"avatar_url": avatar_url, "avatar_url": avatar_url,
"is_admin": is_admin, "is_admin": is_admin,
"preferred_language": profile.preferred_language if profile else None,
} }
+10 -9
View File
@@ -452,16 +452,17 @@ async def update_preferences(
) )
try: try:
# Pre-fetch existing preferences for this user to avoid N+1 queries
existing_prefs = (
db.query(UserNotificationPreference).filter(UserNotificationPreference.owner_id == owner_id).all()
)
# Build a fast lookup dictionary keyed by (event_type, channel_type, target_id)
prefs_dict = {(pref.event_type, pref.channel_type, pref.target_id): pref for pref in existing_prefs}
for item in body.preferences: for item in body.preferences:
existing = prefs_dict.get((item.event_type, item.channel_type, item.target_id)) existing = (
db.query(UserNotificationPreference)
.filter(
UserNotificationPreference.owner_id == owner_id,
UserNotificationPreference.event_type == item.event_type,
UserNotificationPreference.channel_type == item.channel_type,
UserNotificationPreference.target_id == item.target_id,
)
.first()
)
if existing: if existing:
existing.is_enabled = item.is_enabled existing.is_enabled = item.is_enabled
else: else:
+113 -150
View File
@@ -3,10 +3,10 @@ OneDrive API endpoints
""" """
import logging import logging
import os
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Annotated, Optional from typing import Annotated, Optional
import httpx
import requests import requests
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -14,7 +14,6 @@ from sqlalchemy.orm import Session
from app.auth import require_login from app.auth import require_login
from app.config import settings from app.config import settings
from app.database import get_db from app.database import get_db
from app.utils.env_utils import update_env_file
from app.utils.oauth_helper import exchange_oauth_token from app.utils.oauth_helper import exchange_oauth_token
from app.utils.settings_service import save_setting_to_db from app.utils.settings_service import save_setting_to_db
from app.utils.settings_sync import notify_settings_updated from app.utils.settings_sync import notify_settings_updated
@@ -25,17 +24,6 @@ logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
def _require_admin(request: Request) -> dict:
"""Ensure the caller is an admin. Raises 403 otherwise."""
user = request.session.get("user")
if not user or not user.get("is_admin"):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
return user
AdminUser = Annotated[dict, Depends(_require_admin)]
@router.post("/onedrive/exchange-token") @router.post("/onedrive/exchange-token")
@require_login @require_login
async def exchange_onedrive_token( async def exchange_onedrive_token(
@@ -68,7 +56,6 @@ async def exchange_onedrive_token(
# Return just what's needed by the frontend # Return just what's needed by the frontend
return { return {
"refresh_token": token_data["refresh_token"], "refresh_token": token_data["refresh_token"],
"access_token": token_data.get("access_token", ""),
"expires_in": token_data.get("expires_in", 3600), "expires_in": token_data.get("expires_in", 3600),
} }
@@ -105,18 +92,17 @@ async def test_onedrive_token(request: Request):
"scope": "offline_access Files.ReadWrite", "scope": "offline_access Files.ReadWrite",
} }
async with httpx.AsyncClient(timeout=settings.http_request_timeout) as client: response = requests.post(token_url, data=refresh_data, timeout=settings.http_request_timeout)
response = await client.post(token_url, data=refresh_data)
if response.status_code != 200: if response.status_code != 200:
logger.error(f"Failed to refresh OneDrive token: {response.text}") logger.error(f"Failed to refresh OneDrive token: {response.text}")
return { return {
"status": "error", "status": "error",
"message": "Refresh token has expired or is invalid", "message": "Refresh token has expired or is invalid",
"needs_reauth": True, "needs_reauth": True,
} }
token_data = response.json() token_data = response.json()
access_token = token_data.get("access_token") access_token = token_data.get("access_token")
expires_in = token_data.get("expires_in", 3600) # Default to 1 hour if not specified expires_in = token_data.get("expires_in", 3600) # Default to 1 hour if not specified
@@ -129,7 +115,32 @@ async def test_onedrive_token(request: Request):
settings.onedrive_refresh_token = new_refresh_token settings.onedrive_refresh_token = new_refresh_token
# Also try to update .env file if it exists # Also try to update .env file if it exists
update_env_file({"ONEDRIVE_REFRESH_TOKEN": new_refresh_token}) try:
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
if os.path.exists(env_path):
with open(env_path, "r") as f:
env_lines = f.readlines()
updated_lines = []
updated = False
for line in env_lines:
if line.startswith("ONEDRIVE_REFRESH_TOKEN="):
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
updated = True
else:
updated_lines.append(line)
if not updated:
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
with open(env_path, "w") as f:
f.writelines(updated_lines)
logger.info("Updated refresh token in .env file")
except Exception as e:
logger.warning(f"Failed to update refresh token in .env file: {e}")
# Persist the rotated refresh token to the database # Persist the rotated refresh token to the database
try: try:
@@ -153,18 +164,17 @@ async def test_onedrive_token(request: Request):
user_info_url = "https://graph.microsoft.com/v1.0/me" user_info_url = "https://graph.microsoft.com/v1.0/me"
headers = {"Authorization": f"Bearer {access_token}"} headers = {"Authorization": f"Bearer {access_token}"}
async with httpx.AsyncClient(timeout=settings.http_request_timeout) as client: user_response = requests.get(user_info_url, headers=headers, timeout=settings.http_request_timeout)
user_response = await client.get(user_info_url, headers=headers)
if user_response.status_code != 200: if user_response.status_code != 200:
logger.error(f"OneDrive token test failed: {user_response.status_code} {user_response.text}") logger.error(f"OneDrive token test failed: {user_response.status_code} {user_response.text}")
return { return {
"status": "error", "status": "error",
"message": f"Token validation failed with status {user_response.status_code}: {user_response.text}", "message": f"Token validation failed with status {user_response.status_code}: {user_response.text}",
} }
# Get user info # Get user info
user_info = user_response.json() user_info = user_response.json()
display_name = user_info.get("displayName", "Unknown user") display_name = user_info.get("displayName", "Unknown user")
email = user_info.get("userPrincipalName", "Unknown email") email = user_info.get("userPrincipalName", "Unknown email")
@@ -196,102 +206,6 @@ async def test_onedrive_token(request: Request):
return {"status": "error", "message": f"Connection error: {str(e)}"} 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): def format_time_remaining(time_delta):
"""Format a timedelta into a human-readable string.""" """Format a timedelta into a human-readable string."""
if time_delta.total_seconds() <= 0: if time_delta.total_seconds() <= 0:
@@ -313,15 +227,15 @@ def format_time_remaining(time_delta):
@router.post("/onedrive/save-settings") @router.post("/onedrive/save-settings")
@require_login
async def save_onedrive_settings( async def save_onedrive_settings(
request: Request, request: Request,
refresh_token: Annotated[str, Form(...)], refresh_token: Annotated[str, Form(...)],
_admin: AdminUser,
db: Session = Depends(get_db),
client_id: Annotated[Optional[str], Form()] = None, client_id: Annotated[Optional[str], Form()] = None,
client_secret: Annotated[Optional[str], Form()] = None, client_secret: Annotated[Optional[str], Form()] = None,
tenant_id: Annotated[str, Form()] = "common", tenant_id: Annotated[str, Form()] = "common",
folder_path: Annotated[Optional[str], Form()] = None, folder_path: Annotated[Optional[str], Form()] = None,
db: Session = Depends(get_db),
): ):
""" """
Saves to database (primary) and .env file (best-effort). Saves to database (primary) and .env file (best-effort).
@@ -332,26 +246,75 @@ async def save_onedrive_settings(
user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "wizard" user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "wizard"
) )
# Build settings dictionary mapped to database/memory keys # Best-effort .env file write
onedrive_settings = { try:
"onedrive_refresh_token": refresh_token, env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
"onedrive_client_id": client_id, if not os.path.exists(env_path):
"onedrive_client_secret": client_secret, logger.warning(f".env file not found at {env_path}, skipping file write")
"onedrive_tenant_id": tenant_id, else:
"onedrive_folder_path": folder_path, logger.info(f"Updating OneDrive settings in {env_path}")
}
# Filter out None values with open(env_path, "r") as f:
onedrive_settings = {k: v for k, v in onedrive_settings.items() if v is not None} env_lines = f.readlines()
# Best-effort .env file write using the new utility onedrive_settings = {"ONEDRIVE_REFRESH_TOKEN": refresh_token}
env_settings = {k.upper(): v for k, v in onedrive_settings.items()} if client_id:
update_env_file(env_settings) onedrive_settings["ONEDRIVE_CLIENT_ID"] = client_id
if client_secret:
onedrive_settings["ONEDRIVE_CLIENT_SECRET"] = client_secret
if tenant_id:
onedrive_settings["ONEDRIVE_TENANT_ID"] = tenant_id
if folder_path:
onedrive_settings["ONEDRIVE_FOLDER_PATH"] = folder_path
# Update in-memory settings and persist to database dynamically updated = set()
for key, value in onedrive_settings.items(): new_env_lines = []
setattr(settings, key, value) for line in env_lines:
save_setting_to_db(db, key, value, changed_by=changed_by) stripped_line = line.rstrip()
is_updated = False
for key, value in onedrive_settings.items():
if stripped_line.startswith(f"{key}=") or stripped_line.startswith(f"# {key}="):
new_env_lines.append(f"{key}={value}")
updated.add(key)
is_updated = True
break
if not is_updated:
new_env_lines.append(stripped_line)
for key, value in onedrive_settings.items():
if key not in updated:
new_env_lines.append(f"{key}={value}")
with open(env_path, "w") as f:
f.write("\n".join(new_env_lines) + "\n")
logger.info("Successfully updated OneDrive settings in .env file")
except Exception as env_err:
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
# Update the settings in memory
if refresh_token:
settings.onedrive_refresh_token = refresh_token
if client_id:
settings.onedrive_client_id = client_id
if client_secret:
settings.onedrive_client_secret = client_secret
if tenant_id:
settings.onedrive_tenant_id = tenant_id
if folder_path:
settings.onedrive_folder_path = folder_path
# Persist to database (primary)
if refresh_token:
save_setting_to_db(db, "onedrive_refresh_token", refresh_token, changed_by=changed_by)
if client_id:
save_setting_to_db(db, "onedrive_client_id", client_id, changed_by=changed_by)
if client_secret:
save_setting_to_db(db, "onedrive_client_secret", client_secret, changed_by=changed_by)
if tenant_id:
save_setting_to_db(db, "onedrive_tenant_id", tenant_id, changed_by=changed_by)
if folder_path:
save_setting_to_db(db, "onedrive_folder_path", folder_path, changed_by=changed_by)
notify_settings_updated() notify_settings_updated()
+2 -8
View File
@@ -117,14 +117,8 @@ PIPELINE_STEP_TYPES: dict[str, dict[str, Any]] = {
}, },
"classify": { "classify": {
"label": "Document Classification", "label": "Document Classification",
"description": "Classify the document type using built-in and custom rules (filename patterns, content keywords, metadata matching).", "description": "Classify the document type using AI without full metadata extraction.",
"config_schema": { "config_schema": {},
"use_builtin_rules": {
"type": "boolean",
"default": True,
"description": "Include the pre-built classification rules (invoice, contract, receipt, etc.).",
},
},
}, },
} }
+2 -11
View File
@@ -196,20 +196,11 @@ def seed_plans(db: DbSession, _admin: AdminUser) -> dict[str, Any]:
def reorder_plans(body: ReorderBody, db: DbSession, _admin: AdminUser) -> dict[str, Any]: def reorder_plans(body: ReorderBody, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
"""Update sort_order for each plan_id in *body.order* (position = index in list).""" """Update sort_order for each plan_id in *body.order* (position = index in list)."""
updated = 0 updated = 0
for sort_order, plan_id in enumerate(body.order):
# Fetch all requested plans in a single query to avoid N+1 plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
plan_ids = body.order
plans = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id.in_(plan_ids)).all()
# Build a map for fast O(1) lookup
plan_map = {p.plan_id: p for p in plans}
for sort_order, plan_id in enumerate(plan_ids):
plan = plan_map.get(plan_id)
if plan: if plan:
plan.sort_order = sort_order plan.sort_order = sort_order
updated += 1 updated += 1
try: try:
db.commit() db.commit()
except Exception: except Exception:
-18
View File
@@ -99,8 +99,6 @@ class ProfileResponse(BaseModel):
contact_email: str | None contact_email: str | None
preferred_language: str | None preferred_language: str | None
preferred_theme: str | None preferred_theme: str | None
default_document_language: str | None
"""ISO 639-1 code for the user's preferred document translation target language."""
avatar_url: str avatar_url: str
"""Gravatar URL or ``data:`` URI for a custom uploaded avatar.""" """Gravatar URL or ``data:`` URI for a custom uploaded avatar."""
is_local_user: bool is_local_user: bool
@@ -114,10 +112,6 @@ class ProfileUpdateRequest(BaseModel):
contact_email: str | None = Field(default=None, max_length=255, description="Contact / notification e-mail") contact_email: str | None = Field(default=None, max_length=255, description="Contact / notification e-mail")
preferred_language: str | None = Field(default=None, description="ISO 639-1 language code, e.g. 'en', 'de'") preferred_language: str | None = Field(default=None, description="ISO 639-1 language code, e.g. 'en', 'de'")
preferred_theme: str | None = Field(default=None, description="Colour scheme: 'light', 'dark', or 'system'") preferred_theme: str | None = Field(default=None, description="Colour scheme: 'light', 'dark', or 'system'")
default_document_language: str | None = Field(
default=None,
description="ISO 639-1 code for the default document translation target language, e.g. 'en', 'de'",
)
class ChangePasswordRequest(BaseModel): class ChangePasswordRequest(BaseModel):
@@ -155,7 +149,6 @@ async def get_profile(request: Request, db: DbSession) -> ProfileResponse:
contact_email=profile.contact_email, # type: ignore[arg-type] contact_email=profile.contact_email, # type: ignore[arg-type]
preferred_language=profile.preferred_language, # type: ignore[arg-type] preferred_language=profile.preferred_language, # type: ignore[arg-type]
preferred_theme=profile.preferred_theme, # type: ignore[arg-type] preferred_theme=profile.preferred_theme, # type: ignore[arg-type]
default_document_language=profile.default_document_language, # type: ignore[arg-type]
avatar_url=avatar_url, avatar_url=avatar_url,
is_local_user=is_local, is_local_user=is_local,
) )
@@ -208,16 +201,6 @@ async def update_profile(
) )
profile.preferred_theme = theme or None # type: ignore[assignment] profile.preferred_theme = theme or None # type: ignore[assignment]
# Validate default document language
if body.default_document_language is not None:
doc_lang = body.default_document_language.lower().strip()
if doc_lang and doc_lang not in SUPPORTED_LANGUAGE_CODES:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Unsupported language code: {doc_lang}",
)
profile.default_document_language = doc_lang or None # type: ignore[assignment]
if body.display_name is not None: if body.display_name is not None:
profile.display_name = body.display_name.strip() or None # type: ignore[assignment] profile.display_name = body.display_name.strip() or None # type: ignore[assignment]
@@ -242,7 +225,6 @@ async def update_profile(
contact_email=profile.contact_email, # type: ignore[arg-type] contact_email=profile.contact_email, # type: ignore[arg-type]
preferred_language=profile.preferred_language, # type: ignore[arg-type] preferred_language=profile.preferred_language, # type: ignore[arg-type]
preferred_theme=profile.preferred_theme, # type: ignore[arg-type] preferred_theme=profile.preferred_theme, # type: ignore[arg-type]
default_document_language=profile.default_document_language, # type: ignore[arg-type]
avatar_url=avatar_url, avatar_url=avatar_url,
is_local_user=is_local, is_local_user=is_local,
) )
-249
View File
@@ -1,249 +0,0 @@
"""QR code login API endpoints for mobile app authentication.
Provides a secure challenge-response flow for logging into the mobile app
by scanning a QR code displayed in the web interface:
1. **Web user** calls ``POST /qr-auth/challenge`` receives a time-limited
challenge token (encoded in the QR code).
2. **Web UI** polls ``GET /qr-auth/challenge/{id}/status`` to detect when
the mobile app has claimed the challenge.
3. **Mobile app** scans the QR code and calls ``POST /qr-auth/claim`` with
the challenge token + device name receives an API token.
Security properties:
* Challenges expire after a configurable TTL (default 2 minutes).
* Single-use: once claimed, a challenge cannot be reused (replay-safe).
* Cryptographically random 64-byte tokens.
* IP addresses are logged for audit.
"""
from __future__ import annotations
import base64
import io
import logging
from datetime import datetime
from typing import Annotated, Any
import segno
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from app.auth import require_login
from app.config import settings
from app.database import get_db
from app.middleware.audit_log import get_client_ip
from app.utils.session_manager import (
claim_qr_challenge,
create_qr_challenge,
get_challenge_status,
)
from app.utils.user_scope import get_current_owner_id
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/qr-auth", tags=["qr-auth"])
DbSession = Annotated[Session, Depends(get_db)]
# ---------------------------------------------------------------------------
# Auth helper
# ---------------------------------------------------------------------------
def _get_owner_id(request: Request) -> str:
"""Return the current user's owner ID, raising 401 if unauthenticated."""
owner_id = get_current_owner_id(request)
if not owner_id:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
return owner_id
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
# ---------------------------------------------------------------------------
# Request / Response schemas
# ---------------------------------------------------------------------------
class CreateChallengeResponse(BaseModel):
"""Response after creating a QR login challenge."""
challenge_id: int
challenge_token: str
expires_at: datetime
ttl_seconds: int = Field(description="Seconds until the challenge expires (use for client-side countdown).")
qr_payload: str = Field(description="The string to encode in the QR code.")
qr_code_svg: str = Field(description="Base64-encoded SVG data URI of the QR code, ready for use in an <img> src.")
class ChallengeStatusResponse(BaseModel):
"""Response for polling the status of a QR challenge."""
id: int
status: str # "pending", "claimed", "expired", "cancelled"
device_name: str | None = None
claimed_at: datetime | None = None
expires_at: datetime
class ClaimChallengeRequest(BaseModel):
"""Request body for claiming a QR login challenge."""
challenge_token: str = Field(min_length=1, max_length=256)
device_name: str = Field(
default="Mobile App",
min_length=1,
max_length=120,
description="Human-readable device name.",
)
class ClaimChallengeResponse(BaseModel):
"""Response after successfully claiming a QR challenge."""
token: str
token_id: int
name: str
owner_id: str
created_at: datetime
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
# QR code rendering parameters
_QR_ERROR_LEVEL = "M" # Medium error correction (~15% recovery); sufficient for on-screen display
_QR_SCALE = 4 # Each QR module is rendered as 4×4 SVG pixels
def _generate_qr_svg(payload: str) -> str:
"""Generate a QR code for *payload* and return it as a base64 SVG data URI.
Using ``segno`` (pure-Python, no Pillow dependency) and SVG output so the
QR code scales crisply at any resolution without requiring a canvas or any
client-side JavaScript library.
"""
qr = segno.make(payload, error=_QR_ERROR_LEVEL)
buf = io.BytesIO()
qr.save(buf, kind="svg", scale=_QR_SCALE, xmldecl=False, svgclass=None, lineclass=None, omitsize=True)
svg_bytes = buf.getvalue()
return "data:image/svg+xml;base64," + base64.b64encode(svg_bytes).decode("ascii")
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.post("/challenge", status_code=status.HTTP_201_CREATED, response_model=CreateChallengeResponse)
@require_login
async def create_challenge(
request: Request,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, Any]:
"""Create a new QR login challenge.
The returned ``qr_payload`` should be encoded into a QR code and
displayed to the user. The mobile app scans this QR code and
calls the ``/claim`` endpoint.
"""
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)
# The QR payload is a JSON-like string with enough info for the mobile
# app to know the server URL and challenge token.
base_url = str(request.base_url).rstrip("/")
qr_payload = f"docuelevate://qr-login?token={challenge.challenge_token}&server={base_url}"
# Compute the TTL in seconds so the client can run a countdown timer
# without comparing absolute timestamps (which breaks when client and
# server clocks are out of sync).
ttl_seconds = max(0, int((challenge.expires_at - challenge.created_at).total_seconds()))
return {
"challenge_id": challenge.id,
"challenge_token": challenge.challenge_token,
"expires_at": challenge.expires_at,
"ttl_seconds": ttl_seconds,
"qr_payload": qr_payload,
"qr_code_svg": _generate_qr_svg(qr_payload),
}
@router.get("/challenge/{challenge_id}/status", response_model=ChallengeStatusResponse)
@require_login
async def poll_challenge_status(
request: Request,
challenge_id: int,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, Any]:
"""Poll the status of a QR login challenge.
The web UI calls this endpoint every few seconds to check if the
mobile app has scanned the QR code and claimed the challenge.
"""
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")
return result
@router.post("/claim", response_model=ClaimChallengeResponse)
async def claim_challenge(
request: Request,
body: ClaimChallengeRequest,
db: DbSession,
) -> dict[str, Any]:
"""Claim a QR login challenge and receive an API token.
This endpoint is called by the mobile app after scanning a QR code.
It does **not** require authentication the challenge token itself
serves as proof that the user authorized this login from their web
session.
"""
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)
if not result:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid, expired, or already claimed challenge.",
)
try:
from app.utils.audit_service import record_event
record_event(
db,
action="qr_login_claimed",
user=result["owner_id"],
resource_type="session",
ip_address=ip,
details={"device_name": body.device_name, "token_id": result["token_id"]},
severity="info",
)
except Exception:
logger.debug("Failed to write QR login audit event", exc_info=True)
return result
-196
View File
@@ -1,196 +0,0 @@
"""API endpoints for managing user sessions.
Provides endpoints for listing active sessions, revoking individual sessions,
and the "log off everywhere" feature that invalidates all sessions and API
tokens across all devices.
"""
from __future__ import annotations
import logging
from datetime import datetime
from typing import Annotated, Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.auth import require_login
from app.database import get_db
from app.middleware.audit_log import get_client_ip
from app.utils.session_manager import (
get_session_lifetime_days,
list_user_sessions,
revoke_all_sessions,
revoke_session,
)
from app.utils.user_scope import get_current_owner_id
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/sessions", tags=["sessions"])
DbSession = Annotated[Session, Depends(get_db)]
# ---------------------------------------------------------------------------
# Auth helper
# ---------------------------------------------------------------------------
def _get_owner_id(request: Request) -> str:
"""Return the current user's owner ID, raising 401 if unauthenticated."""
owner_id = get_current_owner_id(request)
if not owner_id:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
return owner_id
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
# ---------------------------------------------------------------------------
# Response schemas
# ---------------------------------------------------------------------------
class SessionResponse(BaseModel):
"""Serialised user session for the management UI."""
id: int
device_info: str | None
ip_address: str | None
created_at: datetime
last_active_at: datetime
expires_at: datetime
is_current: bool = False
class SessionListResponse(BaseModel):
"""Response for listing active sessions."""
sessions: list[SessionResponse]
session_lifetime_days: int
class RevokeAllResponse(BaseModel):
"""Response after revoking all sessions."""
revoked_count: int
message: str
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.get("/", response_model=SessionListResponse)
@require_login
async def list_sessions(
request: Request,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, Any]:
"""List all active sessions for the current user."""
sessions = list_user_sessions(db, owner_id)
# Determine which session is the current one
current_token = request.session.get("_session_token")
session_list = []
for s in sessions:
session_list.append(
{
"id": s.id,
"device_info": s.device_info,
"ip_address": s.ip_address,
"created_at": s.created_at,
"last_active_at": s.last_active_at,
"expires_at": s.expires_at,
"is_current": s.session_token == current_token if current_token else False,
}
)
return {
"sessions": session_list,
"session_lifetime_days": get_session_lifetime_days(),
}
@router.delete("/{session_id}", status_code=status.HTTP_204_NO_CONTENT)
@require_login
async def revoke_single_session(
request: Request,
session_id: int,
owner_id: CurrentOwner,
db: DbSession,
) -> None:
"""Revoke a specific session by ID."""
success = revoke_session(db, session_id, owner_id)
if not success:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
try:
from app.utils.audit_service import record_event
record_event(
db,
action="session_revoked",
user=owner_id,
resource_type="session",
resource_id=str(session_id),
ip_address=get_client_ip(request),
severity="info",
)
except Exception:
logger.debug("Failed to write session revocation audit event", exc_info=True)
@router.post("/revoke-all", response_model=RevokeAllResponse)
@require_login
async def revoke_all(
request: Request,
owner_id: CurrentOwner,
db: DbSession,
) -> dict[str, Any]:
"""Revoke all sessions except the current one ("log off everywhere").
Also revokes all active API tokens for the user, which invalidates
mobile app sessions and any programmatic access.
"""
# Find current session to preserve it
current_token = request.session.get("_session_token")
current_session_id = None
if current_token:
from app.models import UserSession
current = db.query(UserSession).filter(UserSession.session_token == current_token).first()
if current:
current_session_id = current.id
count = revoke_all_sessions(
db,
owner_id,
except_session_id=current_session_id,
revoke_api_tokens=True,
)
try:
from app.utils.audit_service import record_event
record_event(
db,
action="revoke_all_sessions",
user=owner_id,
resource_type="session",
ip_address=get_client_ip(request),
details={"revoked_count": count},
severity="warning",
)
except Exception:
logger.debug("Failed to write revoke-all audit event", exc_info=True)
return {
"revoked_count": count,
"message": f"Successfully revoked {count} session(s) and all API tokens.",
}
-62
View File
@@ -55,12 +55,6 @@ class SettingUpdate(BaseModel):
value: Optional[str] = Field(None, description="Setting value (None to delete)") 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): class SettingResponse(BaseModel):
"""Model for setting response""" """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}") @router.delete("/{key}")
async def delete_setting(key: str, request: Request, db: DbSession, admin: AdminUser): async def delete_setting(key: str, request: Request, db: DbSession, admin: AdminUser):
""" """
+5 -7
View File
@@ -313,18 +313,16 @@ async def list_shared_links(
active_only: bool = Query(False, description="When true, only return active (non-revoked) links"), active_only: bool = Query(False, description="When true, only return active (non-revoked) links"),
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""List all shared links created by the authenticated user.""" """List all shared links created by the authenticated user."""
q = ( q = db.query(SharedLink).filter(SharedLink.owner_id == owner_id)
db.query(SharedLink, FileRecord.original_filename)
.outerjoin(FileRecord, SharedLink.file_id == FileRecord.id)
.filter(SharedLink.owner_id == owner_id)
)
if active_only: if active_only:
q = q.filter(SharedLink.is_active.is_(True)) q = q.filter(SharedLink.is_active.is_(True))
links_with_filenames = q.order_by(SharedLink.created_at.desc()).all() links = q.order_by(SharedLink.created_at.desc()).all()
base_url = str(request.base_url).rstrip("/") base_url = str(request.base_url).rstrip("/")
result = [] result = []
for link, filename in links_with_filenames: for link in links:
file_record = db.query(FileRecord).filter(FileRecord.id == link.file_id).first()
filename = file_record.original_filename if file_record else None
result.append(_link_to_dict(link, base_url, filename)) result.append(_link_to_dict(link, base_url, filename))
return result return result
-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
+3 -13
View File
@@ -9,7 +9,7 @@ Public endpoints:
""" """
import logging import logging
from datetime import datetime, time, timedelta, timezone from datetime import datetime, timezone
from typing import Annotated, Any from typing import Annotated, Any
from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
@@ -207,29 +207,19 @@ def platform_stats(request: Request, db: DbSession, _admin: AdminUser) -> dict[s
from app.models import FileRecord, UserProfile from app.models import FileRecord, UserProfile
today = datetime.now(timezone.utc).date() today = datetime.now(timezone.utc).date()
day_start = datetime.combine(today, time.min, tzinfo=timezone.utc)
day_end = day_start + timedelta(days=1)
month_start = day_start.replace(day=1)
if month_start.month == 12:
month_end = month_start.replace(year=month_start.year + 1, month=1)
else:
month_end = month_start.replace(month=month_start.month + 1)
# Total files # Total files
total_files: int = db.query(func.count(FileRecord.id)).scalar() or 0 total_files: int = db.query(func.count(FileRecord.id)).scalar() or 0
# Files today # Files today
files_today: int = ( files_today: int = (
db.query(func.count(FileRecord.id)) db.query(func.count(FileRecord.id)).filter(func.date(FileRecord.created_at) == today).scalar() or 0
.filter(FileRecord.created_at >= day_start, FileRecord.created_at < day_end)
.scalar()
or 0
) )
# Files this month # Files this month
files_this_month: int = ( files_this_month: int = (
db.query(func.count(FileRecord.id)) db.query(func.count(FileRecord.id))
.filter(FileRecord.created_at >= month_start, FileRecord.created_at < month_end) .filter(func.strftime("%Y-%m", FileRecord.created_at) == today.strftime("%Y-%m"))
.scalar() .scalar()
or 0 or 0
) )
-124
View File
@@ -1,124 +0,0 @@
"""
System reset API endpoints for DocuElevate.
Provides admin-only REST endpoints for:
- Full system reset (wipe all user data)
- Reset with re-import (move originals reimport folder, wipe, re-ingest)
Both operations require the ``ENABLE_FACTORY_RESET=True`` feature flag and
admin privileges.
"""
import logging
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.config import settings
from app.database import get_db
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/admin/system-reset", tags=["system-reset"])
def _require_admin(request: Request) -> dict:
"""Ensure the caller is an admin. Raises 403 otherwise."""
user = request.session.get("user")
if not user or not user.get("is_admin"):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
return user
AdminUser = Annotated[dict, Depends(_require_admin)]
def _require_feature_enabled() -> None:
"""Raise 404 when the factory-reset feature flag is off."""
if not settings.enable_factory_reset:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="System reset is not enabled. Set ENABLE_FACTORY_RESET=True to activate.",
)
class ResetRequest(BaseModel):
"""Body for system reset endpoints. Requires explicit confirmation."""
confirmation: str
@router.post("/full")
async def full_reset(
body: ResetRequest,
_admin: AdminUser,
db: Session = Depends(get_db),
) -> dict:
"""Wipe all user data (database + work-files).
The caller must send ``{"confirmation": "DELETE"}`` to proceed.
"""
_require_feature_enabled()
if body.confirmation != "DELETE":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail='Confirmation required: send {"confirmation": "DELETE"} to proceed.',
)
from app.utils.system_reset import perform_full_reset
try:
result = perform_full_reset(db)
except Exception as exc:
logger.exception("Full system reset failed")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"System reset failed: {exc}",
) from exc
return {"status": "ok", "result": result}
@router.post("/reimport")
async def reset_and_reimport(
body: ResetRequest,
_admin: AdminUser,
db: Session = Depends(get_db),
) -> dict:
"""Move original files to a reimport folder, wipe everything, and
configure the reimport folder as a watch folder for automatic
re-ingestion.
The caller must send ``{"confirmation": "REIMPORT"}`` to proceed.
"""
_require_feature_enabled()
if body.confirmation != "REIMPORT":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail='Confirmation required: send {"confirmation": "REIMPORT"} to proceed.',
)
from app.utils.system_reset import perform_reset_and_reimport
try:
result = perform_reset_and_reimport(db)
except Exception as exc:
logger.exception("Reset-and-reimport failed")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Reset and reimport failed: {exc}",
) from exc
return {"status": "ok", "result": result}
@router.get("/status")
async def reset_status(_admin: AdminUser) -> dict:
"""Return whether the system reset feature is enabled."""
return {
"enabled": settings.enable_factory_reset,
"factory_reset_on_startup": settings.factory_reset_on_startup,
}
-156
View File
@@ -1,156 +0,0 @@
"""
API endpoints for document translation.
Provides on-the-fly translation via the AI provider and access to the
persisted default-language translation.
"""
import logging
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from app.auth import require_login
from app.config import settings
from app.database import get_db
from app.models import FileRecord
from app.utils.ai_provider import get_ai_provider
from app.utils.user_scope import apply_owner_filter
logger = logging.getLogger(__name__)
router = APIRouter()
DbSession = Annotated[Session, Depends(get_db)]
# Maximum characters sent to the AI provider for a single translation request.
_MAX_TRANSLATION_INPUT = 50_000
def _get_file_or_404(db: Session, file_id: int, request: Request) -> FileRecord:
"""Fetch a FileRecord visible to the current user or raise 404."""
query = db.query(FileRecord).filter(FileRecord.id == file_id)
query = apply_owner_filter(query, request)
record = query.first()
if not record:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
return record
@router.get("/files/{file_id}/translation/default")
@require_login
def get_default_translation(
request: Request,
file_id: int,
db: DbSession,
) -> JSONResponse:
"""Return the persisted default-language translation for a document.
Returns 404 if no default-language translation has been generated yet
(e.g. because the document is already in the default language).
"""
record = _get_file_or_404(db, file_id, request)
if not record.default_language_text:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="No default-language translation available for this file",
)
return JSONResponse(
content={
"file_id": record.id,
"detected_language": record.detected_language,
"default_language_code": record.default_language_code,
"text": record.default_language_text,
}
)
@router.get("/files/{file_id}/translate")
@require_login
def translate_on_the_fly(
request: Request,
file_id: int,
db: DbSession,
lang: str = Query(..., min_length=2, max_length=10, description="Target language ISO 639-1 code"),
) -> JSONResponse:
"""Translate a document's extracted text into an arbitrary language on the fly.
The translation is generated via the configured AI provider and is **not**
persisted. For the default-language translation, use the
``/files/{file_id}/translation/default`` endpoint instead.
"""
record = _get_file_or_404(db, file_id, request)
source_text = record.ocr_text
if not source_text:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="No extracted text available for this file — translation requires OCR text",
)
# If the requested language matches what is already stored, return it directly.
if record.default_language_code and lang == record.default_language_code and record.default_language_text:
return JSONResponse(
content={
"file_id": record.id,
"source_language": record.detected_language,
"target_language": lang,
"text": record.default_language_text,
"cached": True,
}
)
# If the detected language already matches, return the original text.
detected = record.detected_language
if detected and detected == lang:
return JSONResponse(
content={
"file_id": record.id,
"source_language": detected,
"target_language": lang,
"text": source_text,
"cached": True,
}
)
# Truncate to keep AI costs bounded.
text_to_translate = source_text[:_MAX_TRANSLATION_INPUT]
try:
provider = get_ai_provider()
model = settings.ai_model or settings.openai_model
translated = provider.chat_completion(
messages=[
{
"role": "system",
"content": (
f"You are a professional translator. Translate the following text "
f"into {lang}. Preserve the original formatting, paragraph structure, "
f"and meaning. Do not add any commentary — output ONLY the translated text."
),
},
{"role": "user", "content": text_to_translate},
],
model=model,
temperature=0.3,
)
except Exception as exc:
logger.exception(f"On-the-fly translation failed for file {file_id}: {exc}")
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="Translation failed — the AI provider returned an error",
)
return JSONResponse(
content={
"file_id": record.id,
"source_language": detected or "unknown",
"target_language": lang,
"text": translated,
"cached": False,
}
)
+58 -99
View File
@@ -9,14 +9,12 @@ import urllib.parse
import uuid import uuid
from typing import Optional from typing import Optional
import aiofiles import requests
import httpx from fastapi import APIRouter, HTTPException, Request
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel, HttpUrl, field_validator from pydantic import BaseModel, HttpUrl, field_validator
from app.auth import require_login from app.auth import require_login
from app.config import settings 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.tasks.process_document import process_document
from app.utils.allowed_types import ALLOWED_MIME_TYPES from app.utils.allowed_types import ALLOWED_MIME_TYPES
from app.utils.filename_utils import sanitize_filename from app.utils.filename_utils import sanitize_filename
@@ -28,10 +26,6 @@ logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
class UnsafeRedirectError(httpx.RequestError):
"""Raised when a redirect target fails URL safety checks."""
class URLUploadRequest(BaseModel): class URLUploadRequest(BaseModel):
"""Request model for URL-based file upload""" """Request model for URL-based file upload"""
@@ -110,33 +104,9 @@ def validate_file_type(content_type: str, filename: str) -> bool:
return False return False
async def verify_redirect(response: httpx.Response) -> None:
"""
Event hook to intercept redirects and validate the new destination URL.
Prevents SSRF bypasses via redirects to internal networks or metadata endpoints.
"""
if response.status_code in (301, 302, 303, 307, 308):
location = response.headers.get("Location")
if location:
# Resolve relative redirects
new_url = str(response.url.join(location))
# Validate the new URL
try:
validate_url_safety(new_url)
except HTTPException as e:
raise UnsafeRedirectError(
f"Redirect to unsafe URL blocked: {e.detail}",
request=response.request,
) from e
@router.post("/process-url") @router.post("/process-url")
@require_login @require_login
async def process_url( async def process_url(request: Request, url_request: URLUploadRequest):
request: Request,
url_request: URLUploadRequest,
_rate_ok: None = Depends(require_upload_rate_limit),
):
""" """
Download a file from a URL and enqueue it for processing. Download a file from a URL and enqueue it for processing.
@@ -183,74 +153,67 @@ async def process_url(
logger.info(f"Downloading file from URL: {url}") logger.info(f"Downloading file from URL: {url}")
# Use configured timeout to prevent hanging # Use configured timeout to prevent hanging
async with httpx.AsyncClient( response = requests.get(
url,
timeout=settings.http_request_timeout, timeout=settings.http_request_timeout,
follow_redirects=True, stream=True, # Stream to handle large files
event_hooks={"response": [verify_redirect]}, allow_redirects=True, # Follow redirects
headers={ headers={
"User-Agent": "DocuElevate/1.0", # Identify ourselves "User-Agent": "DocuElevate/1.0", # Identify ourselves
}, },
) as client: )
async with client.stream("GET", url) as response: response.raise_for_status()
response.raise_for_status()
# Validate content type # Validate content type
content_type = response.headers.get("Content-Type", "") content_type = response.headers.get("Content-Type", "")
if not validate_file_type(content_type, safe_filename): if not validate_file_type(content_type, safe_filename):
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail=f"Unsupported file type: {content_type}. " detail=f"Unsupported file type: {content_type}. "
"Supported types: PDF, Office documents, images, plain text", "Supported types: PDF, Office documents, images, plain text",
) )
# Check content length before downloading # Check content length before downloading
content_length = response.headers.get("Content-Length") content_length = response.headers.get("Content-Length")
if content_length: if content_length:
file_size = int(content_length) file_size = int(content_length)
max_size = settings.max_upload_size max_size = settings.max_upload_size
if file_size > max_size: if file_size > max_size:
raise HTTPException(
status_code=413,
detail=f"File too large: {file_size} bytes (max {max_size} bytes)",
)
# Generate unique filename
unique_id = str(uuid.uuid4())
if "." in safe_filename:
file_extension = safe_filename.rsplit(".", 1)[1]
target_filename = f"{unique_id}.{file_extension}"
else:
target_filename = unique_id
target_path = os.path.join(settings.workdir, target_filename)
# Download file in chunks to handle large files
downloaded_size = 0
max_size = settings.max_upload_size
with open(target_path, "wb") as f:
for chunk in response.iter_content(chunk_size=8192):
if chunk:
f.write(chunk)
downloaded_size += len(chunk)
# Check size during download
if downloaded_size > max_size:
# Remove partial file
f.close()
os.remove(target_path)
raise HTTPException( raise HTTPException(
status_code=413, status_code=413,
detail=f"File too large: {file_size} bytes (max {max_size} bytes)", detail=f"File too large: exceeded {max_size} bytes during download",
) )
# Generate unique filename
unique_id = str(uuid.uuid4())
# Check for extension using original_filename to avoid any CodeQL issues
# with safe_filename which is derived from the URL directly.
if "." in original_filename:
_, ext = os.path.splitext(original_filename)
# Strip out the leading dot and any non-alphanumeric chars
clean_ext = "".join(c for c in ext if c.isalnum())
if not clean_ext:
clean_ext = "bin"
target_filename = f"{unique_id}.{clean_ext}"
else:
target_filename = unique_id
target_path = os.path.join(settings.workdir, target_filename)
# Download file in chunks to handle large files
downloaded_size = 0
max_size = settings.max_upload_size
async with aiofiles.open(target_path, "wb") as f:
async for chunk in response.aiter_bytes(chunk_size=8192):
if chunk:
await f.write(chunk)
downloaded_size += len(chunk)
# Check size during download
if downloaded_size > max_size:
# Remove partial file
await f.close()
os.remove(target_path)
raise HTTPException(
status_code=413,
detail=f"File too large: exceeded {max_size} bytes during download",
)
logger.info(f"Downloaded file from URL '{url}' as '{target_filename}' ({downloaded_size} bytes)") logger.info(f"Downloaded file from URL '{url}' as '{target_filename}' ({downloaded_size} bytes)")
# Enqueue for processing # Enqueue for processing
@@ -264,23 +227,19 @@ async def process_url(
"size": downloaded_size, "size": downloaded_size,
} }
except httpx.TimeoutException: except requests.exceptions.Timeout:
logger.error(f"Timeout while downloading file from URL: {url}") logger.error(f"Timeout while downloading file from URL: {url}")
raise HTTPException(status_code=408, detail="Request timeout: server took too long to respond") raise HTTPException(status_code=408, detail="Request timeout: server took too long to respond")
except httpx.ConnectError as e: except requests.exceptions.ConnectionError as e:
logger.error(f"Connection error while downloading file from URL: {url} - {str(e)}") logger.error(f"Connection error while downloading file from URL: {url} - {str(e)}")
raise HTTPException(status_code=502, detail=f"Failed to connect to URL: {str(e)}") raise HTTPException(status_code=502, detail=f"Failed to connect to URL: {str(e)}")
except httpx.HTTPStatusError as e: except requests.exceptions.HTTPError as e:
logger.error(f"HTTP error while downloading file from URL: {url} - {str(e)}") logger.error(f"HTTP error while downloading file from URL: {url} - {str(e)}")
raise HTTPException(status_code=e.response.status_code, detail=f"HTTP error: {str(e)}") raise HTTPException(status_code=e.response.status_code, detail=f"HTTP error: {str(e)}")
except UnsafeRedirectError as e: except requests.exceptions.RequestException as e:
logger.warning(f"Unsafe redirect blocked while downloading file from URL: {url} - {str(e)}")
raise HTTPException(status_code=400, detail=str(e))
except httpx.RequestError as e:
logger.error(f"Error downloading file from URL: {url} - {str(e)}") logger.error(f"Error downloading file from URL: {url} - {str(e)}")
raise HTTPException(status_code=500, detail=f"Failed to download file: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to download file: {str(e)}")
+73 -469
View File
@@ -45,264 +45,78 @@ OAUTH_PROVIDER_NAME = "Single Sign-On"
# Social login providers that are enabled and registered # Social login providers that are enabled and registered
SOCIAL_PROVIDERS: dict[str, dict[str, str]] = {} 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"
# --------------------------------------------------------------------------- # --- Social Login Providers ---------------------------------------------------
# Helpers for dynamic (re-)registration of OAuth 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",
def _register_oauth_client(name: str, **kwargs: object) -> None: client_id=settings.social_auth_google_client_id,
"""Register (or re-register) an authlib OAuth client, clearing any cached instance. client_secret=settings.social_auth_google_client_secret,
server_metadata_url="https://accounts.google.com/.well-known/openid-configuration",
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,
client_kwargs={"scope": "openid profile email"}, client_kwargs={"scope": "openid profile email"},
) )
OAUTH_CONFIGURED = True SOCIAL_PROVIDERS["google"] = {"name": "Google", "icon": "fab fa-google", "color": "red"}
OAUTH_PROVIDER_NAME = settings.oauth_provider_name or "Authentik SSO" 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 AUTH_ENABLED and settings.social_auth_apple_enabled:
if settings.social_auth_google_enabled: if settings.social_auth_apple_client_id and settings.social_auth_apple_team_id:
_google_client_id = settings.social_auth_google_client_id oauth.register(
_google_client_secret = settings.social_auth_google_client_secret name="apple",
if settings.social_auth_google_use_global_credentials and not (_google_client_id and _google_client_secret): client_id=settings.social_auth_apple_client_id,
_google_client_id = settings.google_drive_client_id server_metadata_url="https://appleid.apple.com/.well-known/openid-configuration",
_google_client_secret = settings.google_drive_client_secret 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: if AUTH_ENABLED and settings.social_auth_dropbox_enabled:
_register_oauth_client( if settings.social_auth_dropbox_client_id and settings.social_auth_dropbox_client_secret:
"google", oauth.register(
client_id=_google_client_id, name="dropbox",
client_secret=_google_client_secret, client_id=settings.social_auth_dropbox_client_id,
server_metadata_url="https://accounts.google.com/.well-known/openid-configuration", client_secret=settings.social_auth_dropbox_client_secret,
client_kwargs={"scope": "openid profile email"}, authorize_url="https://www.dropbox.com/oauth2/authorize",
) access_token_url="https://api.dropboxapi.com/oauth2/token",
SOCIAL_PROVIDERS["google"] = {"name": "Google", "icon": "fab fa-google", "color": "red"} userinfo_endpoint="https://api.dropboxapi.com/2/users/get_current_account",
logger.info("Social login provider registered: Google") client_kwargs={"token_endpoint_auth_method": "client_secret_post"},
else: )
logger.warning("SOCIAL_AUTH_GOOGLE_ENABLED=true but client ID/secret not configured") SOCIAL_PROVIDERS["dropbox"] = {"name": "Dropbox", "icon": "fab fa-dropbox", "color": "blue"}
logger.info("Social login provider registered: Dropbox")
# Microsoft else:
if settings.social_auth_microsoft_enabled: logger.warning("SOCIAL_AUTH_DROPBOX_ENABLED=true but client ID/secret not configured")
_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")
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()
router = APIRouter() router = APIRouter()
@@ -311,36 +125,8 @@ def get_current_user(request: Request):
# Check for Bearer token auth first (API tokens) # Check for Bearer token auth first (API tokens)
api_user = getattr(request.state, "api_token_user", None) api_user = getattr(request.state, "api_token_user", None)
if isinstance(api_user, dict): if isinstance(api_user, dict):
logger.debug("[AUTH] get_current_user: resolved from API token (user_id=%s)", api_user.get("id"))
return api_user return api_user
session_user = request.session.get("user") return request.session.get("user")
if session_user:
# Validate server-side session if a session token is present
session_token = request.session.get("_session_token")
if session_token:
try:
from app.database import SessionLocal
from app.utils.session_manager import validate_session
db = SessionLocal()
try:
valid = validate_session(db, session_token)
if not valid:
logger.debug("[AUTH] get_current_user: server-side session invalid — clearing")
request.session.pop("user", None)
request.session.pop("_session_token", None)
return None
finally:
db.close()
except Exception:
logger.debug("[AUTH] get_current_user: session validation error", exc_info=True)
logger.debug(
"[AUTH] get_current_user: resolved from session (user=%s)",
session_user.get("preferred_username") or session_user.get("email") or session_user.get("id"),
)
else:
logger.debug("[AUTH] get_current_user: no user in session or API token")
return session_user
def _resolve_bearer_user(request: Request, db: Session) -> dict | None: def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
@@ -355,12 +141,10 @@ def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
""" """
auth_header = request.headers.get("authorization", "") auth_header = request.headers.get("authorization", "")
if not isinstance(auth_header, str) or not auth_header.startswith("Bearer "): if not isinstance(auth_header, str) or not auth_header.startswith("Bearer "):
logger.debug("[AUTH] _resolve_bearer_user: no Bearer token in Authorization header")
return None return None
raw_token = auth_header[7:] raw_token = auth_header[7:]
if not raw_token or not isinstance(raw_token, str): if not raw_token or not isinstance(raw_token, str):
logger.debug("[AUTH] _resolve_bearer_user: empty or invalid token after 'Bearer ' prefix")
return None return None
from app.api.api_tokens import hash_token from app.api.api_tokens import hash_token
@@ -369,25 +153,8 @@ def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
token_hash = hash_token(raw_token) token_hash = hash_token(raw_token)
db_token = db.query(ApiToken).filter(ApiToken.token_hash == token_hash, ApiToken.is_active.is_(True)).first() db_token = db.query(ApiToken).filter(ApiToken.token_hash == token_hash, ApiToken.is_active.is_(True)).first()
if db_token is None: if db_token is None:
logger.debug("[AUTH] _resolve_bearer_user: no active API token matched the provided hash")
return None 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,
db_token.owner_id,
)
# Update usage tracking # Update usage tracking
try: try:
db_token.last_used_at = datetime.now(timezone.utc) db_token.last_used_at = datetime.now(timezone.utc)
@@ -438,18 +205,16 @@ def require_login(func):
@wraps(func) @wraps(func)
async def wrapper(request: Request, *args, **kwargs): async def wrapper(request: Request, *args, **kwargs):
url_path = urlparse(str(request.url)).path
# Check session auth first # Check session auth first
if request.session.get("user"): if request.session.get("user"):
logger.debug("[AUTH] require_login: session auth OK for %s", url_path)
if inspect.iscoroutinefunction(func): if inspect.iscoroutinefunction(func):
return await func(*args, request=request, **kwargs) return await func(*args, request=request, **kwargs)
else: else:
return func(*args, request=request, **kwargs) return func(*args, request=request, **kwargs)
# Fall back to Bearer token auth for API endpoints # Fall back to Bearer token auth for API endpoints
url_path = urlparse(str(request.url)).path
if url_path.startswith("/api/"): if url_path.startswith("/api/"):
logger.debug("[AUTH] require_login: no session, trying Bearer token for %s", url_path)
try: try:
from app.database import SessionLocal from app.database import SessionLocal
@@ -463,22 +228,17 @@ def require_login(func):
if api_user: if api_user:
request.state.api_token_user = api_user request.state.api_token_user = api_user
logger.debug(
"[AUTH] require_login: Bearer token auth OK for %s (user=%s)", url_path, api_user.get("id")
)
if inspect.iscoroutinefunction(func): if inspect.iscoroutinefunction(func):
return await func(*args, request=request, **kwargs) return await func(*args, request=request, **kwargs)
else: else:
return func(*args, request=request, **kwargs) return func(*args, request=request, **kwargs)
logger.debug("[AUTH] require_login: no valid auth for API endpoint %s — returning 401", url_path)
return JSONResponse( return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED, status_code=status.HTTP_401_UNAUTHORIZED,
content={"error": "Not authenticated"}, content={"error": "Not authenticated"},
) )
# Non-API endpoint with no session — redirect to login # Non-API endpoint with no session — redirect to login
logger.debug("[AUTH] require_login: no session for %s — redirecting to /login", url_path)
request.session["redirect_after_login"] = str(request.url) request.session["redirect_after_login"] = str(request.url)
return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)
@@ -527,21 +287,13 @@ async def login(request: Request):
get_client_ip(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( return templates.TemplateResponse(
request,
"login.html", "login.html",
context={ {
"error": error, "request": request,
"message": message, "error": request.query_params.get("error"),
"show_oauth": show_oauth, "message": request.query_params.get("message"),
"show_oauth": OAUTH_CONFIGURED,
"oauth_provider_name": OAUTH_PROVIDER_NAME, "oauth_provider_name": OAUTH_PROVIDER_NAME,
"social_providers": SOCIAL_PROVIDERS, "social_providers": SOCIAL_PROVIDERS,
"app_version": settings.version, "app_version": settings.version,
@@ -555,15 +307,9 @@ async def login(request: Request):
async def oauth_login(request: Request): async def oauth_login(request: Request):
"""Handle OAuth login flow""" """Handle OAuth login flow"""
if not OAUTH_CONFIGURED: if not OAUTH_CONFIGURED:
logger.debug("[AUTH] oauth_login: OAuth not configured — redirecting to /login")
return RedirectResponse(url="/login?error=OAuth+not+configured", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login?error=OAuth+not+configured", status_code=status.HTTP_302_FOUND)
redirect_uri = request.url_for("oauth_callback") redirect_uri = request.url_for("oauth_callback")
logger.debug(
"[AUTH] oauth_login: initiating Authentik OAuth redirect_uri=%s session_keys=%s",
redirect_uri,
list(request.session.keys()),
)
return await oauth.authentik.authorize_redirect(request, redirect_uri) return await oauth.authentik.authorize_redirect(request, redirect_uri)
@@ -578,23 +324,13 @@ async def social_login(request: Request, provider: str):
A redirect to the provider's authorization page, or back to /login on error. A redirect to the provider's authorization page, or back to /login on error.
""" """
if provider not in SOCIAL_PROVIDERS: if provider not in SOCIAL_PROVIDERS:
logger.debug(
"[AUTH] social_login: unknown provider=%r (registered=%s)", provider, list(SOCIAL_PROVIDERS.keys())
)
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
redirect_uri = request.url_for("social_callback", provider=provider) redirect_uri = request.url_for("social_callback", provider=provider)
oauth_client = getattr(oauth, provider, None) oauth_client = getattr(oauth, provider, None)
if oauth_client is None: if oauth_client is None:
logger.debug("[AUTH] social_login: provider=%r registered but OAuth client not configured", provider)
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
logger.debug(
"[AUTH] social_login: initiating %s OAuth, redirect_uri=%s session_keys=%s",
provider,
redirect_uri,
list(request.session.keys()),
)
return await oauth_client.authorize_redirect(request, redirect_uri) return await oauth_client.authorize_redirect(request, redirect_uri)
@@ -627,17 +363,6 @@ def _normalize_social_userinfo(provider: str, token: dict, raw_userinfo: dict |
"picture": userinfo.get("profile_photo_url", ""), "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) # Standard OIDC providers (Google, Microsoft, Apple)
return { return {
"sub": userinfo.get("sub", ""), "sub": userinfo.get("sub", ""),
@@ -664,39 +389,27 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
A redirect to the user's original destination or the upload page. A redirect to the user's original destination or the upload page.
""" """
if provider not in SOCIAL_PROVIDERS: if provider not in SOCIAL_PROVIDERS:
logger.debug("[AUTH] social_callback: unknown provider=%r", provider)
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
oauth_client = getattr(oauth, provider, None) oauth_client = getattr(oauth, provider, None)
if oauth_client is None: if oauth_client is None:
logger.debug("[AUTH] social_callback: provider=%r not configured", provider)
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
try: try:
logger.debug("[AUTH] social_callback: exchanging auth code for provider=%s", provider)
token = await oauth_client.authorize_access_token(request) token = await oauth_client.authorize_access_token(request)
# Try standard OIDC userinfo first, fall back to token-embedded userinfo # Try standard OIDC userinfo first, fall back to token-embedded userinfo
raw_userinfo = token.get("userinfo") raw_userinfo = token.get("userinfo")
if not raw_userinfo: if not raw_userinfo:
logger.debug("[AUTH] social_callback: no userinfo in token, fetching from userinfo endpoint")
try: try:
resp = await oauth_client.userinfo(token=token) resp = await oauth_client.userinfo(token=token)
raw_userinfo = resp if isinstance(resp, dict) else resp.json() if hasattr(resp, "json") else {} raw_userinfo = resp if isinstance(resp, dict) else resp.json() if hasattr(resp, "json") else {}
except Exception: except Exception:
logger.debug("[AUTH] social_callback: userinfo endpoint failed, using empty dict", exc_info=True)
raw_userinfo = {} raw_userinfo = {}
user_data = _normalize_social_userinfo(provider, token, raw_userinfo) user_data = _normalize_social_userinfo(provider, token, raw_userinfo)
logger.debug(
"[AUTH] social_callback: normalized user_data email=%s sub=%s provider=%s",
user_data.get("email"),
user_data.get("sub"),
provider,
)
if not user_data.get("email"): if not user_data.get("email"):
logger.debug("[AUTH] social_callback: no email in user_data — aborting")
return RedirectResponse( return RedirectResponse(
url="/login?error=Could+not+retrieve+email+from+provider", url="/login?error=Could+not+retrieve+email+from+provider",
status_code=status.HTTP_302_FOUND, status_code=status.HTTP_302_FOUND,
@@ -715,27 +428,6 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
request.session["user"] = user_data request.session["user"] = user_data
# Create server-side session for tracking and revocation
try:
from app.utils.session_manager import create_session
_session_user_id = (
user_data.get("sub")
or user_data.get("preferred_username")
or user_data.get("email")
or user_data.get("id")
)
if _session_user_id:
user_session = create_session(
db,
user_id=_session_user_id,
ip_address=get_client_ip(request),
user_agent=request.headers.get("user-agent"),
)
request.session["_session_token"] = user_session.session_token
except Exception:
logger.debug("[AUTH] Failed to create server-side session for social user", exc_info=True)
# Auto-create or update UserProfile # Auto-create or update UserProfile
_ensure_user_profile(db, user_data, is_admin=False) _ensure_user_profile(db, user_data, is_admin=False)
@@ -750,29 +442,21 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
) )
# Mobile app flow: issue an inline API token and redirect back to the app. # Mobile app flow: issue an inline API token and redirect back to the app.
logger.debug(
"[MOBILE] social_callback: checking for mobile redirect (session has mobile_redirect_uri=%s)",
"mobile_redirect_uri" in request.session,
)
mobile_resp = _create_mobile_redirect(request, db) mobile_resp = _create_mobile_redirect(request, db)
if mobile_resp: if mobile_resp:
logger.info("[MOBILE] social_callback: returning mobile redirect response for provider=%s", provider)
return mobile_resp return mobile_resp
if user_id: if user_id:
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first() profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
if profile and not profile.onboarding_completed: if profile and not profile.onboarding_completed:
logger.debug("[AUTH] social_callback: user=%s needs onboarding, redirecting", user_id)
post_onboarding = request.session.pop("redirect_after_login", "/upload") post_onboarding = request.session.pop("redirect_after_login", "/upload")
request.session["post_onboarding_redirect"] = post_onboarding request.session["post_onboarding_redirect"] = post_onboarding
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
redirect_url = request.session.pop("redirect_after_login", "/upload") redirect_url = request.session.pop("redirect_after_login", "/upload")
logger.debug("[AUTH] social_callback: login complete, redirecting to %s", redirect_url)
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND) return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
except Exception as e: except Exception as e:
logger.warning("[SECURITY] SOCIAL_LOGIN_FAILURE provider=%s error=%s", provider, type(e).__name__) logger.warning("[SECURITY] SOCIAL_LOGIN_FAILURE provider=%s error=%s", provider, type(e).__name__)
logger.debug("[AUTH] social_callback: full exception for provider=%s", provider, exc_info=True)
return RedirectResponse( return RedirectResponse(
url="/login?error=Social+login+failed.+Please+try+again.", status_code=status.HTTP_302_FOUND url="/login?error=Social+login+failed.+Please+try+again.", status_code=status.HTTP_302_FOUND
) )
@@ -877,23 +561,15 @@ def _ensure_user_profile(db: Session, user_data: dict, is_admin: bool = False) -
async def oauth_callback(request: Request, db: Session = Depends(get_db)): async def oauth_callback(request: Request, db: Session = Depends(get_db)):
"""Handle OAuth callback from provider""" """Handle OAuth callback from provider"""
try: try:
logger.debug("[AUTH] oauth_callback: exchanging authorization code for token")
token = await oauth.authentik.authorize_access_token(request) token = await oauth.authentik.authorize_access_token(request)
userinfo = token.get("userinfo") userinfo = token.get("userinfo")
if not userinfo: if not userinfo:
logger.debug("[AUTH] oauth_callback: no userinfo in token response — aborting")
return RedirectResponse( return RedirectResponse(
url="/login?error=Failed+to+retrieve+user+information", status_code=status.HTTP_302_FOUND url="/login?error=Failed+to+retrieve+user+information", status_code=status.HTTP_302_FOUND
) )
# Store user info in session # Store user info in session
user_data = dict(userinfo) user_data = dict(userinfo)
logger.debug(
"[AUTH] oauth_callback: received userinfo email=%s sub=%s groups=%s",
user_data.get("email"),
user_data.get("sub"),
user_data.get("groups", []),
)
# Add Gravatar picture if no picture is provided # Add Gravatar picture if no picture is provided
if not user_data.get("picture") and user_data.get("email"): if not user_data.get("picture") and user_data.get("email"):
@@ -908,39 +584,12 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)):
groups = user_data.get("groups", []) groups = user_data.get("groups", [])
admin_group = (settings.admin_group_name or "admin").strip().lower() admin_group = (settings.admin_group_name or "admin").strip().lower()
is_admin = admin_group in [group.lower() for group in groups] is_admin = admin_group in [group.lower() for group in groups]
logger.debug(
"[AUTH] oauth_callback: admin group check — looking for %r in %s → is_admin=%s",
admin_group,
[g.lower() for g in groups],
is_admin,
)
# Set is_admin flag (defaults to False for OAuth users unless they're in admin group) # Set is_admin flag (defaults to False for OAuth users unless they're in admin group)
user_data["is_admin"] = is_admin user_data["is_admin"] = is_admin
request.session["user"] = user_data request.session["user"] = user_data
# Create server-side session for tracking and revocation
try:
from app.utils.session_manager import create_session
_session_user_id = (
user_data.get("sub")
or user_data.get("preferred_username")
or user_data.get("email")
or user_data.get("id")
)
if _session_user_id:
user_session = create_session(
db,
user_id=_session_user_id,
ip_address=get_client_ip(request),
user_agent=request.headers.get("user-agent"),
)
request.session["_session_token"] = user_session.session_token
except Exception:
logger.debug("[AUTH] Failed to create server-side session for OAuth user", exc_info=True)
# Auto-create or update UserProfile so the user appears in admin user management # Auto-create or update UserProfile so the user appears in admin user management
_ensure_user_profile(db, user_data, is_admin=is_admin) _ensure_user_profile(db, user_data, is_admin=is_admin)
@@ -974,18 +623,15 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)):
if user_id: if user_id:
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first() profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
if profile and not profile.onboarding_completed: if profile and not profile.onboarding_completed:
logger.debug("[AUTH] oauth_callback: user=%s needs onboarding, redirecting", user_id)
post_onboarding = request.session.pop("redirect_after_login", "/upload") post_onboarding = request.session.pop("redirect_after_login", "/upload")
request.session["post_onboarding_redirect"] = post_onboarding request.session["post_onboarding_redirect"] = post_onboarding
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND) return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
# Redirect to original destination or default # Redirect to original destination or default
redirect_url = request.session.pop("redirect_after_login", "/upload") redirect_url = request.session.pop("redirect_after_login", "/upload")
logger.debug("[AUTH] oauth_callback: login complete, redirecting to %s", redirect_url)
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND) return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
except Exception as e: except Exception as e:
logger.warning(f"[SECURITY] OAUTH_LOGIN_FAILURE error={type(e).__name__}") logger.warning(f"[SECURITY] OAUTH_LOGIN_FAILURE error={type(e).__name__}")
logger.debug("[AUTH] oauth_callback: full exception details", exc_info=True)
return RedirectResponse(url=f"/login?error=Authentication+failed:+{str(e)}", status_code=status.HTTP_302_FOUND) return RedirectResponse(url=f"/login?error=Authentication+failed:+{str(e)}", status_code=status.HTTP_302_FOUND)
@@ -1194,19 +840,6 @@ async def auth(request: Request, db: Session = Depends(get_db)):
return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302) return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302)
user_data = _build_session_user(local_user) user_data = _build_session_user(local_user)
request.session["user"] = user_data request.session["user"] = user_data
# Create server-side session for tracking and revocation
try:
from app.utils.session_manager import create_session
user_session = create_session(
db,
user_id=local_user.email,
ip_address=get_client_ip(request),
user_agent=request.headers.get("user-agent"),
)
request.session["_session_token"] = user_session.session_token
except Exception:
logger.debug("[AUTH] Failed to create server-side session", exc_info=True)
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", local_user.email) logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", local_user.email)
_record_login_event(db, request, local_user.email, success=True) _record_login_event(db, request, local_user.email, success=True)
_ensure_user_profile(db, user_data, is_admin=bool(local_user.is_admin)) _ensure_user_profile(db, user_data, is_admin=bool(local_user.is_admin))
@@ -1263,20 +896,6 @@ async def auth(request: Request, db: Session = Depends(get_db)):
"is_admin": True, "is_admin": True,
} }
request.session["user"] = admin_user_data request.session["user"] = admin_user_data
# Create server-side session for tracking and revocation
try:
from app.utils.session_manager import create_session
admin_user_id = settings.admin_username or "admin"
user_session = create_session(
db,
user_id=admin_user_id,
ip_address=get_client_ip(request),
user_agent=request.headers.get("user-agent"),
)
request.session["_session_token"] = user_session.session_token
except Exception:
logger.debug("[AUTH] Failed to create server-side session for admin", exc_info=True)
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", username) logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", username)
_record_login_event(db, request, username, success=True) _record_login_event(db, request, username, success=True)
_ensure_user_profile(db, admin_user_data, is_admin=True) _ensure_user_profile(db, admin_user_data, is_admin=True)
@@ -1312,7 +931,6 @@ async def logout(request: Request, db: Session = Depends(get_db)):
username = "unknown" username = "unknown"
if isinstance(user, dict): if isinstance(user, dict):
username = user.get("preferred_username") or user.get("email") or "unknown" username = user.get("preferred_username") or user.get("email") or "unknown"
logger.debug("[AUTH] logout: clearing session for user=%s client_ip=%s", username, get_client_ip(request))
logger.info(f"[SECURITY] LOGOUT user={username}") logger.info(f"[SECURITY] LOGOUT user={username}")
try: try:
from app.utils.audit_service import record_event from app.utils.audit_service import record_event
@@ -1327,20 +945,6 @@ async def logout(request: Request, db: Session = Depends(get_db)):
) )
except Exception: except Exception:
logger.debug("Failed to write logout audit event for user=%s", username, exc_info=True) logger.debug("Failed to write logout audit event for user=%s", username, exc_info=True)
# Revoke server-side session
session_token = request.session.get("_session_token")
if session_token:
try:
from app.utils.session_manager import validate_session
user_session = validate_session(db, session_token)
if user_session:
user_session.is_revoked = True
user_session.revoked_at = datetime.now(timezone.utc)
db.commit()
except Exception:
logger.debug("[AUTH] Failed to revoke server-side session", exc_info=True)
request.session.pop("_session_token", None)
request.session.pop("user", None) request.session.pop("user", None)
return RedirectResponse(url="/login?message=You+have+been+logged+out+successfully", status_code=302) return RedirectResponse(url="/login?message=You+have+been+logged+out+successfully", status_code=302)
+2 -69
View File
@@ -1,15 +1,10 @@
# app/celery_app.py # app/celery_app.py
import logging
import os
from celery import Celery from celery import Celery
from celery.signals import task_failure, worker_ready from celery.signals import task_failure, worker_ready
from app.config import settings from app.config import settings
logger = logging.getLogger(__name__)
celery = Celery( celery = Celery(
"document_processor", "document_processor",
broker=settings.redis_url, broker=settings.redis_url,
@@ -26,64 +21,6 @@ celery.conf.task_routes = {
"app.tasks.*": {"queue": "document_processor"}, "app.tasks.*": {"queue": "document_processor"},
} }
# Mapping of document pipeline task names to the positional index of ``file_id``
# in their ``args`` tuple. These indices correspond to the task signatures:
# process_with_ocr(filename, file_id, ...) → index 1
# extract_metadata_with_gpt(filename, text, file_id) → index 2
# embed_metadata_into_pdf(path, text, metadata, file_id) → index 3
# Tasks that always pass ``file_id`` as a keyword argument
# (e.g. ``process_document``, ``finalize_document_storage``) are not listed
# here — their ``file_id`` is found via ``kwargs`` instead.
_FILE_ID_ARG_INDEX: dict[str, int] = {
"app.tasks.process_with_ocr.process_with_ocr": 1,
"app.tasks.extract_metadata_with_gpt.extract_metadata_with_gpt": 2,
"app.tasks.embed_metadata_into_pdf.embed_metadata_into_pdf": 3,
}
def _dispatch_user_failure_notification(sender, exception, args: list | None, kwargs: dict | None) -> None:
"""Best-effort per-user failure notification for document pipeline tasks.
Extracts ``file_id`` from the failed task's arguments, looks up the owning
user from the database, and dispatches a ``document.failed`` notification.
"""
from app.database import SessionLocal
from app.models import FileRecord
from app.utils.user_notification import notify_user_document_failed
task_name = sender.name if sender else ""
if not task_name.startswith("app.tasks."):
return
# 1. Resolve file_id from kwargs or positional args
file_id = (kwargs or {}).get("file_id")
if file_id is None:
idx = _FILE_ID_ARG_INDEX.get(task_name)
if idx is not None and args and len(args) > idx:
val = args[idx]
if isinstance(val, int):
file_id = val
if file_id is None:
return
# 2. Look up owner from the database
with SessionLocal() as db:
record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if not record or not record.owner_id:
return
owner_id = record.owner_id
filename = record.original_filename or record.local_filename or "unknown"
# 3. Dispatch per-user notification
error_msg = f"{type(exception).__name__}: {exception}" if exception else "Unknown error"
notify_user_document_failed(
owner_id=owner_id,
filename=os.path.basename(filename),
error=error_msg,
file_id=file_id,
)
@worker_ready.connect @worker_ready.connect
def init_sentry_on_worker_ready(**kwargs): def init_sentry_on_worker_ready(**kwargs):
@@ -111,10 +48,6 @@ def task_failure_handler(
kwargs=kwargs or {}, kwargs=kwargs or {},
) )
except Exception as e: except Exception as e:
logger.exception(f"Failed to send task failure notification: {e}") import logging
# Also dispatch a per-user failure notification for document pipeline tasks logging.exception(f"Failed to send task failure notification: {e}")
try:
_dispatch_user_failure_notification(sender, exception, args, kwargs)
except Exception:
logger.warning("Could not dispatch per-user failure notification", exc_info=True)
-5
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 # Import the shared Celery instance
from app.celery_app import celery from app.celery_app import celery
from app.config import settings 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.backup_tasks import cleanup_old_backups, create_backup # noqa: F401
from app.tasks.batch_tasks import ( # noqa: F401 from app.tasks.batch_tasks import ( # noqa: F401
backfill_missing_metadata, backfill_missing_metadata,
@@ -23,7 +22,6 @@ from app.tasks.batch_tasks import ( # noqa: F401
sync_search_index, sync_search_index,
) )
from app.tasks.check_credentials import check_credentials 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.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_pdf import convert_to_pdf # noqa: F401
from app.tasks.convert_to_pdfa import convert_to_pdfa # noqa: F401 from app.tasks.convert_to_pdfa import convert_to_pdfa # noqa: F401
@@ -41,12 +39,10 @@ from app.tasks.refine_text_with_gpt import refine_text_with_gpt # noqa: F401
from app.tasks.rotate_pdf_pages import rotate_pdf_pages # noqa: F401 from app.tasks.rotate_pdf_pages import rotate_pdf_pages # noqa: F401
from app.tasks.send_to_all import send_to_all_destinations # noqa: F401 from app.tasks.send_to_all import send_to_all_destinations # noqa: F401
from app.tasks.subscription_tasks import apply_pending_subscription_changes_all # noqa: F401 from app.tasks.subscription_tasks import apply_pending_subscription_changes_all # noqa: F401
from app.tasks.translate_to_default_language import translate_to_default_language # noqa: F401
# Import new send tasks # Import new send tasks
from app.tasks.upload_to_dropbox import upload_to_dropbox # noqa: F401 from app.tasks.upload_to_dropbox import upload_to_dropbox # noqa: F401
from app.tasks.upload_to_email import upload_to_email # noqa: F401 from app.tasks.upload_to_email import upload_to_email # noqa: F401
from app.tasks.upload_to_evernote import upload_to_evernote # noqa: F401
from app.tasks.upload_to_ftp import upload_to_ftp # noqa: F401 from app.tasks.upload_to_ftp import upload_to_ftp # noqa: F401
from app.tasks.upload_to_google_drive import upload_to_google_drive # noqa: F401 from app.tasks.upload_to_google_drive import upload_to_google_drive # noqa: F401
from app.tasks.upload_to_icloud import upload_to_icloud # noqa: F401 from app.tasks.upload_to_icloud import upload_to_icloud # noqa: F401
@@ -55,7 +51,6 @@ from app.tasks.upload_to_onedrive import upload_to_onedrive # noqa: F401
from app.tasks.upload_to_paperless import upload_to_paperless # noqa: F401 from app.tasks.upload_to_paperless import upload_to_paperless # noqa: F401
from app.tasks.upload_to_s3 import upload_to_s3 # noqa: F401 from app.tasks.upload_to_s3 import upload_to_s3 # noqa: F401
from app.tasks.upload_to_sftp import upload_to_sftp # noqa: F401 from app.tasks.upload_to_sftp import upload_to_sftp # noqa: F401
from app.tasks.upload_to_sharepoint import upload_to_sharepoint # noqa: F401
from app.tasks.upload_to_user_integration import upload_to_user_integration # noqa: F401 from app.tasks.upload_to_user_integration import upload_to_user_integration # noqa: F401
from app.tasks.upload_to_webdav import upload_to_webdav # noqa: F401 from app.tasks.upload_to_webdav import upload_to_webdav # noqa: F401
from app.tasks.upload_with_rclone import send_to_all_rclone_destinations, upload_with_rclone # noqa: F401 from app.tasks.upload_with_rclone import send_to_all_rclone_destinations, upload_with_rclone # noqa: F401
-293
View File
@@ -13,24 +13,6 @@ class Settings(BaseSettings):
database_url: str database_url: str
redis_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_api_key: str
openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint
openai_model: str = "gpt-4o-mini" # Default model openai_model: str = "gpt-4o-mini" # Default model
@@ -66,51 +48,6 @@ class Settings(BaseSettings):
workdir: str workdir: str
debug: bool = False # Default to False debug: bool = False # Default to False
# Logging level for the application. Accepts standard Python level names:
# DEBUG, INFO, WARNING, ERROR, CRITICAL. When *debug* is True and
# *log_level* has not been explicitly set, the effective level is forced to
# DEBUG so that all ``logger.debug()`` calls produce output.
log_level: str = Field(
default="INFO",
description=(
"Python logging level for the application root logger. "
"Accepts: DEBUG, INFO, WARNING, ERROR, CRITICAL. "
"When DEBUG=True and LOG_LEVEL is not explicitly set, "
"the effective level is automatically lowered to DEBUG."
),
)
# Log output format. ``text`` is the human-readable default.
# ``json`` emits one JSON object per line, ideal for log collectors
# (Promtail, Fluentd, Filebeat, Datadog agent) and SIEM ingestion.
log_format: str = Field(
default="text",
description=(
"Log output format: 'text' (human-readable, default) or "
"'json' (structured JSON lines for SIEM / log aggregation)."
),
)
# Optional syslog forwarding for application logs (not just audit events).
# When enabled, a Python SysLogHandler is added to the root logger so that
# every log message is also sent to the configured syslog receiver.
log_syslog_enabled: bool = Field(
default=False,
description="Forward application logs to a syslog receiver in addition to stdout.",
)
log_syslog_host: str = Field(
default="localhost",
description="Hostname or IP of the syslog receiver for application logs.",
)
log_syslog_port: int = Field(
default=514,
description="Port of the syslog receiver for application logs.",
)
log_syslog_protocol: str = Field(
default="udp",
description="Protocol for syslog transport: 'udp' or 'tcp'.",
)
# Making Dropbox optional # Making Dropbox optional
dropbox_enabled: bool = Field( dropbox_enabled: bool = Field(
default=True, default=True,
@@ -120,16 +57,6 @@ class Settings(BaseSettings):
dropbox_app_secret: Optional[str] = None dropbox_app_secret: Optional[str] = None
dropbox_folder: Optional[str] = None dropbox_folder: Optional[str] = None
dropbox_refresh_token: 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 # Making Nextcloud optional
nextcloud_enabled: bool = Field( nextcloud_enabled: bool = Field(
@@ -154,17 +81,6 @@ class Settings(BaseSettings):
# "language": "Language", "correspondent": "Correspondent"} # "language": "Language", "correspondent": "Correspondent"}
paperless_custom_fields_mapping: Optional[str] = None paperless_custom_fields_mapping: Optional[str] = None
# Evernote destination settings
evernote_enabled: bool = Field(
default=True,
description="Enable Evernote as an upload destination. Set to False to disable uploads even when credentials are configured.",
)
evernote_auth_token: Optional[str] = None
evernote_sandbox: bool = False
evernote_notebook_guid: Optional[str] = None
evernote_default_tags: Optional[str] = None
evernote_include_metadata: bool = True
azure_ai_key: str azure_ai_key: str
azure_region: str azure_region: str
azure_endpoint: str azure_endpoint: str
@@ -204,65 +120,12 @@ class Settings(BaseSettings):
google_docai_processor_id: Optional[str] = None google_docai_processor_id: Optional[str] = None
google_docai_location: str = "us" # Processor location, e.g. "us" or "eu" google_docai_location: str = "us" # Processor location, e.g. "us" or "eu"
external_hostname: str = "localhost" # Default to localhost external_hostname: str = "localhost" # Default to localhost
public_base_url: Optional[str] = Field(
default=None,
description=(
"The full public base URL of the application, including scheme "
"(e.g., 'https://docuelevate.example.com'). "
"When set, this overrides the auto-detected URL for OAuth redirect URIs. "
"This is required when the application is behind a reverse proxy that does "
"not forward X-Forwarded-Proto headers correctly."
),
)
# ---------------------------------------------------------------------------
# Document Translation Settings
# ---------------------------------------------------------------------------
# Default target language for automatic document translation (ISO 639-1 code).
# After OCR / metadata extraction, if the detected document language differs
# from this value the system translates the extracted text into this language
# and stores it alongside the original. Other language translations are
# generated on the fly via the AI provider and are NOT persisted.
# Per-user overrides are stored in UserProfile.default_document_language.
default_document_language: str = Field(
default="en",
description=(
"ISO 639-1 language code for the default translation target "
"(e.g. 'en', 'de', 'fr'). Documents whose detected language "
"differs are automatically translated into this language after "
"processing. Default: 'en' (English)."
),
)
# Authentication settings # Authentication settings
auth_enabled: bool = True # Default to enabled auth_enabled: bool = True # Default to enabled
admin_username: Optional[str] = None admin_username: Optional[str] = None
admin_password: Optional[str] = None admin_password: Optional[str] = None
session_secret: Optional[str] = None session_secret: Optional[str] = None
session_lifetime_days: int = Field(
default=30,
description=(
"Session lifetime in days. Common values: 30, 60, 90. "
"Determines how long a user stays logged in before being required to re-authenticate. "
"Applies to both browser sessions and the session cookie max_age."
),
)
session_lifetime_custom_days: int | None = Field(
default=None,
description=(
"Override session_lifetime_days with a custom value. "
"When set, this takes precedence over session_lifetime_days. "
"Useful for admin-configured non-standard durations."
),
)
qr_login_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).",
)
admin_group_name: str = "admin" admin_group_name: str = "admin"
# Multi-user settings # Multi-user settings
@@ -320,55 +183,12 @@ class Settings(BaseSettings):
authentik_client_secret: Optional[str] = None authentik_client_secret: Optional[str] = None
authentik_config_url: Optional[str] = None authentik_config_url: Optional[str] = None
oauth_provider_name: Optional[str] = None # Name to display for the OAuth provider 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 # Social Login Providers
# Google OAuth2 # Google OAuth2
social_auth_google_enabled: bool = False social_auth_google_enabled: bool = False
social_auth_google_client_id: Optional[str] = None social_auth_google_client_id: Optional[str] = None
social_auth_google_client_secret: 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) # Microsoft OAuth2 (Azure AD / Microsoft Entra ID)
social_auth_microsoft_enabled: bool = False social_auth_microsoft_enabled: bool = False
@@ -383,17 +203,6 @@ class Settings(BaseSettings):
"Default: common." "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 # Apple Sign-In
social_auth_apple_enabled: bool = False social_auth_apple_enabled: bool = False
@@ -406,21 +215,6 @@ class Settings(BaseSettings):
social_auth_dropbox_enabled: bool = False social_auth_dropbox_enabled: bool = False
social_auth_dropbox_client_id: Optional[str] = None social_auth_dropbox_client_id: Optional[str] = None
social_auth_dropbox_client_secret: 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 # Local user signup
allow_local_signup: bool = Field( allow_local_signup: bool = Field(
@@ -731,15 +525,6 @@ class Settings(BaseSettings):
onedrive_refresh_token: Optional[str] = None # Required for personal accounts onedrive_refresh_token: Optional[str] = None # Required for personal accounts
onedrive_folder_path: Optional[str] = None onedrive_folder_path: Optional[str] = None
# SharePoint settings
sharepoint_client_id: Optional[str] = None
sharepoint_client_secret: Optional[str] = None
sharepoint_tenant_id: Optional[str] = "common"
sharepoint_refresh_token: Optional[str] = None
sharepoint_site_url: Optional[str] = None # e.g. https://tenant.sharepoint.com/sites/sitename
sharepoint_document_library: Optional[str] = "Documents" # Document library name
sharepoint_folder_path: Optional[str] = None # Subfolder inside the library
# AWS S3 settings # AWS S3 settings
s3_enabled: bool = Field( s3_enabled: bool = Field(
default=True, default=True,
@@ -790,25 +575,6 @@ class Settings(BaseSettings):
), ),
) )
# System reset / factory reset settings
factory_reset_on_startup: bool = Field(
default=False,
description=(
"When enabled, DocuElevate wipes all user data (database rows and "
"work-files on disk) on every startup so the instance always comes "
"up in a clean, fresh state. Useful for demo or testing environments. "
"Default: False."
),
)
enable_factory_reset: bool = Field(
default=False,
description=(
"Show the 'System Reset' page in the admin UI. When enabled, "
"administrators can trigger a full data wipe or a wipe-and-reimport "
"directly from the web interface. Default: False."
),
)
# PDF/A archival conversion settings # PDF/A archival conversion settings
enable_pdfa_conversion: bool = Field( enable_pdfa_conversion: bool = Field(
default=False, default=False,
@@ -918,11 +684,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 settings
notification_urls: Union[List[str], str] = Field( notification_urls: Union[List[str], str] = Field(
default_factory=list, default_factory=list,
@@ -957,12 +718,6 @@ class Settings(BaseSettings):
description="Enable webhook delivery for document events", 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 / restore settings ──────────────────────────────────────────────
backup_enabled: bool = Field( backup_enabled: bool = Field(
default=True, default=True,
@@ -1248,20 +1003,6 @@ class Settings(BaseSettings):
), ),
) )
# Per-user upload rate limiting (health-aware, Redis-backed sliding window)
upload_rate_limit_per_user: int = Field(
default=20,
description=(
"Maximum number of file uploads allowed per user within the sliding window. "
"The effective limit may be reduced dynamically when the system is under heavy load "
"(high queue depth or CPU usage). Set to 0 to disable per-user upload rate limiting."
),
)
upload_rate_limit_window: int = Field(
default=60,
description="Sliding window size in seconds for per-user upload rate limiting (default: 60).",
)
# Rate Limiting Configuration (see SECURITY_AUDIT.md and docs/API.md) # Rate Limiting Configuration (see SECURITY_AUDIT.md and docs/API.md)
# Protects against DoS attacks and API abuse # Protects against DoS attacks and API abuse
rate_limiting_enabled: bool = Field( rate_limiting_enabled: bool = Field(
@@ -1388,40 +1129,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") @model_validator(mode="before")
@classmethod @classmethod
def strip_outer_quotes(cls, data: Any) -> Any: def strip_outer_quotes(cls, data: Any) -> Any:
+2 -32
View File
@@ -10,7 +10,6 @@ from typing import Any
from sqlalchemy import create_engine, exc from sqlalchemy import create_engine, exc
from sqlalchemy.engine.url import make_url from sqlalchemy.engine.url import make_url
from sqlalchemy.orm import Session, declarative_base, sessionmaker from sqlalchemy.orm import Session, declarative_base, sessionmaker
from sqlalchemy.pool import NullPool, QueuePool
from app.config import settings from app.config import settings
@@ -18,37 +17,9 @@ logger = logging.getLogger(__name__)
Base = declarative_base() Base = declarative_base()
# --------------------------------------------------------------------------- # Parse the DATABASE_URL
# Engine construction
# ---------------------------------------------------------------------------
DB_URL = settings.database_url DB_URL = settings.database_url
_parsed_url = make_url(DB_URL) engine = create_engine(DB_URL, connect_args={"check_same_thread": False})
_connect_args: dict[str, Any] = {}
_engine_kwargs: dict[str, Any] = {
"pool_pre_ping": True, # detect stale / dropped connections before use
}
if _parsed_url.get_backend_name() == "sqlite":
# SQLite does not benefit from connection pooling and is prone to
# QueuePool exhaustion under concurrent access. NullPool opens a fresh
# connection for each request and closes it immediately afterwards,
# completely avoiding the "QueuePool limit reached" TimeoutError.
_connect_args["check_same_thread"] = False
_engine_kwargs["poolclass"] = NullPool
else:
# PostgreSQL / MySQL — use a bounded QueuePool with configurable limits.
_engine_kwargs["poolclass"] = QueuePool
_engine_kwargs.update(
{
"pool_size": settings.db_pool_size,
"max_overflow": settings.db_max_overflow,
"pool_timeout": settings.db_pool_timeout,
"pool_recycle": settings.db_pool_recycle,
}
)
engine = create_engine(DB_URL, connect_args=_connect_args, **_engine_kwargs)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
@@ -300,7 +271,6 @@ def _ensure_indexes(engine: Any, inspector: Any) -> None:
if table not in columns_by_table: if table not in columns_by_table:
columns_by_table[table] = {col["name"] for col in inspector.get_columns(table)} columns_by_table[table] = {col["name"] for col in inspector.get_columns(table)}
if column in columns_by_table[table]: if column in columns_by_table[table]:
# SECURITY: Quoted identifiers to prevent SQL injection during index creation
quoted_idx = preparer.quote(idx_name) quoted_idx = preparer.quote(idx_name)
quoted_table = preparer.quote(table) quoted_table = preparer.quote(table)
quoted_col = preparer.quote(column) quoted_col = preparer.quote(column)
+8 -153
View File
@@ -1,11 +1,8 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
import json as _json_mod
import logging import logging
import os import os
import pathlib import pathlib
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from datetime import datetime as _dt
from datetime import timezone as _tz
from fastapi import FastAPI, HTTPException, Request, status from fastapi import FastAPI, HTTPException, Request, status
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
@@ -39,114 +36,6 @@ from app.views import router as frontend_router
# Explicitly include the files router # Explicitly include the files router
from app.views.files import router as files_router from app.views.files import router as files_router
# ---------------------------------------------------------------------------
# Configure Python root logging level early so that *all* loggers (including
# those already created via ``logging.getLogger(__name__)`` in other modules)
# respect the configured level.
#
# Standard behaviour (matches Django, Flask, 12-factor conventions):
# • ``LOG_LEVEL`` env var takes precedence when explicitly set.
# • When ``DEBUG=True`` and ``LOG_LEVEL`` is **not** set, the effective
# level is automatically lowered to ``DEBUG``.
# • Default (neither flag set): ``INFO``.
#
# ``LOG_FORMAT=json`` enables structured JSON lines on stdout, suitable for
# Promtail, Fluentd, Filebeat, Datadog, Splunk UF, or any log collector.
#
# ``LOG_SYSLOG_ENABLED=true`` adds a Python SysLogHandler so that every log
# message is also forwarded to the configured syslog receiver — useful for
# traditional (non-container) deployments and centralised SIEM ingestion.
#
# Noisy third-party loggers (httpx, httpcore, authlib, etc.) are pinned to
# WARNING when the app-level is DEBUG to keep output useful.
# ---------------------------------------------------------------------------
_explicit_log_level = os.environ.get("LOG_LEVEL")
if settings.debug and _explicit_log_level is None:
_effective_level = "DEBUG"
else:
_effective_level = settings.log_level.upper()
_effective_level_int = getattr(logging, _effective_level, logging.INFO)
class _JsonFormatter(logging.Formatter):
"""Emit one JSON object per log line for machine consumption.
Fields emitted: ``timestamp``, ``level``, ``logger``, ``message``,
``module``, ``funcName``, ``lineno``, and when present ``exc_info``.
Compatible with Grafana Loki, Splunk, ELK, Datadog, and most SIEM tools.
"""
def format(self, record: logging.LogRecord) -> str:
log_entry: dict = {
"timestamp": _dt.fromtimestamp(record.created, tz=_tz.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"module": record.module,
"funcName": record.funcName,
"lineno": record.lineno,
}
if record.exc_info and record.exc_info[1] is not None:
log_entry["exc_info"] = self.formatException(record.exc_info)
return _json_mod.dumps(log_entry, default=str)
# Choose formatter based on LOG_FORMAT setting
if settings.log_format.lower() == "json":
_handler = logging.StreamHandler()
_handler.setFormatter(_JsonFormatter())
logging.root.handlers = [_handler]
logging.root.setLevel(_effective_level_int)
else:
logging.basicConfig(
level=_effective_level_int,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
force=True,
)
# Optional: forward application logs to a syslog receiver
if settings.log_syslog_enabled:
import logging.handlers as _lh
import socket as _socket
_proto = settings.log_syslog_protocol.lower()
_socktype = _socket.SOCK_STREAM if _proto == "tcp" else _socket.SOCK_DGRAM
_syslog_handler = _lh.SysLogHandler(
address=(settings.log_syslog_host, settings.log_syslog_port),
socktype=_socktype,
)
_syslog_handler.setLevel(_effective_level_int)
# Use the same formatter as stdout (text or JSON)
if settings.log_format.lower() == "json":
_syslog_handler.setFormatter(_JsonFormatter())
else:
_syslog_handler.setFormatter(logging.Formatter("%(name)s - %(levelname)s - %(message)s"))
logging.root.addHandler(_syslog_handler)
# Keep noisy third-party loggers quiet at DEBUG level
if _effective_level_int <= logging.DEBUG:
for _noisy in (
"httpx",
"httpcore",
"authlib",
"urllib3",
"hpack",
"multipart",
"watchfiles",
):
logging.getLogger(_noisy).setLevel(logging.WARNING)
_startup_logger = logging.getLogger(__name__)
_startup_logger.info(
"Root logging level set to %s (debug=%s, format=%s, syslog=%s)",
_effective_level,
settings.debug,
settings.log_format,
settings.log_syslog_enabled,
)
# Load configuration from .env for the session key # Load configuration from .env for the session key
config = Config(".env") config = Config(".env")
# Use settings.session_secret which has proper validation # Use settings.session_secret which has proper validation
@@ -170,12 +59,6 @@ async def lifespan(app: FastAPI):
# Startup: Initialize database # Startup: Initialize database
init_db() # Create tables if they don't exist init_db() # Create tables if they don't exist
# Factory reset on startup — wipe all user data before anything else
if settings.factory_reset_on_startup:
from app.utils.system_reset import perform_startup_reset
perform_startup_reset()
# Load settings from database after DB initialization # Load settings from database after DB initialization
from app.database import SessionLocal from app.database import SessionLocal
from app.utils.config_loader import load_settings_from_db from app.utils.config_loader import load_settings_from_db
@@ -189,18 +72,6 @@ async def lifespan(app: FastAPI):
finally: finally:
db.close() 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 # 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. # via the database UI (e.g. SENTRY_DSN) are respected in addition to env vars.
init_sentry() init_sentry()
@@ -294,16 +165,10 @@ async def lifespan(app: FastAPI):
yield yield
# Shutdown: Cleanup tasks # Shutdown: Cleanup tasks
try: logging.info("Application shutting down")
logging.info("Application shutting down")
except Exception:
_startup_logger.exception("Error during shutdown logging")
# Send shutdown notification # Send shutdown notification
try: notify_shutdown()
notify_shutdown()
except Exception:
_startup_logger.exception("Error sending shutdown notification")
app = FastAPI( app = FastAPI(
@@ -342,19 +207,8 @@ app.add_middleware(CSRFMiddleware, config=settings)
# See SECURITY_AUDIT.md Infrastructure Security section # See SECURITY_AUDIT.md Infrastructure Security section
app.add_middleware(AuditLogMiddleware, config=settings) app.add_middleware(AuditLogMiddleware, config=settings)
# 3) Session Middleware (for request.session to work) # 3) Session Middleware (for request.session to work)
def _get_session_max_age() -> int: app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET)
"""Compute session max-age at startup time."""
try:
from app.utils.session_manager import get_session_max_age_seconds
return get_session_max_age_seconds()
except Exception:
return 30 * 86400 # 30 days default fallback
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET, max_age=_get_session_max_age())
# 3a) CORS Middleware - handles cross-origin requests and preflight (OPTIONS) responses. # 3a) CORS Middleware - handles cross-origin requests and preflight (OPTIONS) responses.
# Disabled by default: set CORS_ENABLED=True only when NOT using a reverse proxy # Disabled by default: set CORS_ENABLED=True only when NOT using a reverse proxy
@@ -430,13 +284,15 @@ async def http_exception_handler(request: Request, exc: HTTPException):
# For frontend routes, return appropriate HTML templates # For frontend routes, return appropriate HTML templates
# Handle 404 errors with a custom template # Handle 404 errors with a custom template
if exc.status_code == 404: if exc.status_code == 404:
return _error_templates.TemplateResponse(request, "404.html", status_code=status.HTTP_404_NOT_FOUND) return _error_templates.TemplateResponse(
"404.html", {"request": request}, status_code=status.HTTP_404_NOT_FOUND
)
# For other HTTP errors, we could create specific templates or use a generic one # For other HTTP errors, we could create specific templates or use a generic one
# For now, return a simple error page # For now, return a simple error page
return _error_templates.TemplateResponse( return _error_templates.TemplateResponse(
request,
"404.html", # Reuse 404 template for other errors, or create a generic error template "404.html", # Reuse 404 template for other errors, or create a generic error template
{"request": request},
status_code=exc.status_code, status_code=exc.status_code,
) )
@@ -456,9 +312,8 @@ async def custom_500_handler(request: Request, exc: Exception):
# Serve the 500 template for non-API routes # Serve the 500 template for non-API routes
return _error_templates.TemplateResponse( return _error_templates.TemplateResponse(
request,
"500.html", "500.html",
context={"exc": exc}, {"request": request, "exc": exc},
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
) )
-7
View File
@@ -20,9 +20,6 @@ How it works:
Exempt paths (CSRF is not checked even for state-changing methods): Exempt paths (CSRF is not checked even for state-changing methods):
- ``/oauth-callback`` OAuth 2.0 callback; protected by the ``state`` parameter. - ``/oauth-callback`` OAuth 2.0 callback; protected by the ``state`` parameter.
- ``/api/qr-auth/claim`` Called by the unauthenticated mobile app; the
cryptographically-random, single-use challenge token provides equivalent
protection.
""" """
import logging import logging
@@ -42,10 +39,6 @@ CSRF_PROTECTED_METHODS = {"POST", "PUT", "DELETE", "PATCH"}
# their own replay-protection mechanism). # their own replay-protection mechanism).
CSRF_EXEMPT_PATHS = { CSRF_EXEMPT_PATHS = {
"/oauth-callback", "/oauth-callback",
# The mobile app calls this endpoint without a browser session/CSRF token.
# The cryptographically-random, single-use challenge token already provides
# equivalent protection against cross-site request forgery.
"/api/qr-auth/claim",
} }
-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,
)
-265
View File
@@ -84,19 +84,6 @@ class FileRecord(Base):
# Processing pipeline assigned to this file (NULL = use system default) # Processing pipeline assigned to this file (NULL = use system default)
pipeline_id = Column(Integer, ForeignKey(_PIPELINES_ID_FK), nullable=True, index=True) pipeline_id = Column(Integer, ForeignKey(_PIPELINES_ID_FK), nullable=True, index=True)
# Detected document language (ISO 639-1 code, e.g. "de", "en", "fr")
# Extracted from AI metadata during processing; cached here for fast access.
detected_language = Column(String(10), nullable=True)
# Default-language translation of the extracted text.
# Stored when the detected language differs from the user's/system default
# document language. Only the original text and this translation are persisted;
# other languages are translated on the fly via the AI provider.
default_language_text = Column(Text, nullable=True)
# ISO 639-1 code of the default-language translation stored above (e.g. "en").
default_language_code = Column(String(10), nullable=True)
# Timestamp when we inserted this record # Timestamp when we inserted this record
created_at = Column(DateTime(timezone=True), server_default=func.now(), index=True) created_at = Column(DateTime(timezone=True), server_default=func.now(), index=True)
@@ -211,28 +198,6 @@ class WebhookConfig(Base):
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) 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): class LocalUser(Base):
"""A locally-registered user authenticated by email and bcrypt password. """A locally-registered user authenticated by email and bcrypt password.
@@ -317,12 +282,6 @@ class UserProfile(Base):
# NULL means "auto-detect from browser Accept-Language header" # NULL means "auto-detect from browser Accept-Language header"
preferred_language = Column(String(10), nullable=True) preferred_language = Column(String(10), nullable=True)
# Default document language for translated versions (ISO 639-1 code).
# When a document's detected language differs from this value, the system
# automatically generates and stores a translation into this language.
# NULL means "use the global DEFAULT_DOCUMENT_LANGUAGE setting".
default_document_language = Column(String(10), nullable=True)
# UI colour scheme preference: "light" | "dark" | "system" (NULL = "system") # UI colour scheme preference: "light" | "dark" | "system" (NULL = "system")
preferred_theme = Column(String(10), nullable=True) preferred_theme = Column(String(10), nullable=True)
@@ -625,7 +584,6 @@ class IntegrationType:
EMAIL = "EMAIL" EMAIL = "EMAIL"
PAPERLESS = "PAPERLESS" PAPERLESS = "PAPERLESS"
RCLONE = "RCLONE" RCLONE = "RCLONE"
SHAREPOINT = "SHAREPOINT"
ICLOUD = "ICLOUD" ICLOUD = "ICLOUD"
ALL = { ALL = {
@@ -643,7 +601,6 @@ class IntegrationType:
EMAIL, EMAIL,
PAPERLESS, PAPERLESS,
RCLONE, RCLONE,
SHAREPOINT,
ICLOUD, ICLOUD,
} }
@@ -808,9 +765,6 @@ class ApiToken(Base):
created_at = Column(DateTime(timezone=True), server_default=func.now()) created_at = Column(DateTime(timezone=True), server_default=func.now())
revoked_at = Column(DateTime(timezone=True), nullable=True) 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): class SharedLink(Base):
"""Shareable, time-limited or view-limited document link. """Shareable, time-limited or view-limited document link.
@@ -962,51 +916,6 @@ class ScheduledJob(Base):
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) 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): class MobileDevice(Base):
"""Registered mobile device for push notifications. """Registered mobile device for push notifications.
@@ -1041,86 +950,6 @@ class MobileDevice(Base):
__table_args__ = (UniqueConstraint("owner_id", "push_token", name="uq_mobile_device_owner_token"),) __table_args__ = (UniqueConstraint("owner_id", "push_token", name="uq_mobile_device_owner_token"),)
class UserSession(Base):
"""Server-side session tracking for invalidation and device management.
Each row represents an active browser or app session. The ``session_token``
is stored in the user's cookie and validated on every authenticated request.
Revoking a row (``is_revoked=True``) immediately terminates that session
on the next request.
"""
__tablename__ = "user_sessions"
id = Column(Integer, primary_key=True, index=True)
# Cryptographically random token stored in the session cookie.
session_token = Column(String(128), unique=True, nullable=False, index=True)
# Stable owner identifier — matches FileRecord.owner_id.
user_id = Column(String, nullable=False, index=True)
# Client metadata for display in the session management UI.
ip_address = Column(String(45), nullable=True)
user_agent = Column(String(512), nullable=True)
device_info = Column(String(255), nullable=True)
is_revoked = Column(Boolean, nullable=False, default=False)
created_at = Column(DateTime(timezone=True), server_default=func.now())
last_active_at = Column(DateTime(timezone=True), server_default=func.now())
expires_at = Column(DateTime(timezone=True), nullable=False)
revoked_at = Column(DateTime(timezone=True), nullable=True)
class QRLoginChallenge(Base):
"""Time-limited QR code login challenge for mobile app authentication.
A logged-in web user generates a challenge that produces a QR code. The
mobile app scans the QR code and calls the claim endpoint with the
``challenge_token``. The server verifies the challenge is still valid,
unclaimed, and unexpired, then issues an API token for the mobile app.
Security properties:
* Time-bound (default 2 minutes).
* Single-use (``is_claimed`` prevents replay).
* Cryptographically random 64-byte token.
* Bound to the creating user only that user's mobile device receives a
token.
"""
__tablename__ = "qr_login_challenges"
id = Column(Integer, primary_key=True, index=True)
# Cryptographically random token encoded in the QR code.
challenge_token = Column(String(128), unique=True, nullable=False, index=True)
# The user who created this challenge (from the web session).
user_id = Column(String, nullable=False, index=True)
# Whether the challenge has been successfully claimed by a mobile app.
is_claimed = Column(Boolean, nullable=False, default=False)
# Whether the challenge has been explicitly cancelled or expired.
is_cancelled = Column(Boolean, nullable=False, default=False)
# IP address of the web client that created the challenge.
created_by_ip = Column(String(45), nullable=True)
# IP address of the mobile client that claimed the challenge.
claimed_by_ip = Column(String(45), nullable=True)
# Device name provided by the mobile app when claiming.
device_name = Column(String(255), nullable=True)
# The API token ID that was issued to the mobile app (for audit trail).
issued_token_id = Column(Integer, nullable=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
expires_at = Column(DateTime(timezone=True), nullable=False)
claimed_at = Column(DateTime(timezone=True), nullable=True)
class ComplianceTemplate(Base): class ComplianceTemplate(Base):
"""Pre-built compliance configuration templates (GDPR, HIPAA, SOC2). """Pre-built compliance configuration templates (GDPR, HIPAA, SOC2).
@@ -1192,97 +1021,3 @@ class PipelineRoutingRule(Base):
created_at = Column(DateTime(timezone=True), server_default=func.now()) created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=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
View File
@@ -78,7 +78,6 @@ def _convert_pdf_to_pdfa(input_path: str, output_path: str, pdfa_format: str = "
output_type, output_type,
"--quiet", "--quiet",
"--invalidate-digital-signatures", "--invalidate-digital-signatures",
"--", # end-of-options separator: prevents file paths from being interpreted as options
input_path, input_path,
output_path, output_path,
] ]
-24
View File
@@ -216,30 +216,6 @@ def embed_metadata_into_pdf(self, local_file_path: str, extracted_text: str, met
except Exception as search_exc: except Exception as search_exc:
logger.warning(f"[{task_id}] Meilisearch indexing failed (non-fatal): {search_exc}") logger.warning(f"[{task_id}] Meilisearch indexing failed (non-fatal): {search_exc}")
# Cache the detected language on the FileRecord and trigger
# default-language translation when the document is in a
# different language.
detected_lang = metadata.get("language") if metadata else None
if detected_lang and extracted_text:
try:
file_record.detected_language = detected_lang
db.commit()
from app.tasks.translate_to_default_language import translate_to_default_language
translate_to_default_language.delay(
file_id,
extracted_text,
detected_lang,
owner_id=file_record.owner_id,
)
logger.info(
f"[{task_id}] Queued default-language translation for file {file_id} "
f"(detected: {detected_lang})"
)
except Exception as trans_exc:
logger.warning(f"[{task_id}] Could not queue translation task (non-fatal): {trans_exc}")
# Persist the metadata into a JSON file with the same base name. # Persist the metadata into a JSON file with the same base name.
# Include file path references for traceability # Include file path references for traceability
logger.info(f"[{task_id}] Persisting metadata to JSON") logger.info(f"[{task_id}] Persisting metadata to JSON")
+1 -1
View File
@@ -76,7 +76,7 @@ def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: i
"Your task is to analyze the given text and return a well-structured JSON object.\n\n" "Your task is to analyze the given text and return a well-structured JSON object.\n\n"
"Extract and return the following fields:\n" "Extract and return the following fields:\n"
"1. **filename**: Machine-readable filename " "1. **filename**: Machine-readable filename "
"(YYYY-MM-DD_DescriptiveTitle, use only letters, numbers, periods, and underscores).\n" "(YYYY-MM-DD_DescriptiveTitle, use only letters, numbers, spaces, dashes, periods, and underscores).\n"
'2. **empfaenger**: The recipient, or "Unknown" if not found.\n' '2. **empfaenger**: The recipient, or "Unknown" if not found.\n'
'3. **absender**: The sender, or "Unknown" if not found.\n' '3. **absender**: The sender, or "Unknown" if not found.\n'
"4. **correspondent**: The entity or company that issued the document " "4. **correspondent**: The entity or company that issued the document "
+1 -13
View File
@@ -21,9 +21,8 @@ from app.tasks.send_to_all import (
# Import database and logging utils from main # Import database and logging utils from main
from app.utils import log_task_progress from app.utils import log_task_progress
# Import notification utilities # Import notification utility
from app.utils.notification import notify_file_processed from app.utils.notification import notify_file_processed
from app.utils.user_notification import notify_user_document_processed
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -140,15 +139,4 @@ def finalize_document_storage(self, original_file: str, processed_file: str, met
except Exception as e: except Exception as e:
logger.warning(f"[WARNING] Failed to send file processed notification: {e}") logger.warning(f"[WARNING] Failed to send file processed notification: {e}")
# 6. Send per-user notification
if owner_id:
try:
notify_user_document_processed(
owner_id=owner_id,
filename=os.path.basename(processed_file),
file_id=file_id,
)
except Exception as e:
logger.warning(f"[WARNING] Failed to send per-user processed notification: {e}")
return {"status": "Completed", "file": processed_file} return {"status": "Completed", "file": processed_file}
-6
View File
@@ -18,7 +18,6 @@ from app.utils.allowed_types import (
DEFAULT_CATEGORIES, DEFAULT_CATEGORIES,
get_allowed_types_for_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) # Database session for per-user IMAP accounts (imported lazily to avoid circular imports)
_db_session_factory = None _db_session_factory = None
@@ -406,11 +405,6 @@ def pull_inbox(
) )
processed_emails = load_processed_emails() 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: try:
mail = imaplib.IMAP4_SSL(host, port) if use_ssl else imaplib.IMAP4(host, port) mail = imaplib.IMAP4_SSL(host, port) if use_ssl else imaplib.IMAP4(host, port)
mail.login(username, password) mail.login(username, password)
-31
View File
@@ -10,7 +10,6 @@ from app.models import FileRecord, IntegrationDirection, UserIntegration
from app.tasks.retry_config import BaseTaskWithRetry from app.tasks.retry_config import BaseTaskWithRetry
from app.tasks.upload_to_dropbox import upload_to_dropbox from app.tasks.upload_to_dropbox import upload_to_dropbox
from app.tasks.upload_to_email import upload_to_email from app.tasks.upload_to_email import upload_to_email
from app.tasks.upload_to_evernote import upload_to_evernote
from app.tasks.upload_to_ftp import upload_to_ftp from app.tasks.upload_to_ftp import upload_to_ftp
from app.tasks.upload_to_google_drive import upload_to_google_drive from app.tasks.upload_to_google_drive import upload_to_google_drive
from app.tasks.upload_to_icloud import upload_to_icloud from app.tasks.upload_to_icloud import upload_to_icloud
@@ -19,7 +18,6 @@ from app.tasks.upload_to_onedrive import upload_to_onedrive
from app.tasks.upload_to_paperless import upload_to_paperless from app.tasks.upload_to_paperless import upload_to_paperless
from app.tasks.upload_to_s3 import upload_to_s3 from app.tasks.upload_to_s3 import upload_to_s3
from app.tasks.upload_to_sftp import upload_to_sftp from app.tasks.upload_to_sftp import upload_to_sftp
from app.tasks.upload_to_sharepoint import upload_to_sharepoint
from app.tasks.upload_to_webdav import upload_to_webdav from app.tasks.upload_to_webdav import upload_to_webdav
from app.utils.config_validator import get_provider_status from app.utils.config_validator import get_provider_status
from app.utils.logging import log_task_progress from app.utils.logging import log_task_progress
@@ -101,11 +99,6 @@ def _should_upload_to_email():
) )
def _should_upload_to_evernote():
token = getattr(settings, "evernote_auth_token", None)
return bool(getattr(settings, "evernote_enabled", True) and isinstance(token, str) and token.strip())
def _should_upload_to_onedrive(): def _should_upload_to_onedrive():
return bool( return bool(
getattr(settings, "onedrive_enabled", True) getattr(settings, "onedrive_enabled", True)
@@ -128,18 +121,6 @@ def _should_upload_to_icloud():
return bool(getattr(settings, "icloud_enabled", True) and settings.icloud_username and settings.icloud_password) return bool(getattr(settings, "icloud_enabled", True) and settings.icloud_username and settings.icloud_password)
def _should_upload_to_sharepoint():
return bool(
settings.sharepoint_client_id
and settings.sharepoint_client_secret
and settings.sharepoint_site_url
and (
settings.sharepoint_refresh_token
or (settings.sharepoint_tenant_id and settings.sharepoint_tenant_id != "common")
)
)
def get_configured_services_from_validator(): def get_configured_services_from_validator():
""" """
Use the config validator to determine which services are configured and enabled. Use the config validator to determine which services are configured and enabled.
@@ -157,10 +138,8 @@ def get_configured_services_from_validator():
"FTP Storage": "ftp", "FTP Storage": "ftp",
"SFTP Storage": "sftp", "SFTP Storage": "sftp",
"Email": "email", "Email": "email",
"Evernote": "evernote",
"OneDrive": "onedrive", "OneDrive": "onedrive",
"S3 Storage": "s3", "S3 Storage": "s3",
"SharePoint": "sharepoint",
"iCloud Drive": "icloud", "iCloud Drive": "icloud",
} }
@@ -261,11 +240,6 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
"should_upload": _should_upload_to_email, "should_upload": _should_upload_to_email,
"upload_func": upload_to_email, "upload_func": upload_to_email,
}, },
{
"name": "evernote",
"should_upload": _should_upload_to_evernote,
"upload_func": upload_to_evernote,
},
{ {
"name": "onedrive", "name": "onedrive",
"should_upload": _should_upload_to_onedrive, "should_upload": _should_upload_to_onedrive,
@@ -276,11 +250,6 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
"should_upload": _should_upload_to_s3, "should_upload": _should_upload_to_s3,
"upload_func": upload_to_s3, "upload_func": upload_to_s3,
}, },
{
"name": "sharepoint",
"should_upload": _should_upload_to_sharepoint,
"upload_func": upload_to_sharepoint,
},
{ {
"name": "icloud", "name": "icloud",
"should_upload": _should_upload_to_icloud, "should_upload": _should_upload_to_icloud,
-141
View File
@@ -1,141 +0,0 @@
#!/usr/bin/env python3
"""Celery task to translate extracted document text into the default target language.
This task is triggered after metadata extraction when the detected document
language differs from the user's (or system) default document language. The
translated text is persisted in ``FileRecord.default_language_text`` so that
users can always read a reference copy in their preferred language.
Other ad-hoc translations are generated on the fly via the ``/api/files/{id}/translate``
endpoint and are NOT persisted.
"""
import logging
from app.celery_app import celery
from app.config import settings
from app.database import SessionLocal
from app.models import FileRecord, UserProfile
from app.tasks.retry_config import BaseTaskWithRetry
from app.utils import log_task_progress
from app.utils.ai_provider import get_ai_provider
logger = logging.getLogger(__name__)
def _resolve_default_language(owner_id: str | None) -> str:
"""Return the default document language for the given owner.
Resolution order:
1. ``UserProfile.default_document_language`` (per-user override)
2. ``settings.default_document_language`` (global setting)
"""
if owner_id:
with SessionLocal() as db:
profile = db.query(UserProfile).filter_by(user_id=owner_id).first()
if profile and profile.default_document_language:
return profile.default_document_language
return settings.default_document_language
@celery.task(base=BaseTaskWithRetry, bind=True)
def translate_to_default_language(
self,
file_id: int,
extracted_text: str,
detected_language: str,
owner_id: str | None = None,
) -> dict:
"""Translate *extracted_text* into the default document language and persist the result.
Args:
file_id: Primary key of the :class:`FileRecord`.
extracted_text: The OCR / refined text in the document's original language.
detected_language: ISO 639-1 code of the document's detected language.
owner_id: Owner identifier used to resolve per-user language preference.
Returns:
A dict with ``status``, ``target_language``, and the translated text length.
"""
task_id = self.request.id
target_language = _resolve_default_language(owner_id)
# Nothing to do when the document is already in the target language.
if detected_language == target_language:
logger.info(
f"[{task_id}] Document {file_id} already in target language '{target_language}', skipping translation"
)
log_task_progress(
task_id,
"translate_to_default_language",
"skipped",
f"Document already in {target_language}",
file_id=file_id,
)
return {"status": "skipped", "reason": "already_in_target_language"}
logger.info(f"[{task_id}] Translating document {file_id} from '{detected_language}' to '{target_language}'")
log_task_progress(
task_id,
"translate_to_default_language",
"in_progress",
f"Translating from {detected_language} to {target_language}",
file_id=file_id,
)
try:
provider = get_ai_provider()
model = settings.ai_model or settings.openai_model
translated_text = provider.chat_completion(
messages=[
{
"role": "system",
"content": (
f"You are a professional translator. Translate the following text "
f"from {detected_language} to {target_language}. "
f"Preserve the original formatting, paragraph structure, and meaning. "
f"Do not add any commentary or explanation — output ONLY the translated text."
),
},
{"role": "user", "content": extracted_text},
],
model=model,
temperature=0.3,
)
# Persist the translation.
with SessionLocal() as db:
record = db.query(FileRecord).filter_by(id=file_id).first()
if record:
record.default_language_text = translated_text
record.default_language_code = target_language
record.detected_language = detected_language
db.commit()
logger.info(
f"[{task_id}] Stored default-language translation ({len(translated_text)} chars) for file {file_id}"
)
log_task_progress(
task_id,
"translate_to_default_language",
"success",
f"Translated {len(extracted_text)}{len(translated_text)} chars ({detected_language}{target_language})",
file_id=file_id,
)
return {
"status": "success",
"target_language": target_language,
"translated_length": len(translated_text),
}
except Exception as exc:
logger.exception(f"[{task_id}] Translation failed for file {file_id}: {exc}")
log_task_progress(
task_id,
"translate_to_default_language",
"failure",
f"Exception: {exc}",
file_id=file_id,
)
raise
+5 -34
View File
@@ -11,7 +11,6 @@ from email.mime.image import MIMEImage
from email.mime.multipart import MIMEMultipart from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText from email.mime.text import MIMEText
import pypdf
from jinja2 import Environment, FileSystemLoader, select_autoescape from jinja2 import Environment, FileSystemLoader, select_autoescape
from app.celery_app import celery from app.celery_app import celery
@@ -24,15 +23,6 @@ logger = logging.getLogger(__name__)
# Constants # Constants
_LOGO_FILENAME = "logo.png" _LOGO_FILENAME = "logo.png"
# Mapping from PDF metadata keys (with leading slash stripped) to application-specific names.
# This mirrors the inverse of the mapping used in app/tasks/embed_metadata_into_pdf.py.
_PDF_METADATA_KEY_MAP = {
"Title": "filename",
"Author": "absender",
"Subject": "document_type",
"Keywords": "tags",
}
def get_email_template(template_name="default.html"): def get_email_template(template_name="default.html"):
""" """
@@ -73,12 +63,9 @@ def extract_metadata_from_file(file_path):
""" """
Try to extract metadata from a file using several methods: Try to extract metadata from a file using several methods:
1. Check for a .json metadata file with the same name 1. Check for a .json metadata file with the same name
2. Extract embedded metadata from PDF using pypdf 2. Extract metadata from PDF if it's embedded
JSON metadata takes precedence; embedded PDF metadata fills in any missing Returns a dictionary of metadata or None if not found
fields using the application's standard key mapping (e.g., /Title → filename).
Returns a dictionary of metadata (may be empty if none found).
""" """
metadata = {} metadata = {}
@@ -89,28 +76,12 @@ def extract_metadata_from_file(file_path):
with open(metadata_path, "r", encoding="utf-8") as f: with open(metadata_path, "r", encoding="utf-8") as f:
metadata = json.load(f) metadata = json.load(f)
logger.info(f"Loaded metadata from external JSON file: {metadata_path}") logger.info(f"Loaded metadata from external JSON file: {metadata_path}")
return metadata
except Exception as e: except Exception as e:
logger.warning(f"Failed to load metadata from JSON file: {str(e)}") logger.warning(f"Failed to load metadata from JSON file: {str(e)}")
# Try to extract embedded metadata from PDF # TODO: For PDF files, try to extract embedded metadata using PyPDF2
if file_path.lower().endswith(".pdf") and os.path.exists(file_path): # This would require additional dependencies, so for now we'll just check for external JSON
try:
with open(file_path, "rb") as f:
pdf_reader = pypdf.PdfReader(f)
pdf_metadata = pdf_reader.metadata
if pdf_metadata:
for key, value in pdf_metadata.items():
# Remove the leading slash from PDF metadata keys (e.g., '/Title' -> 'Title')
clean_key = key[1:] if key.startswith("/") else key
# Map to application-specific key names where possible
mapped_key = _PDF_METADATA_KEY_MAP.get(clean_key, clean_key)
# Only set if not already present (JSON metadata takes precedence)
if mapped_key not in metadata:
metadata[mapped_key] = str(value)
logger.info(f"Extracted embedded metadata from PDF: {file_path}")
except Exception as e:
logger.warning(f"Failed to extract metadata from PDF {file_path}: {str(e)}")
return metadata return metadata
-242
View File
@@ -1,242 +0,0 @@
#!/usr/bin/env python3
import hashlib
import json
import logging
import mimetypes
import os
from html import escape
from typing import Any
from app.celery_app import celery
from app.config import settings
from app.tasks.retry_config import UploadTaskWithRetry
from app.utils import log_task_progress
logger = logging.getLogger(__name__)
_UNKNOWN_PLACEHOLDERS = {"", "Unknown", "unknown", "N/A", "n/a", "None", "none", "null"}
_MAX_EVERNOTE_TITLE_LENGTH = 255
def _get_evernote_sdk():
"""Import the Evernote SDK lazily so the missing dependency error is actionable."""
try:
from evernote.edam.notestore import NoteStore
from evernote.edam.type import ttypes as Types
from evernote.edam.userstore import UserStore
from thrift.protocol import TBinaryProtocol
from thrift.transport import THttpClient
except ImportError as exc:
raise RuntimeError("Evernote upload requires the evernote3 package. Install requirements.txt again.") from exc
return NoteStore, Types, UserStore, TBinaryProtocol, THttpClient
def _build_thrift_client(client_cls, url: str, binary_protocol, http_transport):
transport = http_transport.THttpClient(url)
protocol = binary_protocol.TBinaryProtocol(transport)
return client_cls(protocol)
def _get_note_store(auth_token: str):
NoteStore, Types, UserStore, TBinaryProtocol, THttpClient = _get_evernote_sdk()
base_url = (
"https://sandbox.evernote.com" if getattr(settings, "evernote_sandbox", False) else "https://www.evernote.com"
)
user_store = _build_thrift_client(UserStore.Client, f"{base_url}/edam/user", TBinaryProtocol, THttpClient)
user = user_store.getUser(auth_token)
shard_id = getattr(user, "shardId", None)
if not shard_id:
raise RuntimeError("Evernote user response did not include a shard ID")
note_store_url = f"{base_url}/shard/{shard_id}/notestore"
note_store = _build_thrift_client(NoteStore.Client, note_store_url, TBinaryProtocol, THttpClient)
return note_store, Types
def _load_metadata(file_path: str) -> dict[str, Any]:
"""Load extracted DocuElevate metadata from the companion JSON file, when present."""
json_path = os.path.splitext(file_path)[0] + ".json"
if not os.path.exists(json_path):
return {}
try:
with open(json_path, "r", encoding="utf-8") as metadata_file:
data = json.load(metadata_file)
except Exception as exc: # noqa: BLE001
logger.warning("Failed to load Evernote metadata from %s: %s", json_path, exc)
return {}
return data if isinstance(data, dict) else {}
def _normalize_metadata_value(value: Any) -> str:
if value is None:
return ""
if isinstance(value, (list, tuple, set)):
normalized = ", ".join(str(item) for item in value if item is not None)
elif isinstance(value, dict):
normalized = json.dumps(value, ensure_ascii=False, sort_keys=True)
else:
normalized = str(value)
normalized = normalized.strip()
return "" if normalized in _UNKNOWN_PLACEHOLDERS else normalized
def _metadata_rows(metadata: dict[str, Any]) -> list[tuple[str, str]]:
rows = []
for key in sorted(metadata):
value = _normalize_metadata_value(metadata[key])
if value:
rows.append((key, value))
return rows
def _extract_tags(metadata: dict[str, Any]) -> list[str]:
tags: list[str] = []
def add_tag(value: Any) -> None:
normalized = _normalize_metadata_value(value)
if normalized and normalized not in tags:
tags.append(normalized)
default_tags = getattr(settings, "evernote_default_tags", None)
if default_tags:
for tag in str(default_tags).split(","):
add_tag(tag)
metadata_tags = metadata.get("tags")
if isinstance(metadata_tags, str):
for tag in metadata_tags.split(","):
add_tag(tag)
elif isinstance(metadata_tags, (list, tuple, set)):
for tag in metadata_tags:
add_tag(tag)
return tags
def _note_title(file_path: str, metadata: dict[str, Any]) -> str:
title = (
_normalize_metadata_value(metadata.get("title"))
or _normalize_metadata_value(metadata.get("filename"))
or os.path.basename(file_path)
)
return title[:_MAX_EVERNOTE_TITLE_LENGTH]
def _build_enml(metadata: dict[str, Any], resource_hash: str, resource_mime: str, include_metadata: bool) -> str:
body_parts = ['<?xml version="1.0" encoding="UTF-8"?>']
body_parts.append('<!DOCTYPE en-note SYSTEM "http://xml.evernote.com/pub/enml2.dtd">')
body_parts.append("<en-note>")
if include_metadata:
rows = _metadata_rows(metadata)
if rows:
body_parts.append("<div><b>DocuElevate metadata</b></div>")
for key, value in rows:
body_parts.append(f"<div><b>{escape(key)}:</b> {escape(value)}</div>")
body_parts.append("<br/>")
body_parts.append(f'<en-media type="{escape(resource_mime)}" hash="{resource_hash}"/>')
body_parts.append("</en-note>")
return "".join(body_parts)
def _create_evernote_note(file_path: str, metadata: dict[str, Any], task_id: str):
auth_token = getattr(settings, "evernote_auth_token", None)
if not auth_token:
raise ValueError("Evernote auth token is not configured (EVERNOTE_AUTH_TOKEN)")
note_store, Types = _get_note_store(auth_token)
filename = os.path.basename(file_path)
with open(file_path, "rb") as pdf_file:
resource_body = pdf_file.read()
body_hash = hashlib.md5(resource_body).digest() # noqa: S324 - Evernote API requires MD5 resource hashes.
body_hash_hex = hashlib.md5(resource_body).hexdigest() # noqa: S324 - Evernote ENML references MD5 hashes.
resource_mime = mimetypes.guess_type(filename)[0] or "application/pdf"
data = Types.Data()
data.size = len(resource_body)
data.bodyHash = body_hash
data.body = resource_body
resource = Types.Resource()
resource.mime = resource_mime
resource.data = data
resource.attributes = Types.ResourceAttributes(fileName=filename)
note = Types.Note()
note.title = _note_title(file_path, metadata)
note.content = _build_enml(
metadata,
body_hash_hex,
resource_mime,
include_metadata=getattr(settings, "evernote_include_metadata", True),
)
note.resources = [resource]
notebook_guid = getattr(settings, "evernote_notebook_guid", None)
if notebook_guid:
note.notebookGuid = notebook_guid
tag_names = _extract_tags(metadata)
if tag_names:
note.tagNames = tag_names
created_note = note_store.createNote(auth_token, note)
logger.info("[%s] Created Evernote note %s for %s", task_id, getattr(created_note, "guid", None), file_path)
return created_note
@celery.task(base=UploadTaskWithRetry, bind=True)
def upload_to_evernote(self, file_path: str, file_id: int = None):
"""
Upload a document to Evernote by creating a note with metadata and a PDF attachment.
Args:
file_path: Path to the PDF file to upload
file_id: Optional file ID to associate with logs
"""
task_id = self.request.id
filename = os.path.basename(file_path)
logger.info("[%s] Starting Evernote upload: %s", task_id, file_path)
log_task_progress(
task_id, "upload_to_evernote", "in_progress", f"Uploading to Evernote: {filename}", file_id=file_id
)
if not os.path.exists(file_path):
error_msg = f"File not found: {file_path}"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_evernote", "failure", error_msg, file_id=file_id)
raise FileNotFoundError(error_msg)
if not getattr(settings, "evernote_auth_token", None):
error_msg = "Evernote auth token is not configured (EVERNOTE_AUTH_TOKEN)"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_evernote", "failure", error_msg, file_id=file_id)
raise ValueError(error_msg)
try:
metadata = _load_metadata(file_path)
created_note = _create_evernote_note(file_path, metadata, task_id)
except Exception as exc:
error_msg = f"Failed to upload to Evernote: {exc}"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_evernote", "failure", error_msg, file_id=file_id)
raise
note_guid = getattr(created_note, "guid", None)
log_task_progress(task_id, "upload_to_evernote", "success", f"Uploaded to Evernote: {note_guid}", file_id=file_id)
return {
"status": "Completed",
"file_path": file_path,
"evernote_note_guid": note_guid,
"evernote_title": getattr(created_note, "title", None),
"evernote_notebook_guid": getattr(created_note, "notebookGuid", None),
}
-338
View File
@@ -1,338 +0,0 @@
#!/usr/bin/env python3
"""Upload documents to Microsoft SharePoint via the Microsoft Graph API.
This module authenticates using MSAL (same OAuth2 flow as OneDrive) and
uploads files to a configurable SharePoint Online document library using
the chunked upload session approach for reliability with large files.
Key differences from the OneDrive provider:
- Uses ``/sites/{siteId}/drives/{driveId}`` instead of ``/me/drive``
- Requires a SharePoint site URL to resolve the site and drive IDs
- Targets a named document library (default: ``Documents``)
"""
import logging
import os
import time
import urllib.parse
import msal
import requests
from app.celery_app import celery
from app.config import settings
from app.tasks.retry_config import UploadTaskWithRetry
from app.utils import log_task_progress
logger = logging.getLogger(__name__)
def get_sharepoint_token() -> str:
"""Acquire a Microsoft Graph API access token for SharePoint.
Uses MSAL ``ConfidentialClientApplication`` with the refresh-token flow
(delegated permissions) or the client-credentials flow (application
permissions) depending on configuration.
Returns:
A valid access token string.
Raises:
ValueError: When required settings are missing or token acquisition fails.
"""
if not settings.sharepoint_client_id or not settings.sharepoint_client_secret:
raise ValueError("SharePoint client ID and client secret must be configured")
tenant = settings.sharepoint_tenant_id or "common"
logger.info("Using SharePoint tenant: %s", tenant)
scopes = ["https://graph.microsoft.com/.default"]
if settings.sharepoint_refresh_token:
app = msal.ConfidentialClientApplication(
client_id=settings.sharepoint_client_id,
client_credential=settings.sharepoint_client_secret,
authority=f"https://login.microsoftonline.com/{tenant}",
)
logger.info("Attempting to acquire SharePoint token using refresh token")
token_response = app.acquire_token_by_refresh_token(
refresh_token=settings.sharepoint_refresh_token, scopes=scopes
)
if "access_token" not in token_response:
error = token_response.get("error", "")
error_desc = token_response.get("error_description", "Unknown error")
logger.error("Failed to get SharePoint access token: %s - %s", error, error_desc)
raise ValueError(f"Failed to get SharePoint access token: {error} - {error_desc}")
if "refresh_token" in token_response:
settings.sharepoint_refresh_token = token_response["refresh_token"]
logger.info("Updated SharePoint refresh token in memory")
return token_response["access_token"]
elif settings.sharepoint_tenant_id and settings.sharepoint_tenant_id != "common":
authority = f"https://login.microsoftonline.com/{settings.sharepoint_tenant_id}"
app = msal.ConfidentialClientApplication(
client_id=settings.sharepoint_client_id,
client_credential=settings.sharepoint_client_secret,
authority=authority,
)
token_response = app.acquire_token_for_client(scopes=scopes)
if "access_token" not in token_response:
error = token_response.get("error", "")
error_desc = token_response.get("error_description", "Unknown error")
raise ValueError(f"Failed to get SharePoint access token: {error} - {error_desc}")
return token_response["access_token"]
else:
raise ValueError("For SharePoint, either a refresh token or a non-'common' tenant ID is required")
def resolve_sharepoint_drive(access_token: str, site_url: str, library_name: str) -> tuple[str, str]:
"""Resolve the Graph API site ID and drive ID for a SharePoint site.
Args:
access_token: Valid Microsoft Graph API token.
site_url: Full SharePoint site URL, e.g.
``https://tenant.sharepoint.com/sites/sitename``.
library_name: Display name of the document library (e.g. ``Documents``).
Returns:
A ``(site_id, drive_id)`` tuple.
Raises:
ValueError: When the site URL cannot be parsed.
RuntimeError: When the Graph API call fails.
"""
parsed = urllib.parse.urlparse(site_url)
hostname = parsed.hostname
site_path = parsed.path.rstrip("/")
if not hostname or not site_path:
raise ValueError(
f"Invalid SharePoint site URL '{site_url}'. Expected format: https://tenant.sharepoint.com/sites/sitename"
)
headers = {"Authorization": f"Bearer {access_token}"}
# Resolve site ID
site_api_url = f"https://graph.microsoft.com/v1.0/sites/{hostname}:{site_path}"
logger.info("Resolving SharePoint site: %s", site_api_url)
resp = requests.get(site_api_url, headers=headers, timeout=settings.http_request_timeout)
if resp.status_code != 200:
raise RuntimeError(f"Failed to resolve SharePoint site: {resp.status_code} - {resp.text}")
site_id = resp.json()["id"]
logger.info("Resolved SharePoint site ID: %s", site_id)
# Resolve drive ID from the document library name
drives_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives"
resp = requests.get(drives_url, headers=headers, timeout=settings.http_request_timeout)
if resp.status_code != 200:
raise RuntimeError(f"Failed to list SharePoint drives: {resp.status_code} - {resp.text}")
drives = resp.json().get("value", [])
drive_id = None
for drive in drives:
if drive.get("name", "").lower() == library_name.lower():
drive_id = drive["id"]
break
if not drive_id:
available = [d.get("name") for d in drives]
raise RuntimeError(f"Document library '{library_name}' not found on site. Available libraries: {available}")
logger.info("Resolved SharePoint drive ID: %s (library: %s)", drive_id, library_name)
return site_id, drive_id
def create_sharepoint_upload_session(
filename: str, folder_path: str | None, drive_id: str, site_id: str, access_token: str
) -> str:
"""Create a resumable upload session on a SharePoint document library.
Args:
filename: Name of the file to upload.
folder_path: Optional subfolder path inside the library.
drive_id: Graph API drive ID of the document library.
site_id: Graph API site ID.
access_token: Valid access token.
Returns:
The upload session URL for chunked PUT requests.
Raises:
RuntimeError: When session creation fails.
"""
base_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives/{drive_id}"
if folder_path:
folder_path = folder_path.strip("/")
path_components = folder_path.split("/")
encoded_path = "/".join(urllib.parse.quote(component) for component in path_components)
encoded_filename = urllib.parse.quote(filename)
item_path = f"/root:/{encoded_path}/{encoded_filename}:/createUploadSession"
else:
encoded_filename = urllib.parse.quote(filename)
item_path = f"/root:/{encoded_filename}:/createUploadSession"
url = f"{base_url}{item_path}"
request_body = {"item": {"@microsoft.graph.conflictBehavior": "replace"}}
headers = {"Authorization": f"Bearer {access_token}", "Content-Type": "application/json"}
logger.info("Creating SharePoint upload session for %s at path %s", filename, folder_path)
response = requests.post(url, headers=headers, json=request_body, timeout=settings.http_request_timeout)
if response.status_code == 200:
upload_url = response.json().get("uploadUrl")
logger.info("SharePoint upload session created for %s", filename)
return upload_url
else:
raise RuntimeError(f"Failed to create SharePoint upload session: {response.status_code} - {response.text}")
def upload_large_file_sharepoint(file_path: str, upload_url: str) -> dict:
"""Upload a file to SharePoint using a chunked upload session.
Args:
file_path: Local path to the file.
upload_url: The upload session URL from ``create_sharepoint_upload_session``.
Returns:
The Graph API response dict containing file metadata.
Raises:
RuntimeError: When a chunk upload fails after retries.
"""
file_size = os.path.getsize(file_path)
chunk_size = 10 * 1024 * 1024 # 10 MB
response = None
with open(file_path, "rb") as f:
chunk_number = 0
while True:
chunk = f.read(chunk_size)
if not chunk:
break
chunk_start = chunk_number * chunk_size
chunk_end = chunk_start + len(chunk) - 1
content_range = f"bytes {chunk_start}-{chunk_end}/{file_size}"
headers = {"Content-Length": str(len(chunk)), "Content-Range": content_range}
max_retries = 3
retry_delay = 2
for attempt in range(max_retries):
try:
response = requests.put(
upload_url, headers=headers, data=chunk, timeout=settings.http_request_timeout
)
if response.status_code in (201, 202):
break
else:
logger.warning(
"SharePoint chunk upload failed (attempt %d): %d", attempt + 1, response.status_code
)
if attempt < max_retries - 1:
time.sleep(retry_delay * (attempt + 1))
except Exception as e:
logger.warning("SharePoint chunk upload error (attempt %d): %s", attempt + 1, str(e))
if attempt < max_retries - 1:
time.sleep(retry_delay * (attempt + 1))
if response is None or response.status_code not in (201, 202):
status = response.status_code if response else "no response"
text = response.text if response else ""
raise RuntimeError(f"Failed to upload chunk after {max_retries} attempts: {status} - {text}")
chunk_number += 1
return response.json() if response else {}
@celery.task(base=UploadTaskWithRetry, bind=True)
def upload_to_sharepoint(self, file_path: str, file_id: int = None, folder_override: str = None):
"""Upload a file to SharePoint Online.
Args:
file_path: Path to the file to upload.
file_id: Optional file ID to associate with logs.
folder_override: Optional folder path override.
Returns:
A dict with upload status and file details.
Raises:
FileNotFoundError: When the file does not exist.
ValueError: When SharePoint is not configured.
RuntimeError: When the upload fails.
"""
task_id = self.request.id
logger.info("[%s] Starting SharePoint upload: %s", task_id, file_path)
log_task_progress(
task_id,
"upload_to_sharepoint",
"in_progress",
f"Uploading to SharePoint: {os.path.basename(file_path)}",
file_id=file_id,
)
if not os.path.exists(file_path):
error_msg = f"File not found: {file_path}"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
raise FileNotFoundError(error_msg)
filename = os.path.basename(file_path)
if not settings.sharepoint_client_id:
error_msg = "SharePoint client ID is not configured"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
raise ValueError(error_msg)
if not settings.sharepoint_site_url:
error_msg = "SharePoint site URL is not configured"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
raise ValueError(error_msg)
try:
access_token = get_sharepoint_token()
library_name = settings.sharepoint_document_library or "Documents"
site_id, drive_id = resolve_sharepoint_drive(access_token, settings.sharepoint_site_url, library_name)
folder_path = folder_override if folder_override is not None else settings.sharepoint_folder_path
upload_url = create_sharepoint_upload_session(filename, folder_path, drive_id, site_id, access_token)
result = upload_large_file_sharepoint(file_path, upload_url)
web_url = result.get("webUrl", "Not available")
logger.info("[%s] Successfully uploaded %s to SharePoint", task_id, filename)
logger.info("[%s] File accessible at: %s", task_id, web_url)
log_task_progress(
task_id, "upload_to_sharepoint", "success", f"Uploaded to SharePoint: {filename}", file_id=file_id
)
return {
"status": "Completed",
"file_path": file_path,
"sharepoint_path": f"{folder_path or ''}/{filename}",
"web_url": web_url,
}
except Exception as e:
error_msg = f"Failed to upload {filename} to SharePoint: {str(e)}"
logger.error("[%s] %s", task_id, error_msg)
log_task_progress(task_id, "upload_to_sharepoint", "failure", error_msg, file_id=file_id)
raise RuntimeError(error_msg) from e
-108
View File
@@ -571,113 +571,6 @@ def _upload_rclone(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], t
return {"status": "Completed", "rclone_dest": dest} return {"status": "Completed", "rclone_dest": dest}
def _upload_sharepoint(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
"""Upload *file_path* to SharePoint using per-user MSAL credentials."""
import urllib.parse
import msal
import requests as _requests
client_id = creds.get("client_id") or ""
client_secret = creds.get("client_secret") or ""
refresh_token = creds.get("refresh_token") or ""
tenant = cfg.get("tenant_id") or "common"
site_url = cfg.get("site_url") or ""
library_name = cfg.get("document_library") or "Documents"
folder_path = cfg.get("folder_path") or ""
if not (client_id and client_secret):
raise ValueError("SharePoint integration is missing client_id or client_secret in credentials")
if not site_url:
raise ValueError("SharePoint integration is missing site_url in config")
scopes = ["https://graph.microsoft.com/.default"]
msal_app = msal.ConfidentialClientApplication(
client_id=client_id,
client_credential=client_secret,
authority=f"https://login.microsoftonline.com/{tenant}",
)
if refresh_token:
token_resp = msal_app.acquire_token_by_refresh_token(refresh_token=refresh_token, scopes=scopes)
else:
token_resp = msal_app.acquire_token_for_client(scopes=scopes)
if "access_token" not in token_resp:
raise ValueError(f"SharePoint token acquisition failed: {token_resp.get('error_description', 'unknown')}")
access_token = token_resp["access_token"]
headers = {"Authorization": f"Bearer {access_token}"}
# Resolve site ID
parsed = urllib.parse.urlparse(site_url)
hostname = parsed.hostname
site_path = parsed.path.rstrip("/")
if not hostname or not site_path:
raise ValueError(f"Invalid SharePoint site URL: {site_url}")
resp = _requests.get(f"https://graph.microsoft.com/v1.0/sites/{hostname}:{site_path}", headers=headers, timeout=30)
resp.raise_for_status()
site_id = resp.json()["id"]
# Resolve drive ID
resp = _requests.get(f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives", headers=headers, timeout=30)
resp.raise_for_status()
drive_id = None
for drive in resp.json().get("value", []):
if drive.get("name", "").lower() == library_name.lower():
drive_id = drive["id"]
break
if not drive_id:
raise RuntimeError(f"Document library '{library_name}' not found on SharePoint site")
filename = os.path.basename(file_path)
# Build upload-session URL
base_url = f"https://graph.microsoft.com/v1.0/sites/{site_id}/drives/{drive_id}"
if folder_path:
folder_path = folder_path.strip("/")
encoded_path = "/".join(urllib.parse.quote(p) for p in folder_path.split("/"))
encoded_file = urllib.parse.quote(filename)
item_path = f"/root:/{encoded_path}/{encoded_file}:/createUploadSession"
else:
encoded_file = urllib.parse.quote(filename)
item_path = f"/root:/{encoded_file}:/createUploadSession"
session_url = f"{base_url}{item_path}"
session_headers = {"Authorization": f"Bearer {access_token}", "Content-Type": "application/json"}
resp = _requests.post(
session_url,
headers=session_headers,
json={"item": {"@microsoft.graph.conflictBehavior": "replace"}},
timeout=30,
)
resp.raise_for_status()
upload_url = resp.json()["uploadUrl"]
file_size = os.path.getsize(file_path)
chunk_size = 10 * 1024 * 1024
with open(file_path, "rb") as fh:
chunk_num = 0
while True:
chunk = fh.read(chunk_size)
if not chunk:
break
start = chunk_num * chunk_size
end = start + len(chunk) - 1
upload_headers = {
"Content-Length": str(len(chunk)),
"Content-Range": f"bytes {start}-{end}/{file_size}",
}
upload_resp = _requests.put(upload_url, headers=upload_headers, data=chunk, timeout=120)
if upload_resp.status_code not in (201, 202):
raise RuntimeError(f"SharePoint chunk upload failed: {upload_resp.status_code}")
chunk_num += 1
logger.info("[%s] SharePoint upload complete: %s/%s", task_id, folder_path, filename)
return {"status": "Completed", "sharepoint_folder": folder_path, "filename": filename}
def _upload_icloud(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]: def _upload_icloud(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
"""Upload *file_path* to iCloud Drive using per-user credentials. """Upload *file_path* to iCloud Drive using per-user credentials.
@@ -722,7 +615,6 @@ _UPLOAD_HANDLERS = {
IntegrationType.PAPERLESS: _upload_paperless, IntegrationType.PAPERLESS: _upload_paperless,
IntegrationType.EMAIL: _upload_email, IntegrationType.EMAIL: _upload_email,
IntegrationType.RCLONE: _upload_rclone, IntegrationType.RCLONE: _upload_rclone,
IntegrationType.SHAREPOINT: _upload_sharepoint,
IntegrationType.ICLOUD: _upload_icloud, IntegrationType.ICLOUD: _upload_icloud,
} }
+3 -3
View File
@@ -55,12 +55,12 @@ def upload_with_rclone(self, file_path: str, destination: str):
try: try:
# Ensure the remote path exists (create folders if needed) # Ensure the remote path exists (create folders if needed)
mkdir_cmd = ["rclone", "mkdir", "--config", rclone_config_path, "--", destination] mkdir_cmd = ["rclone", "mkdir", "--config", rclone_config_path, destination]
subprocess.run(mkdir_cmd, check=True, capture_output=True) # noqa: S603 subprocess.run(mkdir_cmd, check=True, capture_output=True) # noqa: S603
# Construct the upload command # Construct the upload command
upload_cmd = ["rclone", "copy", "--config", rclone_config_path, "--progress", "--", file_path, destination] upload_cmd = ["rclone", "copy", "--config", rclone_config_path, file_path, destination, "--progress"]
log_task_progress(task_id, "rclone_upload", "in_progress", f"Executing rclone copy to {destination}") log_task_progress(task_id, "rclone_upload", "in_progress", f"Executing rclone copy to {destination}")
@@ -71,7 +71,7 @@ def upload_with_rclone(self, file_path: str, destination: str):
if result.returncode == 0: if result.returncode == 0:
# Try to get a public link if possible # Try to get a public link if possible
try: try:
link_cmd = ["rclone", "link", "--config", rclone_config_path, "--", f"{destination}/{filename}"] link_cmd = ["rclone", "link", "--config", rclone_config_path, f"{destination}/{filename}"]
link_result = subprocess.run(link_cmd, capture_output=True, text=True, check=False) # noqa: S603 link_result = subprocess.run(link_cmd, capture_output=True, text=True, check=False) # noqa: S603
public_url = link_result.stdout.strip() if link_result.returncode == 0 else None public_url = link_result.stdout.strip() if link_result.returncode == 0 else None
except (subprocess.SubprocessError, OSError) as e: except (subprocess.SubprocessError, OSError) as e:
+1 -9
View File
@@ -68,8 +68,6 @@ IMAGE_MIME_TYPES: set[str] = {
"image/tiff", "image/tiff",
"image/webp", "image/webp",
"image/svg+xml", "image/svg+xml",
"image/heic",
"image/heif",
} }
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -126,8 +124,6 @@ ALLOWED_EXTENSIONS: set[str] = {
".tif", ".tif",
".webp", ".webp",
".svg", ".svg",
".heic",
".heif",
# Web # Web
".html", ".html",
".htm", ".htm",
@@ -238,7 +234,7 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
}, },
"images": { "images": {
"label": "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( "mime_types": frozenset(
{ {
"image/jpeg", "image/jpeg",
@@ -249,8 +245,6 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
"image/tiff", "image/tiff",
"image/webp", "image/webp",
"image/svg+xml", "image/svg+xml",
"image/heic",
"image/heif",
} }
), ),
"extensions": frozenset( "extensions": frozenset(
@@ -264,8 +258,6 @@ FILE_TYPE_CATEGORIES: dict[str, dict] = {
".tif", ".tif",
".webp", ".webp",
".svg", ".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),
)
-37
View File
@@ -174,21 +174,6 @@ def get_provider_status() -> dict[str, dict[str, object]]:
}, },
} }
providers["Evernote"] = {
"name": "Evernote",
"icon": "fa-brands fa-evernote",
"configured": bool(getattr(settings, "evernote_auth_token", None)),
"enabled": getattr(settings, "evernote_enabled", True),
"description": "Create Evernote notes with document metadata and PDF attachments",
"details": {
"auth_token": mask_sensitive_value(getattr(settings, "evernote_auth_token", None)),
"sandbox": getattr(settings, "evernote_sandbox", False),
"notebook_guid": getattr(settings, "evernote_notebook_guid", "Not set"),
"default_tags": getattr(settings, "evernote_default_tags", "Not set"),
"include_metadata": getattr(settings, "evernote_include_metadata", True),
},
}
# Add FTP configuration to providers # Add FTP configuration to providers
providers["FTP Storage"] = { providers["FTP Storage"] = {
"name": "FTP Storage", "name": "FTP Storage",
@@ -311,28 +296,6 @@ def get_provider_status() -> dict[str, dict[str, object]]:
}, },
} }
# Check SharePoint configuration
providers["SharePoint"] = {
"name": "SharePoint",
"icon": "fa-brands fa-microsoft",
"configured": bool(
getattr(settings, "sharepoint_client_id", None)
and getattr(settings, "sharepoint_client_secret", None)
and getattr(settings, "sharepoint_site_url", None)
),
"enabled": True,
"description": "Store documents in Microsoft SharePoint Online",
"details": {
"client_id": getattr(settings, "sharepoint_client_id", "Not set"),
"client_secret": mask_sensitive_value(getattr(settings, "sharepoint_client_secret", None)),
"tenant_id": getattr(settings, "sharepoint_tenant_id", "Not set"),
"refresh_token": mask_sensitive_value(getattr(settings, "sharepoint_refresh_token", None)),
"site_url": getattr(settings, "sharepoint_site_url", "Not set"),
"document_library": getattr(settings, "sharepoint_document_library", "Not set"),
"folder_path": getattr(settings, "sharepoint_folder_path", "Not set"),
},
}
# Check S3 configuration # Check S3 configuration
providers["S3 Storage"] = { providers["S3 Storage"] = {
"name": "S3 Storage", "name": "S3 Storage",
-6
View File
@@ -145,12 +145,6 @@ def validate_storage_configs() -> dict[str, list[str]]:
email_issues.append("DEST_EMAIL_DEFAULT_RECIPIENT is not configured") email_issues.append("DEST_EMAIL_DEFAULT_RECIPIENT is not configured")
issues["email"] = email_issues issues["email"] = email_issues
# Validate Evernote
evernote_issues = []
if not getattr(settings, "evernote_auth_token", None):
evernote_issues.append("EVERNOTE_AUTH_TOKEN is not configured")
issues["evernote"] = evernote_issues
# Validate S3 # Validate S3
s3_issues = [] s3_issues = []
if not getattr(settings, "s3_bucket_name", None): if not getattr(settings, "s3_bucket_name", None):
+3 -8
View File
@@ -12,10 +12,9 @@ The utility:
""" """
import logging import logging
import re
from typing import Any from typing import Any
from sqlalchemy import MetaData, create_engine, func, inspect, select, table from sqlalchemy import MetaData, create_engine, inspect, text
from sqlalchemy.engine import Engine from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import make_url from sqlalchemy.engine.url import make_url
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
@@ -85,13 +84,9 @@ def preview_migration(source_url: str) -> dict[str, Any]:
total = 0 total = 0
with src_engine.connect() as conn: with src_engine.connect() as conn:
for table_name in tables: for table_name in tables:
if not re.match(r"^[a-zA-Z0-9_]+$", table_name):
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 # table_name is safe — sourced from inspect().get_table_names(), not user input
t = table(table_name) quoted_table = conn.dialect.identifier_preparer.quote(table_name)
query = select(func.count()).select_from(t) row = conn.execute(text(f"SELECT COUNT(*) FROM {quoted_table}")).fetchone() # noqa: S608
row = conn.execute(query).fetchone()
count = row[0] if row else 0 count = row[0] if row else 0
result.append({"name": table_name, "row_count": count}) result.append({"name": table_name, "row_count": count})
total += count total += count
-54
View File
@@ -1,54 +0,0 @@
import logging
import os
logger = logging.getLogger(__name__)
def update_env_file(settings_to_update: dict[str, str]) -> bool:
"""
Updates the .env file with the given settings (best-effort).
Creates or modifies existing keys.
Args:
settings_to_update: A dictionary mapping uppercase env var names to their new string values.
Returns:
True if the file was successfully updated, False otherwise.
"""
try:
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
if not os.path.exists(env_path):
logger.warning(f".env file not found at {env_path}, skipping file write")
return False
logger.info(f"Updating settings in {env_path}")
with open(env_path, "r") as f:
env_lines = f.readlines()
updated = set()
new_env_lines = []
for line in env_lines:
stripped_line = line.rstrip()
is_updated = False
for key, value in settings_to_update.items():
if stripped_line.startswith(f"{key}=") or stripped_line.startswith(f"# {key}="):
new_env_lines.append(f"{key}={value}")
updated.add(key)
is_updated = True
break
if not is_updated:
new_env_lines.append(stripped_line)
for key, value in settings_to_update.items():
if key not in updated:
new_env_lines.append(f"{key}={value}")
with open(env_path, "w") as f:
f.write("\n".join(new_env_lines) + "\n")
logger.info("Successfully updated settings in .env file")
return True
except Exception as env_err:
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
return False
+1 -12
View File
@@ -7,19 +7,8 @@ def hash_file(filepath: str | Path, chunk_size: int = 65536) -> str:
Returns the SHA-256 hash of the file at 'filepath'. Returns the SHA-256 hash of the file at 'filepath'.
Reads the file in chunks to handle large files efficiently. 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() sha256 = hashlib.sha256()
with open(filepath_obj, "rb") as f: with open(filepath, "rb") as f:
while True: while True:
data = f.read(chunk_size) data = f.read(chunk_size)
if not data: if not data:
+77 -79
View File
@@ -34,91 +34,89 @@ logger = logging.getLogger(__name__)
SUPPORTED_LANGUAGES: list[dict[str, str]] = [ SUPPORTED_LANGUAGES: list[dict[str, str]] = [
# --- Tier 1: Primary European languages --- # --- Tier 1: Primary European languages ---
# flag: lowercase ISO 3166-1 alpha-2 country code used with the flag-icons CSS library {"code": "en", "name": "English", "native": "English", "flag": "🇬🇧"},
# (e.g. "gb" → <span class="fi fi-gb">). Regional codes like "gb-wls" are also supported. {"code": "de", "name": "German", "native": "Deutsch", "flag": "🇩🇪"},
{"code": "en", "name": "English", "native": "English", "flag": "gb"}, {"code": "fr", "name": "French", "native": "Français", "flag": "🇫🇷"},
{"code": "de", "name": "German", "native": "Deutsch", "flag": "de"}, {"code": "es", "name": "Spanish", "native": "Español", "flag": "🇪🇸"},
{"code": "fr", "name": "French", "native": "Français", "flag": "fr"}, {"code": "it", "name": "Italian", "native": "Italiano", "flag": "🇮🇹"},
{"code": "es", "name": "Spanish", "native": "Español", "flag": "es"}, {"code": "pt", "name": "Portuguese", "native": "Português", "flag": "🇵🇹"},
{"code": "it", "name": "Italian", "native": "Italiano", "flag": "it"},
{"code": "pt", "name": "Portuguese", "native": "Português", "flag": "pt"},
# --- Tier 2: Western & Northern European --- # --- Tier 2: Western & Northern European ---
{"code": "nl", "name": "Dutch", "native": "Nederlands", "flag": "nl"}, {"code": "nl", "name": "Dutch", "native": "Nederlands", "flag": "🇳🇱"},
{"code": "nb", "name": "Norwegian Bokmål", "native": "Norsk bokmål", "flag": "no"}, {"code": "nb", "name": "Norwegian Bokmål", "native": "Norsk bokmål", "flag": "🇳🇴"},
{"code": "no", "name": "Norwegian", "native": "Norsk", "flag": "no"}, {"code": "no", "name": "Norwegian", "native": "Norsk", "flag": "🇳🇴"},
{"code": "da", "name": "Danish", "native": "Dansk", "flag": "dk"}, {"code": "da", "name": "Danish", "native": "Dansk", "flag": "🇩🇰"},
{"code": "sv", "name": "Swedish", "native": "Svenska", "flag": "se"}, {"code": "sv", "name": "Swedish", "native": "Svenska", "flag": "🇸🇪"},
{"code": "fi", "name": "Finnish", "native": "Suomi", "flag": "fi"}, {"code": "fi", "name": "Finnish", "native": "Suomi", "flag": "🇫🇮"},
{"code": "is", "name": "Icelandic", "native": "Íslenska", "flag": "is"}, {"code": "is", "name": "Icelandic", "native": "Íslenska", "flag": "🇮🇸"},
{"code": "ga", "name": "Irish", "native": "Gaeilge", "flag": "ie"}, {"code": "ga", "name": "Irish", "native": "Gaeilge", "flag": "🇮🇪"},
{"code": "lb", "name": "Luxembourgish", "native": "Lëtzebuergesch", "flag": "lu"}, {"code": "lb", "name": "Luxembourgish", "native": "Lëtzebuergesch", "flag": "🇱🇺"},
{"code": "ca", "name": "Catalan", "native": "Català", "flag": "es"}, # no dedicated ISO flag; use Spain {"code": "ca", "name": "Catalan", "native": "Català", "flag": "🏴"},
{"code": "cy", "name": "Welsh", "native": "Cymraeg", "flag": "gb-wls"}, # flag-icons GB region code {"code": "cy", "name": "Welsh", "native": "Cymraeg", "flag": "🏴󠁧󠁢󠁷󠁬󠁳󠁿"}, # Wales subdivision flag (U+1F3F4 + tag chars)
{"code": "fy", "name": "Western Frisian", "native": "Frysk", "flag": "nl"}, {"code": "fy", "name": "Western Frisian", "native": "Frysk", "flag": "🇳🇱"},
{"code": "gl", "name": "Galician", "native": "Galego", "flag": "es"}, {"code": "gl", "name": "Galician", "native": "Galego", "flag": "🇪🇸"},
{"code": "li", "name": "Limburgish", "native": "Limburgs", "flag": "nl"}, {"code": "li", "name": "Limburgish", "native": "Limburgs", "flag": "🇳🇱"},
{"code": "vls", "name": "Flemish", "native": "West-Vlams", "flag": "be"}, {"code": "vls", "name": "Flemish", "native": "West-Vlams", "flag": "🇧🇪"},
{"code": "nds", "name": "Low German", "native": "Plattdüütsch", "flag": "de"}, {"code": "nds", "name": "Low German", "native": "Plattdüütsch", "flag": "🇩🇪"},
# --- Tier 3: Central & Eastern European --- # --- Tier 3: Central & Eastern European ---
{"code": "pl", "name": "Polish", "native": "Polski", "flag": "pl"}, {"code": "pl", "name": "Polish", "native": "Polski", "flag": "🇵🇱"},
{"code": "cs", "name": "Czech", "native": "Čeština", "flag": "cz"}, {"code": "cs", "name": "Czech", "native": "Čeština", "flag": "🇨🇿"},
{"code": "sk", "name": "Slovak", "native": "Slovenčina", "flag": "sk"}, {"code": "sk", "name": "Slovak", "native": "Slovenčina", "flag": "🇸🇰"},
{"code": "hu", "name": "Hungarian", "native": "Magyar", "flag": "hu"}, {"code": "hu", "name": "Hungarian", "native": "Magyar", "flag": "🇭🇺"},
{"code": "sl", "name": "Slovenian", "native": "Slovenščina", "flag": "si"}, {"code": "sl", "name": "Slovenian", "native": "Slovenščina", "flag": "🇸🇮"},
{"code": "hr", "name": "Croatian", "native": "Hrvatski", "flag": "hr"}, {"code": "hr", "name": "Croatian", "native": "Hrvatski", "flag": "🇭🇷"},
{"code": "ro", "name": "Romanian", "native": "Română", "flag": "ro"}, {"code": "ro", "name": "Romanian", "native": "Română", "flag": "🇷🇴"},
{"code": "bg", "name": "Bulgarian", "native": "Български", "flag": "bg"}, {"code": "bg", "name": "Bulgarian", "native": "Български", "flag": "🇧🇬"},
{"code": "el", "name": "Greek", "native": "Ελληνικά", "flag": "gr"}, {"code": "el", "name": "Greek", "native": "Ελληνικά", "flag": "🇬🇷"},
{"code": "et", "name": "Estonian", "native": "Eesti", "flag": "ee"}, {"code": "et", "name": "Estonian", "native": "Eesti", "flag": "🇪🇪"},
{"code": "lv", "name": "Latvian", "native": "Latviešu", "flag": "lv"}, {"code": "lv", "name": "Latvian", "native": "Latviešu", "flag": "🇱🇻"},
{"code": "lt", "name": "Lithuanian", "native": "Lietuvių", "flag": "lt"}, {"code": "lt", "name": "Lithuanian", "native": "Lietuvių", "flag": "🇱🇹"},
{"code": "sr", "name": "Serbian", "native": "Српски", "flag": "rs"}, {"code": "sr", "name": "Serbian", "native": "Српски", "flag": "🇷🇸"},
# --- Tier 4: Non-EU European, Middle Eastern & African --- # --- Tier 4: Non-EU European, Middle Eastern & African ---
{"code": "tr", "name": "Turkish", "native": "Türkçe", "flag": "tr"}, {"code": "tr", "name": "Turkish", "native": "Türkçe", "flag": "🇹🇷"},
{"code": "uk", "name": "Ukrainian", "native": "Українська", "flag": "ua"}, {"code": "uk", "name": "Ukrainian", "native": "Українська", "flag": "🇺🇦"},
{"code": "he", "name": "Hebrew", "native": "עברית", "flag": "il"}, {"code": "he", "name": "Hebrew", "native": "עברית", "flag": "🇮🇱"},
{"code": "ar", "name": "Arabic", "native": "العربية", "flag": "sa"}, {"code": "ar", "name": "Arabic", "native": "العربية", "flag": "🇸🇦"},
{"code": "fa", "name": "Persian", "native": "فارسی", "flag": "ir"}, {"code": "fa", "name": "Persian", "native": "فارسی", "flag": "🇮🇷"},
{"code": "af", "name": "Afrikaans", "native": "Afrikaans", "flag": "za"}, {"code": "af", "name": "Afrikaans", "native": "Afrikaans", "flag": "🇿🇦"},
# --- Tier 5: Asian languages --- # --- Tier 5: Asian languages ---
{"code": "zh", "name": "Chinese", "native": "中文", "flag": "cn"}, {"code": "zh", "name": "Chinese", "native": "中文", "flag": "🇨🇳"},
{"code": "zh-TW", "name": "Traditional Chinese", "native": "繁體中文", "flag": "tw"}, {"code": "zh-TW", "name": "Traditional Chinese", "native": "繁體中文", "flag": "🇹🇼"},
{"code": "ja", "name": "Japanese", "native": "日本語", "flag": "jp"}, {"code": "ja", "name": "Japanese", "native": "日本語", "flag": "🇯🇵"},
{"code": "ko", "name": "Korean", "native": "한국어", "flag": "kr"}, {"code": "ko", "name": "Korean", "native": "한국어", "flag": "🇰🇷"},
{"code": "vi", "name": "Vietnamese", "native": "Tiếng Việt", "flag": "vn"}, {"code": "vi", "name": "Vietnamese", "native": "Tiếng Việt", "flag": "🇻🇳"},
{"code": "pa", "name": "Punjabi", "native": "ਪੰਜਾਬੀ", "flag": "in"}, {"code": "pa", "name": "Punjabi", "native": "ਪੰਜਾਬੀ", "flag": "🇮🇳"},
{"code": "kn", "name": "Kannada", "native": "ಕನ್ನಡ", "flag": "in"}, {"code": "kn", "name": "Kannada", "native": "ಕನ್ನಡ", "flag": "🇮🇳"},
{"code": "hi", "name": "Hindi", "native": "हिन्दी", "flag": "in"}, {"code": "hi", "name": "Hindi", "native": "हिन्दी", "flag": "🇮🇳"},
{"code": "bn", "name": "Bengali", "native": "বাংলা", "flag": "bd"}, {"code": "bn", "name": "Bengali", "native": "বাংলা", "flag": "🇧🇩"},
{"code": "gu", "name": "Gujarati", "native": "ગુજરાતી", "flag": "in"}, {"code": "gu", "name": "Gujarati", "native": "ગુજરાતી", "flag": "🇮🇳"},
{"code": "ml", "name": "Malayalam", "native": "മലയാളം", "flag": "in"}, {"code": "ml", "name": "Malayalam", "native": "മലയാളം", "flag": "🇮🇳"},
{"code": "mr", "name": "Marathi", "native": "मराठी", "flag": "in"}, {"code": "mr", "name": "Marathi", "native": "मराठी", "flag": "🇮🇳"},
{"code": "ta", "name": "Tamil", "native": "தமிழ்", "flag": "in"}, {"code": "ta", "name": "Tamil", "native": "தமிழ்", "flag": "🇮🇳"},
{"code": "te", "name": "Telugu", "native": "తెలుగు", "flag": "in"}, {"code": "te", "name": "Telugu", "native": "తెలుగు", "flag": "🇮🇳"},
{"code": "ur", "name": "Urdu", "native": "اردو", "flag": "pk"}, {"code": "ur", "name": "Urdu", "native": "اردو", "flag": "🇵🇰"},
{"code": "si", "name": "Sinhala", "native": "සිංහල", "flag": "lk"}, {"code": "si", "name": "Sinhala", "native": "සිංහල", "flag": "🇱🇰"},
{"code": "ne", "name": "Nepali", "native": "नेपाली", "flag": "np"}, {"code": "ne", "name": "Nepali", "native": "नेपाली", "flag": "🇳🇵"},
{"code": "th", "name": "Thai", "native": "ไทย", "flag": "th"}, {"code": "th", "name": "Thai", "native": "ไทย", "flag": "🇹🇭"},
{"code": "km", "name": "Khmer", "native": "ខ្មែរ", "flag": "kh"}, {"code": "km", "name": "Khmer", "native": "ខ្មែរ", "flag": "🇰🇭"},
{"code": "id", "name": "Indonesian", "native": "Bahasa Indonesia", "flag": "id"}, {"code": "id", "name": "Indonesian", "native": "Bahasa Indonesia", "flag": "🇮🇩"},
{"code": "ms", "name": "Malay", "native": "Bahasa Melayu", "flag": "my"}, {"code": "ms", "name": "Malay", "native": "Bahasa Melayu", "flag": "🇲🇾"},
{"code": "jv", "name": "Javanese", "native": "Basa Jawa", "flag": "id"}, {"code": "jv", "name": "Javanese", "native": "Basa Jawa", "flag": "🇮🇩"},
{"code": "tl", "name": "Tagalog", "native": "Filipino", "flag": "ph"}, {"code": "tl", "name": "Tagalog", "native": "Filipino", "flag": "🇵🇭"},
{"code": "mn", "name": "Mongolian", "native": "Монгол", "flag": "mn"}, {"code": "mn", "name": "Mongolian", "native": "Монгол", "flag": "🇲🇳"},
{"code": "kk", "name": "Kazakh", "native": "Қазақ тілі", "flag": "kz"}, {"code": "kk", "name": "Kazakh", "native": "Қазақ тілі", "flag": "🇰🇿"},
{"code": "uz", "name": "Uzbek", "native": "Oʻzbekcha", "flag": "uz"}, {"code": "uz", "name": "Uzbek", "native": "Oʻzbekcha", "flag": "🇺🇿"},
{"code": "az", "name": "Azerbaijani", "native": "Azərbaycan dili", "flag": "az"}, {"code": "az", "name": "Azerbaijani", "native": "Azərbaycan dili", "flag": "🇦🇿"},
{"code": "hy", "name": "Armenian", "native": "Հայերեն", "flag": "am"}, {"code": "hy", "name": "Armenian", "native": "Հայերեն", "flag": "🇦🇲"},
{"code": "ka", "name": "Georgian", "native": "ქართული", "flag": "ge"}, {"code": "ka", "name": "Georgian", "native": "ქართული", "flag": "🇬🇪"},
# --- Tier 6: African languages --- # --- Tier 6: African languages ---
{"code": "sw", "name": "Swahili", "native": "Kiswahili", "flag": "ke"}, {"code": "sw", "name": "Swahili", "native": "Kiswahili", "flag": "🇰🇪"},
{"code": "am", "name": "Amharic", "native": "አማርኛ", "flag": "et"}, {"code": "am", "name": "Amharic", "native": "አማርኛ", "flag": "🇪🇹"},
{"code": "ha", "name": "Hausa", "native": "Hausa", "flag": "ng"}, {"code": "ha", "name": "Hausa", "native": "Hausa", "flag": "🇳🇬"},
{"code": "yo", "name": "Yoruba", "native": "Yorùbá", "flag": "ng"}, {"code": "yo", "name": "Yoruba", "native": "Yorùbá", "flag": "🇳🇬"},
{"code": "ig", "name": "Igbo", "native": "Igbo", "flag": "ng"}, {"code": "ig", "name": "Igbo", "native": "Igbo", "flag": "🇳🇬"},
{"code": "zu", "name": "Zulu", "native": "isiZulu", "flag": "za"}, {"code": "zu", "name": "Zulu", "native": "isiZulu", "flag": "🇿🇦"},
# --- Tier 7: Constructed & other languages --- # --- Tier 7: Constructed & other languages ---
{"code": "eo", "name": "Esperanto", "native": "Esperanto", "flag": "un"}, # UN flag for international language {"code": "eo", "name": "Esperanto", "native": "Esperanto", "flag": "🌍"},
] ]
SUPPORTED_LANGUAGE_CODES: set[str] = {lang["code"] for lang in SUPPORTED_LANGUAGES} SUPPORTED_LANGUAGE_CODES: set[str] = {lang["code"] for lang in SUPPORTED_LANGUAGES}
+5 -31
View File
@@ -1,7 +1,6 @@
import ipaddress import ipaddress
import logging import logging
import socket import socket
from urllib.parse import urlsplit, urlunsplit
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -28,33 +27,8 @@ def is_private_ip(hostname: str) -> bool:
return True return True
return False return False
except (socket.gaierror, socket.error): except (socket.gaierror, socket.error):
# Cannot resolve. # Cannot resolve - allow for testing/development
# Fail securely: block unresolved domains to prevent DNS rebinding # In production, DNS should work properly
# and SSRF bypasses via unresolvable addresses. # Log this for debugging
logger.warning(f"Could not resolve hostname (blocking securely): {hostname}") logger.warning(f"Could not resolve hostname: {hostname}")
return True return False # Changed from True to False to allow external domains in tests
def join_url(base: str, *parts: str) -> str:
"""
Safely join a base URL with one or more path parts.
Uses urllib.parse to correctly handle scheme/netloc/query/fragment so that
only the path component is modified. Leading and trailing slashes are
stripped from each part before joining, preventing double-slash sequences
at segment boundaries without touching the scheme separator or query string.
Examples:
join_url("https://example.com/dav/", "/remote/", "file.pdf")
-> "https://example.com/dav/remote/file.pdf"
"""
parsed = urlsplit(base)
# Strip each part once and filter out empty segments; use walrus operator
# to avoid calling strip twice per iteration.
stripped_parts = [s for p in parts if (s := p.strip("/"))]
base_path = parsed.path.rstrip("/")
new_path = base_path + "/" + "/".join(stripped_parts) if stripped_parts else base_path
# Ensure path is non-empty so the reconstructed URL is valid.
if not new_path:
new_path = "/"
return urlunsplit((parsed.scheme, parsed.netloc, new_path, parsed.query, parsed.fragment))
-505
View File
@@ -1,505 +0,0 @@
"""Server-side session management utilities.
Provides helpers for creating, validating, and revoking user sessions.
Sessions are tracked in the ``user_sessions`` table and referenced by a
cryptographically random token stored in the browser cookie. This enables
the "log off everywhere" feature and per-session revocation.
"""
from __future__ import annotations
import logging
import secrets
from datetime import datetime, timedelta, timezone
from sqlalchemy.orm import Session
from app.config import settings
from app.models import ApiToken, QRLoginChallenge, UserSession
logger = logging.getLogger(__name__)
def _ensure_tz_aware(dt: datetime | None) -> datetime | None:
"""Return *dt* with UTC tzinfo if it is naive, or unchanged if already aware.
SQLite does not persist timezone information, so datetimes read back from
the database are offset-naive. This helper normalises them for safe
comparison with ``datetime.now(timezone.utc)``.
"""
if dt is not None and dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
return dt
def get_session_lifetime_days() -> int:
"""Return the effective session lifetime in days.
If ``session_lifetime_custom_days`` is set it takes precedence over
``session_lifetime_days``.
"""
custom = getattr(settings, "session_lifetime_custom_days", None)
if custom is not None and isinstance(custom, int) and custom > 0:
return custom
return max(1, getattr(settings, "session_lifetime_days", 30))
def get_session_max_age_seconds() -> int:
"""Return the session max-age in seconds for the cookie."""
return get_session_lifetime_days() * 86400
def create_session(
db: Session,
user_id: str,
ip_address: str | None = None,
user_agent: str | None = None,
) -> UserSession:
"""Create a new server-side session record.
Args:
db: Database session.
user_id: Stable owner identifier.
ip_address: Client IP address.
user_agent: Client User-Agent header.
Returns:
The newly created ``UserSession`` instance.
"""
session_token = secrets.token_urlsafe(64)
now = datetime.now(timezone.utc)
lifetime_days = get_session_lifetime_days()
expires_at = now + timedelta(days=lifetime_days)
device_info = _parse_device_info(user_agent)
user_session = UserSession(
session_token=session_token,
user_id=user_id,
ip_address=ip_address,
user_agent=(user_agent or "")[:512],
device_info=device_info,
created_at=now,
last_active_at=now,
expires_at=expires_at,
)
try:
db.add(user_session)
db.commit()
db.refresh(user_session)
except Exception:
db.rollback()
logger.exception("Failed to create session for user_id=%s", user_id)
raise
logger.info(
"[SESSION] Created session id=%s user=%s device=%r expires=%s",
user_session.id,
user_id,
device_info,
expires_at.isoformat(),
)
return user_session
def validate_session(db: Session, session_token: str) -> UserSession | None:
"""Validate a session token and return the session if valid.
A session is valid when:
* It exists in the database.
* ``is_revoked`` is ``False``.
* ``expires_at`` is in the future.
Side-effect: updates ``last_active_at`` on valid sessions.
Returns:
The ``UserSession`` if valid, else ``None``.
"""
if not session_token:
return None
now = datetime.now(timezone.utc)
user_session = db.query(UserSession).filter(UserSession.session_token == session_token).first()
if not user_session:
logger.debug("[SESSION] Token not found in database")
return None
if user_session.is_revoked:
logger.debug("[SESSION] Session id=%s is revoked", user_session.id)
return None
if user_session.expires_at:
expires = _ensure_tz_aware(user_session.expires_at)
if expires < now:
logger.debug("[SESSION] Session id=%s has expired", user_session.id)
return None
# Update last_active_at (throttled to avoid excessive writes)
last_active = _ensure_tz_aware(user_session.last_active_at)
if not last_active or (now - last_active).total_seconds() > 60:
try:
user_session.last_active_at = now
db.commit()
except Exception:
db.rollback()
logger.debug("[SESSION] Failed to update last_active_at for session id=%s", user_session.id)
return user_session
def revoke_session(db: Session, session_id: int, user_id: str) -> bool:
"""Revoke a single session by ID.
Args:
db: Database session.
session_id: The session record ID to revoke.
user_id: The owner ensures a user can only revoke their own sessions.
Returns:
``True`` if the session was found and revoked, ``False`` otherwise.
"""
user_session = db.get(UserSession, session_id)
if not user_session or user_session.user_id != user_id:
return False
now = datetime.now(timezone.utc)
user_session.is_revoked = True
user_session.revoked_at = now
try:
db.commit()
except Exception:
db.rollback()
raise
logger.info("[SESSION] Revoked session id=%s user=%s", session_id, user_id)
return True
def revoke_all_sessions(
db: Session,
user_id: str,
*,
except_session_id: int | None = None,
revoke_api_tokens: bool = True,
) -> int:
"""Revoke all active sessions for a user ("log off everywhere").
Args:
db: Database session.
user_id: The owner whose sessions should be revoked.
except_session_id: If provided, keep this session active (the
current browser session).
revoke_api_tokens: If ``True``, also revoke all active API tokens.
Returns:
Number of sessions revoked.
"""
now = datetime.now(timezone.utc)
query = db.query(UserSession).filter(
UserSession.user_id == user_id,
UserSession.is_revoked.is_(False),
)
if except_session_id is not None:
query = query.filter(UserSession.id != except_session_id)
sessions = query.all()
count = 0
for s in sessions:
s.is_revoked = True
s.revoked_at = now
count += 1
if revoke_api_tokens:
tokens = (
db.query(ApiToken)
.filter(
ApiToken.owner_id == user_id,
ApiToken.is_active.is_(True),
)
.all()
)
for t in tokens:
t.is_active = False
t.revoked_at = now
try:
db.commit()
except Exception:
db.rollback()
raise
logger.info(
"[SESSION] Revoked all sessions for user=%s (count=%d, except_session_id=%s, tokens_revoked=%s)",
user_id,
count,
except_session_id,
revoke_api_tokens,
)
return count
def list_user_sessions(db: Session, user_id: str) -> list[UserSession]:
"""Return all non-revoked, non-expired sessions for a user.
Results are ordered by most recently active first.
"""
now = datetime.now(timezone.utc)
sessions = (
db.query(UserSession)
.filter(
UserSession.user_id == user_id,
UserSession.is_revoked.is_(False),
)
.order_by(UserSession.last_active_at.desc())
.all()
)
# Filter expired sessions in Python to handle timezone-naive datetimes (SQLite)
result = []
for s in sessions:
expires = _ensure_tz_aware(s.expires_at)
if expires and expires > now:
result.append(s)
return result
def cleanup_expired_sessions(db: Session) -> int:
"""Delete sessions that expired more than 7 days ago.
Intended to be called periodically (e.g. via Celery beat) to keep the
table from growing unbounded.
Returns:
Number of rows deleted.
"""
cutoff = datetime.now(timezone.utc) - timedelta(days=7)
count = db.query(UserSession).filter(UserSession.expires_at < cutoff).delete(synchronize_session=False)
try:
db.commit()
except Exception:
db.rollback()
raise
if count:
logger.info("[SESSION] Cleaned up %d expired sessions", count)
return count
# ---------------------------------------------------------------------------
# QR login helpers
# ---------------------------------------------------------------------------
def create_qr_challenge(db: Session, user_id: str, ip_address: str | None = None) -> QRLoginChallenge:
"""Create a new QR login challenge.
Args:
db: Database session.
user_id: The authenticated web user creating the challenge.
ip_address: IP address of the web client.
Returns:
The newly created ``QRLoginChallenge``.
"""
token = secrets.token_urlsafe(64)
ttl = getattr(settings, "qr_login_challenge_ttl_seconds", 120)
now = datetime.now(timezone.utc)
expires_at = now + timedelta(seconds=ttl)
challenge = QRLoginChallenge(
challenge_token=token,
user_id=user_id,
created_by_ip=ip_address,
created_at=now,
expires_at=expires_at,
)
try:
db.add(challenge)
db.commit()
db.refresh(challenge)
except Exception:
db.rollback()
logger.exception("Failed to create QR login challenge for user_id=%s", user_id)
raise
logger.info("[QR_AUTH] Challenge created: id=%s user=%s expires=%s", challenge.id, user_id, expires_at.isoformat())
return challenge
def validate_qr_challenge(db: Session, challenge_token: str) -> QRLoginChallenge | None:
"""Validate a QR challenge token without claiming it.
Returns the challenge if it exists, is not expired, not claimed,
and not cancelled. Returns ``None`` otherwise.
"""
if not challenge_token:
return None
now = datetime.now(timezone.utc)
challenge = db.query(QRLoginChallenge).filter(QRLoginChallenge.challenge_token == challenge_token).first()
if not challenge:
return None
if challenge.is_claimed or challenge.is_cancelled:
return None
expires = _ensure_tz_aware(challenge.expires_at)
if expires and expires < now:
return None
return challenge
def claim_qr_challenge(
db: Session,
challenge_token: str,
device_name: str = "Mobile App",
ip_address: str | None = None,
) -> dict | None:
"""Claim a QR challenge and issue an API token.
This is the critical security path. The challenge is validated,
marked as claimed atomically, and an API token is issued for the
user who created the challenge.
Args:
db: Database session.
challenge_token: The token from the QR code.
device_name: Name provided by the mobile app.
ip_address: IP address of the claiming mobile device.
Returns:
Dict with ``token`` (plaintext), ``token_id``, ``name``, ``owner_id``
and ``created_at`` on success, or ``None`` if the challenge is invalid.
"""
from app.api.api_tokens import generate_api_token, hash_token
challenge = validate_qr_challenge(db, challenge_token)
if not challenge:
logger.warning("[QR_AUTH] Invalid or expired challenge token attempted")
return None
now = datetime.now(timezone.utc)
# Mark as claimed first to prevent race conditions
challenge.is_claimed = True
challenge.claimed_at = now
challenge.claimed_by_ip = ip_address
challenge.device_name = device_name
# Generate API token for the mobile app
token_name = f"Mobile App (QR) {device_name}"
plaintext = generate_api_token()
token_hash_value = hash_token(plaintext)
prefix = plaintext[:12]
db_token = ApiToken(
owner_id=challenge.user_id,
name=token_name,
token_hash=token_hash_value,
token_prefix=prefix,
)
try:
db.add(db_token)
db.flush()
challenge.issued_token_id = db_token.id
db.commit()
db.refresh(db_token)
except Exception:
db.rollback()
logger.exception("[QR_AUTH] Failed to issue token for challenge id=%s", challenge.id)
raise
logger.info(
"[QR_AUTH] Challenge claimed: id=%s user=%s device=%r token_id=%s",
challenge.id,
challenge.user_id,
device_name,
db_token.id,
)
return {
"token": plaintext,
"token_id": db_token.id,
"name": token_name,
"owner_id": challenge.user_id,
"created_at": db_token.created_at,
}
def get_challenge_status(db: Session, challenge_id: int, user_id: str) -> dict | None:
"""Get the current status of a QR challenge (for polling from the web UI).
Returns:
Dict with ``status`` ("pending", "claimed", "expired", "cancelled")
and metadata, or ``None`` if the challenge doesn't belong to the user.
"""
challenge = db.get(QRLoginChallenge, challenge_id)
if not challenge or challenge.user_id != user_id:
return None
now = datetime.now(timezone.utc)
expires = _ensure_tz_aware(challenge.expires_at)
if challenge.is_claimed:
status = "claimed"
elif challenge.is_cancelled:
status = "cancelled"
elif expires and expires < now:
status = "expired"
else:
status = "pending"
return {
"id": challenge.id,
"status": status,
"device_name": challenge.device_name,
"claimed_at": challenge.claimed_at,
"expires_at": challenge.expires_at,
}
def _parse_device_info(user_agent: str | None) -> str | None:
"""Extract a human-readable device description from User-Agent.
This is a lightweight parser not a full UA library that covers
the most common browsers and platforms.
"""
if not user_agent:
return None
ua = user_agent.lower()
# Platform detection
platform = "Unknown"
if "iphone" in ua:
platform = "iPhone"
elif "ipad" in ua:
platform = "iPad"
elif "android" in ua:
platform = "Android"
elif "macintosh" in ua or "mac os" in ua:
platform = "macOS"
elif "windows" in ua:
platform = "Windows"
elif "linux" in ua:
platform = "Linux"
elif "cros" in ua:
platform = "ChromeOS"
# Browser detection
browser = "Unknown Browser"
if "edg/" in ua or "edge/" in ua:
browser = "Edge"
elif "opr/" in ua or "opera" in ua:
browser = "Opera"
elif "chrome/" in ua and "safari/" in ua:
browser = "Chrome"
elif "safari/" in ua and "chrome/" not in ua:
browser = "Safari"
elif "firefox/" in ua:
browser = "Firefox"
elif "docuelevate" in ua:
browser = "DocuElevate App"
return f"{browser} on {platform}"
+2 -634
View File
@@ -39,50 +39,6 @@ SETTING_METADATA = {
"required": True, "required": True,
"restart_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": { "workdir": {
"category": "Core", "category": "Core",
"description": "Working directory for file storage and processing", "description": "Working directory for file storage and processing",
@@ -99,18 +55,6 @@ SETTING_METADATA = {
"required": True, # Required for OAuth redirects and external URLs "required": True, # Required for OAuth redirects and external URLs
"restart_required": True, "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": { "debug": {
"category": "Core", "category": "Core",
"description": "Enable debug mode for verbose logging", "description": "Enable debug mode for verbose logging",
@@ -190,38 +134,6 @@ SETTING_METADATA = {
"required": True, # Required when auth_enabled=True (validated in config.py) "required": True, # Required when auth_enabled=True (validated in config.py)
"restart_required": True, "restart_required": True,
}, },
"session_lifetime_days": {
"category": "Authentication",
"description": "Session lifetime in days (default 30). Determines how long a user stays logged in.",
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"session_lifetime_custom_days": {
"category": "Authentication",
"description": "Override session_lifetime_days with a custom value. Takes precedence when set.",
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"qr_login_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).",
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": False,
},
"admin_username": { "admin_username": {
"category": "Authentication", "category": "Authentication",
"description": "Admin username for local authentication", "description": "Admin username for local authentication",
@@ -270,17 +182,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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 Login Providers
"social_auth_google_enabled": { "social_auth_google_enabled": {
"category": "Social Login", "category": "Social Login",
@@ -311,20 +212,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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": { "social_auth_microsoft_enabled": {
"category": "Social Login", "category": "Social Login",
"description": ( "description": (
@@ -367,20 +254,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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": { "social_auth_apple_enabled": {
"category": "Social Login", "category": "Social Login",
"description": ( "description": (
@@ -429,19 +302,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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": { "social_auth_dropbox_enabled": {
"category": "Social Login", "category": "Social Login",
"description": ( "description": (
@@ -469,182 +329,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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 # AI Services
"openai_api_key": { "openai_api_key": {
"category": "AI Services", "category": "AI Services",
@@ -838,19 +522,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
# Document Translation
"default_document_language": {
"category": "AI Services",
"description": (
"ISO 639-1 language code for the default document translation target "
"(e.g. 'en', 'de', 'fr'). Documents whose detected language differs "
"are automatically translated into this language after processing."
),
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
# OCR Engine Configuration # OCR Engine Configuration
"ocr_providers": { "ocr_providers": {
"category": "OCR Engines", "category": "OCR Engines",
@@ -1011,18 +682,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_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 # Storage Providers - Nextcloud
"nextcloud_enabled": { "nextcloud_enabled": {
"category": "Storage Providers", "category": "Storage Providers",
@@ -1089,55 +748,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
# Storage Providers - Evernote
"evernote_enabled": {
"category": "Storage Providers",
"description": "Enable Evernote as an upload destination. When disabled, no documents will be sent to Evernote even if credentials are configured.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
"evernote_auth_token": {
"category": "Storage Providers",
"description": "Evernote developer token for note creation",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"evernote_sandbox": {
"category": "Storage Providers",
"description": "Use the Evernote sandbox environment instead of production",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
"evernote_notebook_guid": {
"category": "Storage Providers",
"description": "Optional Evernote notebook GUID for uploaded notes",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"evernote_default_tags": {
"category": "Storage Providers",
"description": "Comma-separated Evernote tags to apply to uploaded notes",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"evernote_include_metadata": {
"category": "Storage Providers",
"description": "Include extracted document metadata in Evernote note content",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Storage Providers - Google Drive # Storage Providers - Google Drive
"google_drive_enabled": { "google_drive_enabled": {
"category": "Storage Providers", "category": "Storage Providers",
@@ -1252,63 +862,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
# Storage Providers - SharePoint
"sharepoint_client_id": {
"category": "Storage Providers",
"description": "SharePoint Azure AD application (client) ID",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"sharepoint_client_secret": {
"category": "Storage Providers",
"description": "SharePoint Azure AD client secret",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"sharepoint_tenant_id": {
"category": "Storage Providers",
"description": "SharePoint Azure AD tenant ID (use 'common' for multi-tenant apps)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"sharepoint_refresh_token": {
"category": "Storage Providers",
"description": "SharePoint OAuth refresh token",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"sharepoint_site_url": {
"category": "Storage Providers",
"description": "SharePoint site URL (e.g. https://tenant.sharepoint.com/sites/sitename)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"sharepoint_document_library": {
"category": "Storage Providers",
"description": "SharePoint document library name (default: 'Documents')",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"sharepoint_folder_path": {
"category": "Storage Providers",
"description": "Subfolder path inside the SharePoint document library",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Storage Providers - WebDAV # Storage Providers - WebDAV
"webdav_enabled": { "webdav_enabled": {
"category": "Storage Providers", "category": "Storage Providers",
@@ -2224,30 +1777,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_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 # Notifications Settings
"notification_urls": { "notification_urls": {
"category": "Notifications", "category": "Notifications",
@@ -2346,18 +1875,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_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": { "compliance_enabled": {
"category": "Feature Flags", "category": "Feature Flags",
"description": ( "description": (
@@ -2370,28 +1887,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
"factory_reset_on_startup": {
"category": "Feature Flags",
"description": (
"Wipe all user data on every startup so the instance always starts fresh. "
"Useful for demo/testing environments. Default: False."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"enable_factory_reset": {
"category": "Feature Flags",
"description": (
"Show the System Reset page in the admin UI. Allows administrators to "
"trigger a full data wipe or a wipe-and-reimport from the web interface. Default: False."
),
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Backup / Restore # Backup / Restore
"backup_enabled": { "backup_enabled": {
"category": "Backup", "category": "Backup",
@@ -2415,26 +1910,14 @@ SETTING_METADATA = {
"category": "Backup", "category": "Backup",
"description": ( "description": (
"Storage provider for remote backup copies. " "Storage provider for remote backup copies. "
"Accepted values: s3, dropbox, google_drive, onedrive, sharepoint, nextcloud, webdav, ftp, sftp, email. " "Accepted values: s3, dropbox, google_drive, onedrive, nextcloud, webdav, ftp, sftp, email. "
"Leave empty to keep backups local only." "Leave empty to keep backups local only."
), ),
"type": "string", "type": "string",
"sensitive": False, "sensitive": False,
"required": False, "required": False,
"restart_required": False, "restart_required": False,
"options": [ "options": ["", "s3", "dropbox", "google_drive", "onedrive", "nextcloud", "webdav", "ftp", "sftp", "email"],
"",
"s3",
"dropbox",
"google_drive",
"onedrive",
"sharepoint",
"nextcloud",
"webdav",
"ftp",
"sftp",
"email",
],
}, },
"backup_remote_folder": { "backup_remote_folder": {
"category": "Backup", "category": "Backup",
@@ -2946,27 +2429,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
# Per-user upload rate limiting
"upload_rate_limit_per_user": {
"category": "Security",
"description": (
"Maximum number of uploads a single user may submit within upload_rate_limit_window seconds. "
"The health-aware limiter may reduce this dynamically under high Redis queue depth or CPU load. "
"Default: 20."
),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": False,
},
"upload_rate_limit_window": {
"category": "Security",
"description": ("Sliding window in seconds over which upload_rate_limit_per_user is enforced. Default: 60."),
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Rate Limiting # Rate Limiting
"rate_limiting_enabled": { "rate_limiting_enabled": {
"category": "Security", "category": "Security",
@@ -3166,63 +2628,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": False, "restart_required": False,
}, },
# Logging
"log_level": {
"category": "Observability",
"description": (
"Python logging level for the application root logger. "
"Accepts: DEBUG, INFO, WARNING, ERROR, CRITICAL. "
"When DEBUG=True and LOG_LEVEL is not explicitly set, "
"the effective level is automatically lowered to DEBUG."
),
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"log_format": {
"category": "Observability",
"description": (
"Log output format: 'text' (human-readable, default) or "
"'json' (structured JSON lines for SIEM / log aggregation)."
),
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"log_syslog_enabled": {
"category": "Observability",
"description": "Forward application logs to a syslog receiver in addition to stdout.",
"type": "boolean",
"sensitive": False,
"required": False,
"restart_required": True,
},
"log_syslog_host": {
"category": "Observability",
"description": "Hostname or IP of the syslog receiver for application logs.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
"log_syslog_port": {
"category": "Observability",
"description": "Port of the syslog receiver for application logs.",
"type": "integer",
"sensitive": False,
"required": False,
"restart_required": True,
},
"log_syslog_protocol": {
"category": "Observability",
"description": "Protocol for syslog transport: 'udp' or 'tcp'.",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": True,
},
# Observability Sentry # Observability Sentry
"sentry_dsn": { "sentry_dsn": {
"category": "Observability", "category": "Observability",
@@ -3281,43 +2686,6 @@ SETTING_METADATA = {
"required": False, "required": False,
"restart_required": True, "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: except Exception as exc:
logger.warning(f"Could not reload in-process settings: {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 # Re-check OCR language availability in the background whenever settings
# are updated. This ensures that if a user changes tesseract_language or # are updated. This ensures that if a user changes tesseract_language or
# easyocr_languages via the UI, the new language data is downloaded without # easyocr_languages via the UI, the new language data is downloaded without
+5 -21
View File
@@ -29,7 +29,7 @@ At average usage (~40 % of quota) margins improve to 55-65 % after tax.
from __future__ import annotations from __future__ import annotations
import logging import logging
from datetime import date, datetime, time, timedelta, timezone from datetime import date, datetime, timezone
from typing import Any from typing import Any
from sqlalchemy import func from sqlalchemy import func
@@ -317,20 +317,6 @@ def _today_utc() -> date:
return datetime.now(timezone.utc).date() return datetime.now(timezone.utc).date()
def _day_bounds_utc(day: date) -> tuple[datetime, datetime]:
start = datetime.combine(day, time.min, tzinfo=timezone.utc)
return start, start + timedelta(days=1)
def _month_bounds_utc(day: date) -> tuple[datetime, datetime]:
start = datetime.combine(day.replace(day=1), time.min, tzinfo=timezone.utc)
if start.month == 12:
end = start.replace(year=start.year + 1, month=1)
else:
end = start.replace(month=start.month + 1)
return start, end
def _scalar_count(query: Any) -> int: def _scalar_count(query: Any) -> int:
"""Execute a count query and return an int, defaulting to 0 for NULL.""" """Execute a count query and return an int, defaulting to 0 for NULL."""
return query.scalar() or 0 return query.scalar() or 0
@@ -349,13 +335,12 @@ def get_today_file_count(db: Session, owner_id: str) -> int:
"""Files processed by this user today (UTC, not counting duplicates).""" """Files processed by this user today (UTC, not counting duplicates)."""
from app.models import FileRecord from app.models import FileRecord
day_start, day_end = _day_bounds_utc(_today_utc()) today = _today_utc()
return _scalar_count( return _scalar_count(
db.query(func.count(FileRecord.id)).filter( db.query(func.count(FileRecord.id)).filter(
FileRecord.owner_id == owner_id, FileRecord.owner_id == owner_id,
FileRecord.is_duplicate.is_(False), FileRecord.is_duplicate.is_(False),
FileRecord.created_at >= day_start, func.date(FileRecord.created_at) == today,
FileRecord.created_at < day_end,
) )
) )
@@ -364,13 +349,12 @@ def get_month_file_count(db: Session, owner_id: str) -> int:
"""Files processed by this user this calendar month (UTC, not counting duplicates).""" """Files processed by this user this calendar month (UTC, not counting duplicates)."""
from app.models import FileRecord from app.models import FileRecord
month_start, month_end = _month_bounds_utc(_today_utc()) today = _today_utc()
return _scalar_count( return _scalar_count(
db.query(func.count(FileRecord.id)).filter( db.query(func.count(FileRecord.id)).filter(
FileRecord.owner_id == owner_id, FileRecord.owner_id == owner_id,
FileRecord.is_duplicate.is_(False), FileRecord.is_duplicate.is_(False),
FileRecord.created_at >= month_start, func.strftime("%Y-%m", FileRecord.created_at) == today.strftime("%Y-%m"),
FileRecord.created_at < month_end,
) )
) )
-296
View File
@@ -1,296 +0,0 @@
"""
System reset utilities for DocuElevate.
Provides functions to:
- Wipe all user data (database rows + work-files on disk) for a fresh start.
- Wipe with re-import: move original files to a dedicated folder, wipe
everything, then let the watch-folder mechanism re-ingest the files.
Security: All public functions in this module require admin-level access.
They MUST only be invoked from admin-guarded API/view endpoints.
"""
import logging
import shutil
from pathlib import Path
from sqlalchemy.orm import Session
from app.config import settings
logger = logging.getLogger(__name__)
# Subdirectories inside *workdir* that contain user-generated data.
# Everything else (app code, static assets, config) is left untouched.
_USER_DATA_SUBDIRS = ("original", "processed", "tmp", "pdfa", "backups")
# JSON cache files written by watch-folder / ingest tasks.
_CACHE_FILES = (
"watch_folder_processed.json",
"ftp_ingest_processed.json",
"sftp_ingest_processed.json",
"dropbox_ingest_processed.json",
"gdrive_ingest_processed.json",
"onedrive_ingest_processed.json",
"nextcloud_ingest_processed.json",
"s3_ingest_processed.json",
"webdav_ingest_processed.json",
"processed_mails.json",
"credential_failures.json",
)
# The folder name used for storing files prior to re-import.
REIMPORT_FOLDER_NAME = "reimport"
def _wipe_workdir_data(workdir: str) -> dict[str, int]:
"""Delete user data subdirectories and cache files inside *workdir*.
Leaves the workdir directory itself intact so the application can
continue to write into it. Also leaves any files that do not belong
to the known data subdirectories or caches.
Returns:
A dict with counts of deleted directories and files.
"""
workdir_path = Path(workdir)
deleted_dirs = 0
deleted_files = 0
# Remove data subdirectories
for subdir in _USER_DATA_SUBDIRS:
target = workdir_path / subdir
if target.is_dir():
shutil.rmtree(target)
logger.info("Deleted data directory: %s", target)
deleted_dirs += 1
# Remove cache / state JSON files
for cache_file in _CACHE_FILES:
target = workdir_path / cache_file
if target.is_file():
target.unlink()
logger.info("Deleted cache file: %s", target)
deleted_files += 1
# Also remove user_wf_*.json files (per-user watch folder caches)
for f in workdir_path.glob("user_wf_*.json"):
f.unlink()
logger.info("Deleted user watch-folder cache: %s", f)
deleted_files += 1
# Remove loose files in workdir root that are user uploads (uuid-named
# files like "a1b2c3d4-…pdf") but NOT application config files.
for entry in workdir_path.iterdir():
if entry.is_file() and entry.suffix.lower() in {
".pdf",
".png",
".jpg",
".jpeg",
".tiff",
".tif",
".docx",
".doc",
".xlsx",
".xls",
".pptx",
".heic",
".heif",
".webp",
".bmp",
".gif",
".txt",
".rtf",
".odt",
".ods",
".odp",
".csv",
".pages",
".numbers",
".keynote",
}:
entry.unlink()
logger.info("Deleted loose workdir file: %s", entry)
deleted_files += 1
return {"deleted_dirs": deleted_dirs, "deleted_files": deleted_files}
def _wipe_database(db: Session) -> dict[str, int]:
"""Delete all user-generated rows from the database.
Preserves schema (tables, migrations) and system-seeded rows that will
be re-created on the next startup (subscription plans, default pipeline,
scheduled jobs, compliance templates).
Returns:
A dict mapping table name number of rows deleted.
"""
from app.models import (
AuditLog,
BackupRecord,
DocumentMetadata,
FileProcessingStep,
FileRecord,
InAppNotification,
ProcessingLog,
SavedSearch,
SettingsAuditLog,
SharedLink,
UserImapAccount,
UserIntegration,
UserNotificationPreference,
UserNotificationTarget,
)
# Order matters: delete children before parents to respect FK constraints.
tables_to_wipe: list[tuple[str, type]] = [
("file_processing_steps", FileProcessingStep),
("processing_logs", ProcessingLog),
("shared_links", SharedLink),
("in_app_notifications", InAppNotification),
("user_notification_preferences", UserNotificationPreference),
("user_notification_targets", UserNotificationTarget),
("user_imap_accounts", UserImapAccount),
("user_integrations", UserIntegration),
("saved_searches", SavedSearch),
("settings_audit_log", SettingsAuditLog),
("audit_logs", AuditLog),
("backup_records", BackupRecord),
("document_metadata", DocumentMetadata),
("files", FileRecord),
]
result: dict[str, int] = {}
for table_name, model in tables_to_wipe:
try:
count = db.query(model).delete()
result[table_name] = count
logger.info("Wiped %d rows from %s", count, table_name)
except Exception:
logger.exception("Failed to wipe table %s during system reset", table_name)
db.rollback()
raise
db.commit()
return result
def perform_full_reset(db: Session) -> dict:
"""Perform a complete system reset: wipe database rows + work-files.
Args:
db: An active SQLAlchemy session.
Returns:
Summary dict with ``database`` and ``filesystem`` sub-dicts.
"""
logger.warning(">>> SYSTEM RESET: wiping all user data <<<")
db_result = _wipe_database(db)
fs_result = _wipe_workdir_data(settings.workdir)
logger.warning(">>> SYSTEM RESET complete <<<")
return {"database": db_result, "filesystem": fs_result}
def perform_reset_and_reimport(db: Session) -> dict:
"""Move original files to a reimport folder, wipe everything, then
configure the reimport folder as a watch folder for re-ingestion.
The watch-folder scanner (``scan_all_watch_folders``) will pick up
the files on its next periodic run and process them exactly as if
they had been freshly uploaded respecting the same backoff
strategy, size limits, and rate limits.
Args:
db: An active SQLAlchemy session.
Returns:
Summary dict with ``database``, ``filesystem``, and ``reimport`` sub-dicts.
"""
workdir_path = Path(settings.workdir)
reimport_dir = workdir_path / REIMPORT_FOLDER_NAME
original_dir = workdir_path / "original"
# 1. Collect original files
files_moved = 0
reimport_dir.mkdir(parents=True, exist_ok=True)
if original_dir.is_dir():
for entry in original_dir.iterdir():
if entry.is_file():
# Validate the resolved path stays within original_dir (path traversal guard)
try:
entry.resolve().relative_to(original_dir.resolve())
except ValueError:
logger.warning("Skipping file outside original dir: %s", entry)
continue
dest = reimport_dir / entry.name
# Avoid overwriting: append counter if name clash
if dest.exists():
stem = dest.stem
suffix = dest.suffix
counter = 1
while dest.exists():
dest = reimport_dir / f"{stem}_{counter}{suffix}"
counter += 1
shutil.copy2(str(entry), str(dest))
files_moved += 1
logger.info("Copied %d original files to reimport folder: %s", files_moved, reimport_dir)
# 2. Perform the full reset (wipe DB + other workdir data)
reset_result = perform_full_reset(db)
# 3. Ensure the reimport folder survived the wipe (it's not in _USER_DATA_SUBDIRS)
# and set up watch folder config to point at it.
_configure_reimport_watch_folder(str(reimport_dir))
reset_result["reimport"] = {
"files_moved": files_moved,
"reimport_folder": str(reimport_dir),
}
logger.warning(">>> SYSTEM RESET with re-import configured — %d files staged <<<", files_moved)
return reset_result
def _configure_reimport_watch_folder(reimport_path: str) -> None:
"""Append *reimport_path* to the application's watch-folder list.
The watch-folder scanner uses ``settings.watch_folders`` (a
comma-separated string). We mutate the runtime setting so the
next scan picks up the folder. We also set
``watch_folder_delete_after_process = True`` so files are cleaned
up after successful processing.
"""
current = getattr(settings, "watch_folders", None) or ""
folders = [f.strip() for f in current.split(",") if f.strip()]
if reimport_path not in folders:
folders.append(reimport_path)
# Mutate runtime settings (not persisted to .env — ephemeral)
object.__setattr__(settings, "watch_folders", ",".join(folders))
object.__setattr__(settings, "watch_folder_delete_after_process", True)
logger.info("Configured reimport watch folder: %s", reimport_path)
def perform_startup_reset() -> None:
"""Called during application startup when ``FACTORY_RESET_ON_STARTUP=True``.
Wipes database and filesystem data so the instance starts completely
fresh. Uses its own DB session so it runs before the normal lifespan
seeding logic.
"""
from app.database import SessionLocal
logger.warning("FACTORY_RESET_ON_STARTUP is enabled — wiping all data")
db = SessionLocal()
try:
perform_full_reset(db)
except Exception:
logger.exception("Factory reset on startup failed")
db.rollback()
finally:
db.close()
-21
View File
@@ -12,13 +12,11 @@ import smtplib
from email.mime.multipart import MIMEMultipart from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText from email.mime.text import MIMEText
from typing import Any from typing import Any
from urllib.parse import urlparse
import httpx import httpx
from app.database import SessionLocal from app.database import SessionLocal
from app.models import InAppNotification, UserNotificationPreference, UserNotificationTarget from app.models import InAppNotification, UserNotificationPreference, UserNotificationTarget
from app.utils.network import is_private_ip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -30,11 +28,6 @@ USER_EVENT_LABELS: dict[str, str] = {
EVENT_DOCUMENT_PROCESSED: "Document Processed", EVENT_DOCUMENT_PROCESSED: "Document Processed",
EVENT_DOCUMENT_FAILED: "Document Processing Failed", EVENT_DOCUMENT_FAILED: "Document Processing Failed",
} }
METADATA_ENDPOINTS = {
"169.254.169.254",
"169.254.169.253",
"metadata.google.internal",
}
def create_in_app_notification( def create_in_app_notification(
@@ -135,20 +128,6 @@ def _send_webhook_notification(target_config: dict[str, Any], event_type: str, t
logger.warning("Webhook notification target missing url") logger.warning("Webhook notification target missing url")
return False return False
parsed_url = urlparse(url)
if parsed_url.scheme not in {"http", "https"}:
logger.warning("Webhook notification to %s blocked: invalid scheme %s", url, parsed_url.scheme)
return False
hostname = parsed_url.hostname
if not hostname:
logger.warning("Webhook notification to %s blocked: missing hostname", url)
return False
if hostname in METADATA_ENDPOINTS or is_private_ip(hostname):
logger.warning("Webhook notification to %s blocked: private or metadata endpoint", url)
return False
payload = { payload = {
"event": event_type, "event": event_type,
"title": title, "title": title,
+13 -147
View File
@@ -11,47 +11,22 @@ import logging
from fastapi import Request from fastapi import Request
from sqlalchemy import or_ from sqlalchemy import or_
from sqlalchemy.orm import Query, Session from sqlalchemy.orm import Query
from sqlalchemy.sql import false from sqlalchemy.sql import false
from app.config import settings 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__) 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.
Priority: ``sub`` (OAuth subject) ``preferred_username`` ``email`` ``id``.
"""
return user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
def get_current_owner_id(request: Request) -> str | None: def get_current_owner_id(request: Request) -> str | None:
"""Extract the owner identifier for the current authenticated user. """Extract the owner identifier for the current authenticated user.
The owner ID is derived from the user's session data or, when no session The owner ID is derived from the user's session data. It uses the
is present, from a valid Bearer API token in the ``Authorization`` header. ``sub`` claim (OAuth subject) when available, falling back to
This ensures that both browser-based (session cookie) and mobile/API ``preferred_username`` or ``email``. Returns ``None`` when no user
(Bearer token) requests are correctly identified. is authenticated.
Priority for user resolution:
1. Session ``user`` dict (set by OAuth or local login).
2. ``request.state.api_token_user`` (set by ``require_login`` or an
earlier call to this function during the same request).
3. Direct Bearer token look-up against the database.
Within the resolved user dict the owner ID is chosen as:
``sub`` ``preferred_username`` ``email`` ``id``.
Args: Args:
request: The current FastAPI request with session data. request: The current FastAPI request with session data.
@@ -59,50 +34,19 @@ def get_current_owner_id(request: Request) -> str | None:
Returns: Returns:
A stable string identifier for the user, or ``None``. A stable string identifier for the user, or ``None``.
""" """
# 1. Session-based auth (most common for web UI)
user = request.session.get("user") user = request.session.get("user")
if user and isinstance(user, dict): if not user or not isinstance(user, dict):
return _owner_id_from_user(user) return None
# Prefer 'sub' (OAuth subject), then 'preferred_username', then 'email', then 'id'
# 2. Already-resolved API token user (cached by require_login or a return user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
# prior dependency call during this request)
api_user = getattr(request.state, "api_token_user", None)
if isinstance(api_user, dict):
return _owner_id_from_user(api_user)
# 3. Direct Bearer token resolution necessary when this function is
# invoked as a FastAPI dependency (via Depends) which runs *before*
# the @require_login decorator wrapper has had a chance to resolve
# the token and populate request.state.api_token_user.
auth_header = request.headers.get("authorization", "")
if isinstance(auth_header, str) and auth_header.startswith("Bearer "):
try:
from app.auth import _resolve_bearer_user
from app.database import SessionLocal
db = SessionLocal()
try:
resolved = _resolve_bearer_user(request, db)
finally:
db.close()
if resolved:
# Cache so subsequent calls (and require_login) skip the DB
request.state.api_token_user = resolved
return _owner_id_from_user(resolved)
except Exception:
logger.debug("Bearer token resolution failed in get_current_owner_id", exc_info=True)
return None
def apply_owner_filter(query: Query, request: Request) -> Query: def apply_owner_filter(query: Query, request: Request) -> Query:
"""Conditionally filter a ``FileRecord`` query by the current user. """Conditionally filter a ``FileRecord`` query by the current user.
When multi-user mode is enabled, only files whose ``owner_id`` When multi-user mode is enabled, only files whose ``owner_id``
matches the authenticated user are returned, **plus** any files that matches the authenticated user are returned. Admin users bypass
have been explicitly shared with the user via ``FileShare``. Admin the filter and see all documents.
users bypass the filter and see all documents.
When ``unowned_docs_visible_to_all`` is ``True`` (default), documents When ``unowned_docs_visible_to_all`` is ``True`` (default), documents
with ``owner_id IS NULL`` (unclaimed) are also included for every with ``owner_id IS NULL`` (unclaimed) are also included for every
@@ -130,89 +74,11 @@ def apply_owner_filter(query: Query, request: Request) -> Query:
# No authenticated user — return empty result set # No authenticated user — return empty result set
return query.filter(false()) 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] 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 # Optionally include unclaimed (owner_id IS NULL) documents
if settings.unowned_docs_visible_to_all: if settings.unowned_docs_visible_to_all:
conditions.append(FileRecord.owner_id.is_(None)) conditions.append(FileRecord.owner_id.is_(None))
return query.filter(or_(*conditions)) 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 -40
View File
@@ -18,13 +18,11 @@ import json
import logging import logging
import time import time
from typing import Any from typing import Any
from urllib.parse import urlparse
import requests import requests
from app.database import SessionLocal from app.database import SessionLocal
from app.models import WebhookConfig from app.models import WebhookConfig
from app.utils.network import is_private_ip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -42,11 +40,6 @@ VALID_EVENTS: frozenset[str] = frozenset(
#: Timeout (seconds) for outgoing webhook HTTP requests. #: Timeout (seconds) for outgoing webhook HTTP requests.
WEBHOOK_TIMEOUT = 10 WEBHOOK_TIMEOUT = 10
METADATA_ENDPOINTS = {
"169.254.169.254",
"169.254.169.253",
"metadata.google.internal",
}
def compute_signature(payload_bytes: bytes, secret: str) -> str: def compute_signature(payload_bytes: bytes, secret: str) -> str:
@@ -74,20 +67,6 @@ def deliver_webhook(url: str, payload: dict[str, Any], secret: str | None = None
Returns: Returns:
``True`` when the remote server responds with a 2xx status. ``True`` when the remote server responds with a 2xx status.
""" """
parsed_url = urlparse(url)
if parsed_url.scheme not in {"http", "https"}:
logger.warning("Webhook to %s blocked: invalid scheme %s", url, parsed_url.scheme)
return False
hostname = parsed_url.hostname
if not hostname:
logger.warning("Webhook to %s blocked: missing hostname", url)
return False
if hostname in METADATA_ENDPOINTS or is_private_ip(hostname):
logger.warning("Webhook to %s blocked: private or metadata endpoint", url)
return False
body = json.dumps(payload, default=str, sort_keys=True) body = json.dumps(payload, default=str, sort_keys=True)
body_bytes = body.encode("utf-8") body_bytes = body.encode("utf-8")
@@ -166,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 It delegates to :func:`deliver_webhook_task` (Celery) for each matching
webhook so delivery happens asynchronously with automatic retries. webhook so delivery happens asynchronously with automatic retries.
Also dispatches to automation hooks (Zapier / Make.com) if enabled.
Args: Args:
event: Event name (must be in :data:`VALID_EVENTS`). event: Event name (must be in :data:`VALID_EVENTS`).
data: Event-specific payload data. data: Event-specific payload data.
@@ -179,23 +156,16 @@ def dispatch_webhook_event(event: str, data: dict[str, Any]) -> None:
webhooks = get_active_webhooks_for_event(event) webhooks = get_active_webhooks_for_event(event)
if not webhooks: if not webhooks:
logger.debug("No active webhooks for event %s", event) logger.debug("No active webhooks for event %s", event)
else: return
payload = build_payload(event, data)
# Import here to avoid circular dependency with celery_app payload = build_payload(event, data)
from app.tasks.webhook_tasks import deliver_webhook_task
for wh in webhooks: # Import here to avoid circular dependency with celery_app
try: from app.tasks.webhook_tasks import deliver_webhook_task
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)
# Also fan-out to Zapier / Make.com automation hooks for wh in webhooks:
try: try:
from app.utils.automation_hooks import dispatch_automation_hooks deliver_webhook_task.delay(wh["url"], payload, wh["secret"])
logger.debug("Queued webhook delivery to %s for event %s", wh["url"], event)
dispatch_automation_hooks(event, data) except Exception as exc:
except Exception as exc: logger.error("Failed to queue webhook to %s: %s", wh["url"], exc)
logger.error("Failed to dispatch automation hooks for event %s: %s", event, exc)
-6
View File
@@ -10,7 +10,6 @@ from app.views.audit_logs import router as audit_logs_router
from app.views.backup import router as backup_router from app.views.backup import router as backup_router
from app.views.compliance import router as compliance_router from app.views.compliance import router as compliance_router
from app.views.db_wizard import router as db_wizard_router from app.views.db_wizard import router as db_wizard_router
from app.views.devices import router as devices_router # Mobile devices dashboard
from app.views.dropbox import router as dropbox_router from app.views.dropbox import router as dropbox_router
from app.views.filemanager import router as filemanager_router from app.views.filemanager import router as filemanager_router
@@ -27,7 +26,6 @@ from app.views.onedrive import router as onedrive_router
from app.views.pipelines import router as pipelines_router # Processing pipelines from app.views.pipelines import router as pipelines_router # Processing pipelines
from app.views.plans import router as plans_router # Admin Plan Designer from app.views.plans import router as plans_router # Admin Plan Designer
from app.views.profile import router as profile_router # User self-service profile from app.views.profile import router as profile_router # User self-service profile
from app.views.qr_login import router as qr_login_router # QR code mobile login
from app.views.queue import router as queue_router from app.views.queue import router as queue_router
from app.views.scheduled_jobs import router as scheduled_jobs_router # Scheduled batch jobs from app.views.scheduled_jobs import router as scheduled_jobs_router # Scheduled batch jobs
from app.views.search import router as search_router from app.views.search import router as search_router
@@ -36,7 +34,6 @@ from app.views.share import router as share_router
from app.views.shared_links import router as shared_links_router from app.views.shared_links import router as shared_links_router
from app.views.status import router as status_router from app.views.status import router as status_router
from app.views.subscriptions import router as subscriptions_router # Pricing + subscription pages from app.views.subscriptions import router as subscriptions_router # Pricing + subscription pages
from app.views.system_reset import router as system_reset_router # System reset / factory reset
from app.views.wizard import router as wizard_router from app.views.wizard import router as wizard_router
# Create a main router that includes all the view routers # Create a main router that includes all the view routers
@@ -63,7 +60,6 @@ router.include_router(plans_router) # Admin Plan Designer
router.include_router(onboarding_router) # User onboarding wizard router.include_router(onboarding_router) # User onboarding wizard
router.include_router(pipelines_router) # Processing pipelines router.include_router(pipelines_router) # Processing pipelines
router.include_router(profile_router) # User self-service profile settings router.include_router(profile_router) # User self-service profile settings
router.include_router(qr_login_router) # QR code mobile login page
router.include_router(imap_accounts_router) # Per-user IMAP ingestion accounts router.include_router(imap_accounts_router) # Per-user IMAP ingestion accounts
router.include_router(integrations_router) # Unified integrations dashboard router.include_router(integrations_router) # Unified integrations dashboard
router.include_router(notifications_router) # User notification dashboard router.include_router(notifications_router) # User notification dashboard
@@ -71,5 +67,3 @@ router.include_router(scheduled_jobs_router) # Admin scheduled batch jobs
router.include_router(audit_logs_router) # Comprehensive audit log viewer router.include_router(audit_logs_router) # Comprehensive audit log viewer
router.include_router(help_router) # Built-in help / How-To docs router.include_router(help_router) # Built-in help / How-To docs
router.include_router(compliance_router) # Compliance templates dashboard router.include_router(compliance_router) # Compliance templates dashboard
router.include_router(devices_router) # Mobile devices dashboard
router.include_router(system_reset_router) # System reset / factory reset
+5 -45
View File
@@ -94,22 +94,6 @@ def _inject_global_context(ctx: dict) -> None:
"allow_signup", "allow_signup",
getattr(settings, "multi_user_enabled", False) and getattr(settings, "allow_local_signup", False), getattr(settings, "multi_user_enabled", False) and getattr(settings, "allow_local_signup", False),
) )
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") req = ctx.get("request")
if req is not None: if req is not None:
@@ -162,36 +146,12 @@ def _inject_global_context(ctx: dict) -> None:
def template_response_with_version(*args, **kwargs): def template_response_with_version(*args, **kwargs):
"""Wrapper for TemplateResponse to include version and CSRF token in all templates. """Wrapper for TemplateResponse to include version and CSRF token in all templates"""
# If context dict is provided, add version to it
Handles both old-style and new-style Starlette TemplateResponse calls: if len(args) >= 2 and isinstance(args[1], dict):
- Old-style (Starlette <1.0): TemplateResponse(name, {"request": req, ...}, ...) _inject_global_context(args[1])
- New-style (Starlette 1.0+): TemplateResponse(request, name, context={...}, ...) elif "context" in kwargs and isinstance(kwargs["context"], dict):
"""
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"]) _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) return original_template_response(*args, **kwargs)
-25
View File
@@ -1,25 +0,0 @@
"""View route for the Devices management page.
Renders the ``devices.html`` template where users can see their registered
mobile devices, mobile API tokens (created via the mobile SSO flow or QR
code login), and revoke access per-device.
"""
import logging
from fastapi import APIRouter, Request
from app.views.base import require_login, templates
logger = logging.getLogger(__name__)
router = APIRouter()
@router.get("/devices", include_in_schema=False)
@require_login
async def devices_page(request: Request):
"""Render the Devices management page."""
return templates.TemplateResponse(
"devices.html",
{"request": request, "page_title": "Devices"},
)
+1 -27
View File
@@ -14,19 +14,6 @@ from app.views.base import APIRouter, Depends, get_db, require_login, settings,
router = APIRouter() 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") @router.get("/dropbox-setup")
@require_login @require_login
async def dropbox_setup_page( 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 path from the integration's existing config is pre-populated; global
admin credentials are never exposed in this mode. admin credentials are never exposed in this mode.
""" """
callback_url = _get_dropbox_callback_url(request)
if integration_id is not None: if integration_id is not None:
owner_id = get_current_owner_id(request) owner_id = get_current_owner_id(request)
integration = ( integration = (
@@ -61,12 +46,6 @@ async def dropbox_setup_page(
cfg = {} cfg = {}
# Support both "folder" (DROPBOX destination) and "folder_path" (WATCH_FOLDER source) # Support both "folder" (DROPBOX destination) and "folder_path" (WATCH_FOLDER source)
folder_path = cfg.get("folder", cfg.get("folder_path", "")) 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( return templates.TemplateResponse(
"dropbox.html", "dropbox.html",
{ {
@@ -77,12 +56,9 @@ async def dropbox_setup_page(
"integration_name": integration.name, "integration_name": integration.name,
"integration_type": integration.integration_type, "integration_type": integration.integration_type,
"folder_path": folder_path, "folder_path": folder_path,
# Only expose the public app key (not the secret) when global creds are allowed "app_key_value": "",
"app_key_value": settings.dropbox_app_key if global_creds_available else "",
"app_secret_value": "", "app_secret_value": "",
"refresh_token_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_id": integration_id,
"integration_name": None, "integration_name": None,
"integration_type": 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_key_value": "", # The callback will prioritize sessionStorage values
"app_secret_value": "", # The callback will prioritize sessionStorage values "app_secret_value": "", # The callback will prioritize sessionStorage values
"folder_path": "", # The callback will prioritize sessionStorage values "folder_path": "", # The callback will prioritize sessionStorage values
"callback_url": _get_dropbox_callback_url(request),
}, },
) )
+7 -228
View File
@@ -19,43 +19,6 @@ router = APIRouter()
_FILE_NOT_FOUND = "File not found" _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") @router.get("/files")
@require_login @require_login
def files_page( def files_page(
@@ -243,93 +206,10 @@ def files_page(
@router.get("/files/{file_id}") @router.get("/files/{file_id}")
@require_login @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)): 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: try:
import json 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) # Resolve the pipeline assigned to this file (explicit or system default)
pipeline_info = _resolve_pipeline(db, file_record) pipeline_info = _resolve_pipeline(db, file_record)
owner_ctx = _resolve_owner_context(request, file_record, db)
return templates.TemplateResponse( return templates.TemplateResponse(
"file_view.html", "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, "processed_file_exists": processed_file_exists,
"step_summary": step_summary, "step_summary": step_summary,
"pipeline_info": pipeline_info, "pipeline_info": pipeline_info,
**owner_ctx,
}, },
) )
except Exception as e: 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)}) 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 @require_login
def file_detail_page(request: Request, file_id: int, db: Session = Depends(get_db)): 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: try:
import json 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)}) 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 # Pipeline ↔ Celery-log stage mapping
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -589,7 +396,9 @@ _STEP_TYPE_TO_STAGES: dict[str, list[str]] = {
"embed_metadata": ["embed_metadata_into_pdf"], "embed_metadata": ["embed_metadata_into_pdf"],
"compute_embedding": ["compute_embedding"], "compute_embedding": ["compute_embedding"],
"send_to_destinations": ["finalize_document_storage", "send_to_all_destinations"], "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 # These internal bookkeeping stages are always shown in the flow regardless of
@@ -717,7 +526,6 @@ def _compute_processing_flow(logs, pipeline_steps=None):
"upload_to_ftp": "FTP Storage", "upload_to_ftp": "FTP Storage",
"upload_to_sftp": "SFTP Storage", "upload_to_sftp": "SFTP Storage",
"upload_to_email": "Email", "upload_to_email": "Email",
"upload_to_evernote": "Evernote",
"queue_dropbox": "Dropbox", "queue_dropbox": "Dropbox",
"queue_nextcloud": "Nextcloud", "queue_nextcloud": "Nextcloud",
"queue_paperless": "Paperless-ngx", "queue_paperless": "Paperless-ngx",
@@ -728,7 +536,6 @@ def _compute_processing_flow(logs, pipeline_steps=None):
"queue_ftp": "FTP Storage", "queue_ftp": "FTP Storage",
"queue_sftp": "SFTP Storage", "queue_sftp": "SFTP Storage",
"queue_email": "Email", "queue_email": "Email",
"queue_evernote": "Evernote",
} }
# Create a map of step names to their log entries # Create a map of step names to their log entries
@@ -1025,34 +832,6 @@ def get_processed_text(request: Request, file_id: int, db: Session = Depends(get
) )
@router.get("/files/{file_id}/text/default-language")
@require_login
def get_default_language_text(request: Request, file_id: int, db: Session = Depends(get_db)):
"""Return the persisted default-language translation for the file view."""
from fastapi import status
from fastapi.responses import JSONResponse
from app.models import FileRecord
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 file_record.default_language_text:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="No default-language translation available",
)
return JSONResponse(
content={
"text": file_record.default_language_text,
"language_code": file_record.default_language_code,
"detected_language": file_record.detected_language,
}
)
@router.get("/duplicates") @router.get("/duplicates")
@require_login @require_login
def duplicates_page( def duplicates_page(
+3 -13
View File
@@ -2,7 +2,7 @@
General routes for the application homepage and basic pages. General routes for the application homepage and basic pages.
""" """
from datetime import date, datetime, time, timedelta, timezone from datetime import date, datetime, timezone
from pathlib import Path from pathlib import Path
from fastapi import Depends, HTTPException, Request from fastapi import Depends, HTTPException, Request
@@ -64,27 +64,17 @@ async def serve_index(request: Request, db: Session = Depends(get_db)):
user = request.session.get("user") or {} user = request.session.get("user") or {}
is_admin = user.get("is_admin", False) is_admin = user.get("is_admin", False)
day_start = datetime.combine(today, time.min, tzinfo=timezone.utc)
day_end = day_start + timedelta(days=1)
month_start = day_start.replace(day=1)
if month_start.month == 12:
month_end = month_start.replace(year=month_start.year + 1, month=1)
else:
month_end = month_start.replace(month=month_start.month + 1)
try: try:
total_files: int = db.query(func.count(FileRecord.id)).scalar() or 0 total_files: int = db.query(func.count(FileRecord.id)).scalar() or 0
files_today: int = ( files_today: int = (
db.query(func.count(FileRecord.id)) db.query(func.count(FileRecord.id)).filter(func.date(FileRecord.created_at) == today).scalar() or 0
.filter(FileRecord.created_at >= day_start, FileRecord.created_at < day_end)
.scalar()
or 0
) )
files_month: int = ( files_month: int = (
db.query(func.count(FileRecord.id)) db.query(func.count(FileRecord.id))
.filter(FileRecord.created_at >= month_start, FileRecord.created_at < month_end) .filter(func.strftime("%Y-%m", FileRecord.created_at) == today.strftime("%Y-%m"))
.scalar() .scalar()
or 0 or 0
) )
+4 -11
View File
@@ -45,9 +45,6 @@ async def google_drive_setup_page(
except (json.JSONDecodeError, TypeError): except (json.JSONDecodeError, TypeError):
cfg = {} cfg = {}
folder_id = cfg.get("folder_id", "") 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( return templates.TemplateResponse(
"google_drive.html", "google_drive.html",
{ {
@@ -61,13 +58,10 @@ async def google_drive_setup_page(
"use_oauth": True, "use_oauth": True,
"oauth_configured": bool(integration.credentials), "oauth_configured": bool(integration.credentials),
"sa_configured": False, "sa_configured": False,
"has_system_credentials": has_system_credentials, "client_id": False,
"client_id": bool(settings.google_drive_client_id) if has_system_credentials else False, "client_id_value": "",
"client_id_value": (settings.google_drive_client_id or "" if has_system_credentials else ""), "client_secret": False,
"client_secret": bool(settings.google_drive_client_secret) if has_system_credentials else False, "client_secret_value": "",
"client_secret_value": (
settings.google_drive_client_secret or "" if has_system_credentials else ""
),
"refresh_token": False, "refresh_token": False,
"refresh_token_value": "", "refresh_token_value": "",
"has_credentials_json": False, "has_credentials_json": False,
@@ -96,7 +90,6 @@ async def google_drive_setup_page(
"use_oauth": use_oauth, "use_oauth": use_oauth,
"oauth_configured": oauth_configured, "oauth_configured": oauth_configured,
"sa_configured": sa_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": bool(settings.google_drive_client_id),
"client_id_value": settings.google_drive_client_id or "", "client_id_value": settings.google_drive_client_id or "",
"client_secret": bool(settings.google_drive_client_secret), "client_secret": bool(settings.google_drive_client_secret),
+5 -10
View File
@@ -44,9 +44,6 @@ async def onedrive_setup_page(
cfg = {} cfg = {}
# Support both "folder_path" (WATCH_FOLDER / ONEDRIVE destination) # Support both "folder_path" (WATCH_FOLDER / ONEDRIVE destination)
folder_path = cfg.get("folder_path", cfg.get("folder", "")) 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( return templates.TemplateResponse(
"onedrive.html", "onedrive.html",
{ {
@@ -57,12 +54,11 @@ async def onedrive_setup_page(
"integration_name": integration.name, "integration_name": integration.name,
"integration_type": integration.integration_type, "integration_type": integration.integration_type,
"folder_path": folder_path, "folder_path": folder_path,
"has_system_credentials": has_system_credentials, "client_id": False,
"client_id": bool(settings.onedrive_client_id) if has_system_credentials else False, "client_id_value": "",
"client_id_value": settings.onedrive_client_id or "" if has_system_credentials else "", "client_secret": False,
"client_secret": bool(settings.onedrive_client_secret) if has_system_credentials else False, "client_secret_value": "",
"client_secret_value": (settings.onedrive_client_secret or "" if has_system_credentials else ""), "tenant_id": "common",
"tenant_id": settings.onedrive_tenant_id or "common",
"refresh_token": False, "refresh_token": False,
"refresh_token_value": "", "refresh_token_value": "",
}, },
@@ -79,7 +75,6 @@ async def onedrive_setup_page(
"request": request, "request": request,
"user_mode": False, "user_mode": False,
"is_configured": is_configured, "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": bool(settings.onedrive_client_id),
"client_id_value": settings.onedrive_client_id or "", "client_id_value": settings.onedrive_client_id or "",
"client_secret": bool(settings.onedrive_client_secret), "client_secret": bool(settings.onedrive_client_secret),
-26
View File
@@ -1,26 +0,0 @@
"""View route for the QR code mobile login page.
Route:
GET /qr-login renders the QR login page (requires login)
"""
from __future__ import annotations
import logging
from fastapi import Request
from app.views.base import APIRouter, require_login, templates
logger = logging.getLogger(__name__)
router = APIRouter()
@router.get("/qr-login", include_in_schema=False)
@require_login
async def qr_login_page(request: Request):
"""Serve the QR code login page for mobile app authentication."""
return templates.TemplateResponse(
"qr_login.html",
{"request": request},
)
-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") @router.get("/admin/settings/audit-log")
@require_login @require_login
@require_admin_access @require_admin_access
+1 -2
View File
@@ -23,7 +23,6 @@ templates = Jinja2Templates(directory=str(_templates_dir))
async def shared_link_view(request: Request, token: str): async def shared_link_view(request: Request, token: str):
"""Render the public share landing page for a given token.""" """Render the public share landing page for a given token."""
return templates.TemplateResponse( return templates.TemplateResponse(
request,
"shared_link_view.html", "shared_link_view.html",
context={"token": token}, {"request": request, "token": token},
) )
-40
View File
@@ -1,40 +0,0 @@
"""
System reset view admin-only UI page.
Renders a confirmation-heavy page that allows administrators to:
1. **Full Reset** wipe all user data (DB + disk) for a fresh start.
2. **Reset & Re-import** move originals to a reimport folder, wipe,
and let the watch-folder mechanism re-ingest them.
Both options are gated behind the ``ENABLE_FACTORY_RESET`` feature flag.
"""
import logging
from fastapi import Depends, Request
from fastapi.responses import RedirectResponse, Response
from sqlalchemy.orm import Session
from app.config import settings
from app.views.base import APIRouter, get_db, require_login, templates
from app.views.settings import require_admin_access
logger = logging.getLogger(__name__)
router = APIRouter()
@router.get("/admin/system-reset")
@require_login
@require_admin_access
async def system_reset_page(request: Request, db: Session = Depends(get_db)) -> Response:
"""Render the system reset administration page."""
if not settings.enable_factory_reset:
return RedirectResponse(url="/settings", status_code=302)
return templates.TemplateResponse(
"system_reset.html",
{
"request": request,
"factory_reset_on_startup": settings.factory_reset_on_startup,
},
)
-101
View File
@@ -1,101 +0,0 @@
import time
import os
import sys
import asyncio
# Mock settings before app imports to bypass validation
os.environ["DATABASE_URL"] = "sqlite:///:memory:"
os.environ["REDIS_URL"] = "redis://localhost:6379"
os.environ["OPENAI_API_KEY"] = "mock_key"
os.environ["WORKDIR"] = "/tmp/workdir"
os.environ["AZURE_AI_KEY"] = "mock"
os.environ["AZURE_REGION"] = "mock"
os.environ["AZURE_ENDPOINT"] = "http://mock"
os.environ["GOTENBERG_URL"] = "http://mock"
os.environ["SESSION_SECRET"] = "mock_secret_mock_secret_mock_secret_mock_secret"
# Ensure app package is accessible
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.database import Base
from app.models import FileRecord
from app.api.duplicates import list_duplicate_groups
# Mocking Request object
class MockRequest:
def __init__(self):
self.session = {"user": {"username": "testuser"}}
self.state = type('State', (), {'user': {"username": "testuser"}})()
class MockURL:
def include_query_params(self, **kwargs):
return f"http://testserver/api/duplicates?page={kwargs.get('page')}"
url = MockURL()
def setup_db():
engine = create_engine('sqlite:///:memory:')
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine)
db = Session()
return db
def populate_data(db, num_groups, duplicates_per_group):
for i in range(num_groups):
filehash = f"hash_{i}"
# Original
original = FileRecord(
filehash=filehash,
local_filename=f"orig_{i}.txt",
file_size=100,
is_duplicate=False
)
db.add(original)
# Duplicates
for j in range(duplicates_per_group):
dup = FileRecord(
filehash=filehash,
local_filename=f"dup_{i}_{j}.txt",
file_size=100,
is_duplicate=True
)
db.add(dup)
db.commit()
async def run_benchmark(db):
request = MockRequest()
start_time = time.time()
# Run the function we want to benchmark
result = list_duplicate_groups(request=request, db=db, page=1, per_page=500)
if asyncio.iscoroutine(result):
result = await result
end_time = time.time()
return end_time - start_time, result
async def main():
db = setup_db()
# 500 groups, each with 20 duplicates = 10500 records total
print("Populating data...")
populate_data(db, 500, 20)
print("Data populated. Running baseline benchmark...")
# Warmup
result = list_duplicate_groups(request=MockRequest(), db=db, page=1, per_page=500)
if asyncio.iscoroutine(result):
await result
# Benchmark
total_time = 0
iterations = 10
for _ in range(iterations):
time_taken, _ = await run_benchmark(db)
total_time += time_taken
avg_time = total_time / iterations
print(f"Average time over {iterations} iterations: {avg_time:.4f} seconds")
if __name__ == "__main__":
asyncio.run(main())
-50
View File
@@ -1,50 +0,0 @@
import json
import time
import pytest
from app.database import get_db
from app.models import UserNotificationTarget, UserNotificationPreference
from app.main import app
from tests.test_notifications_api import _make_client, _OWNER, _cleanup
import statistics
def run_benchmark(notif_engine, notif_session, client, items_count, iterations=5):
# Setup
target = UserNotificationTarget(
owner_id=_OWNER,
channel_type="webhook",
name="My Webhook",
config=json.dumps({"url": "https://x.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
# Generate big payload
preferences = []
for i in range(items_count):
preferences.append({
"event_type": f"event.type.{i}",
"channel_type": "webhook",
"is_enabled": True,
"target_id": target.id,
})
payload = {"preferences": preferences}
# Warm up
client.put("/api/user-notifications/preferences", json=payload)
times = []
for _ in range(iterations):
# Alter the values a bit so it's a real update
for p in payload["preferences"]:
p["is_enabled"] = not p["is_enabled"]
start = time.time()
resp = client.put("/api/user-notifications/preferences", json=payload)
end = time.time()
assert resp.status_code == 200
times.append(end - start)
return statistics.mean(times)

Some files were not shown because too many files have changed in this diff Show More