Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 10df0b21bf |
@@ -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
|
|
||||||
*~
|
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -1 +1 @@
|
|||||||
2026-06-01T03:41:15Z
|
2026-03-15T21:39:26Z
|
||||||
|
|||||||
-1768
File diff suppressed because it is too large
Load Diff
+18
-53
@@ -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
@@ -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
|
||||||
|
|
||||||
|
|||||||
+69
-115
@@ -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 |
|
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
||||||
==============================
|
==============================
|
||||||
|
|||||||
@@ -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
@@ -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)
|
|
||||||
|
|||||||
@@ -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
@@ -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})
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -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)
|
|
||||||
@@ -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
|
|
||||||
]
|
|
||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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,
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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.",
|
|
||||||
}
|
|
||||||
@@ -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):
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
|
||||||
}
|
|
||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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"),)
|
|
||||||
|
|||||||
@@ -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")
|
|
||||||
@@ -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
|
|
||||||
@@ -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,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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 "
|
||||||
|
|||||||
@@ -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}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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),
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
@@ -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,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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",
|
|
||||||
}
|
}
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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)
|
|
||||||
@@ -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),
|
|
||||||
)
|
|
||||||
@@ -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",
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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
@@ -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
@@ -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))
|
|
||||||
|
|||||||
@@ -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}"
|
|
||||||
@@ -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.0–1.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.0–1.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.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."
|
|
||||||
),
|
|
||||||
"type": "float",
|
|
||||||
"sensitive": False,
|
|
||||||
"required": False,
|
|
||||||
"restart_required": True,
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
|
||||||
@@ -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
@@ -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
@@ -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)
|
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
@@ -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),
|
||||||
|
|||||||
@@ -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},
|
|
||||||
)
|
|
||||||
@@ -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
@@ -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},
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
@@ -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())
|
|
||||||
@@ -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
Reference in New Issue
Block a user