mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-10-01 22:00:31 +00:00
Compare commits
45
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
da3de299fc | ||
|
|
3a022e2dc4 | ||
|
|
58bda48682 | ||
|
|
b2a2e15ead | ||
|
|
e865185425 | ||
|
|
d02c19a09e | ||
|
|
4a106454ce | ||
|
|
4254100fae | ||
|
|
74f727b1f3 | ||
|
|
fb6fad952e | ||
|
|
f9b4e0c8e0 | ||
|
|
e5041097af | ||
|
|
61f36eb253 | ||
|
|
e50bcc2efb | ||
|
|
88000aab45 | ||
|
|
01b4ab392c | ||
|
|
52aa8758e3 | ||
|
|
4b7cc1c0c3 | ||
|
|
f02d7455be | ||
|
|
a61c35fcd8 | ||
|
|
64552d5d5b | ||
|
|
a59edf3e5b | ||
|
|
6a5d06ee26 | ||
|
|
126ec414a8 | ||
|
|
773744d04b | ||
|
|
7b98085b31 | ||
|
|
b3bf220f56 | ||
|
|
d4429c3cc7 | ||
|
|
4f777d7438 | ||
|
|
2a44d8b5ba | ||
|
|
abf5050ea7 | ||
|
|
091ddf7c45 | ||
|
|
c9f7f2cfbe | ||
|
|
b457610ffb | ||
|
|
969c2ea0e2 | ||
|
|
31b806a285 | ||
|
|
99851b418c | ||
|
|
e34eda07bb | ||
|
|
793459b416 | ||
|
|
04297fd02c | ||
|
|
1b277dd8e1 | ||
|
|
3c20abeb4c | ||
|
|
b11f1f8459 | ||
|
|
7424e7ce0b | ||
|
|
a53a3d3769 |
@@ -59,7 +59,6 @@ updates:
|
|||||||
- "drf-*"
|
- "drf-*"
|
||||||
- "djangorestframework"
|
- "djangorestframework"
|
||||||
- "whitenoise"
|
- "whitenoise"
|
||||||
- "bleach"
|
|
||||||
- "jinja2"
|
- "jinja2"
|
||||||
# Async, Task Queuing & Caching
|
# Async, Task Queuing & Caching
|
||||||
async-tasks:
|
async-tasks:
|
||||||
|
|||||||
@@ -102,7 +102,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: "${{ matrix.python-version }}"
|
python-version: "${{ matrix.python-version }}"
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
@@ -139,13 +139,13 @@ jobs:
|
|||||||
pytest
|
pytest
|
||||||
- name: Upload test results to Codecov
|
- name: Upload test results to Codecov
|
||||||
if: always()
|
if: always()
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0
|
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
|
||||||
with:
|
with:
|
||||||
flags: backend-python-${{ matrix.python-version }}
|
flags: backend-python-${{ matrix.python-version }}
|
||||||
files: junit.xml
|
files: junit.xml
|
||||||
report_type: test_results
|
report_type: test_results
|
||||||
- name: Upload coverage to Codecov
|
- name: Upload coverage to Codecov
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0
|
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
|
||||||
with:
|
with:
|
||||||
flags: backend-python-${{ matrix.python-version }}
|
flags: backend-python-${{ matrix.python-version }}
|
||||||
files: coverage.xml
|
files: coverage.xml
|
||||||
@@ -176,7 +176,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: "${{ env.DEFAULT_PYTHON }}"
|
python-version: "${{ env.DEFAULT_PYTHON }}"
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
|
|||||||
@@ -106,7 +106,7 @@ jobs:
|
|||||||
echo "repository=${repo_name}"
|
echo "repository=${repo_name}"
|
||||||
echo "name=${repo_name}" >> $GITHUB_OUTPUT
|
echo "name=${repo_name}" >> $GITHUB_OUTPUT
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@37fe631027851001ddb9b187196cc803df7f5f0e # v4.3.0
|
uses: docker/setup-buildx-action@f87e5991a6d7451dcb8d9637bfbc97413f497069 # v4.4.1
|
||||||
- name: Login to GitHub Container Registry
|
- name: Login to GitHub Container Registry
|
||||||
uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
|
uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
|
||||||
with:
|
with:
|
||||||
@@ -132,7 +132,7 @@ jobs:
|
|||||||
type=semver,pattern={{major}}.{{minor}}
|
type=semver,pattern={{major}}.{{minor}}
|
||||||
- name: Build and push by digest
|
- name: Build and push by digest
|
||||||
id: build
|
id: build
|
||||||
uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0
|
uses: docker/build-push-action@c3c9e263c25d99ce0380d002d59b67737d91b0dc # v7.4.0
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./Dockerfile
|
file: ./Dockerfile
|
||||||
@@ -182,7 +182,7 @@ jobs:
|
|||||||
echo "Downloaded digests:"
|
echo "Downloaded digests:"
|
||||||
ls -la /tmp/digests/
|
ls -la /tmp/digests/
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@37fe631027851001ddb9b187196cc803df7f5f0e # v4.3.0
|
uses: docker/setup-buildx-action@f87e5991a6d7451dcb8d9637bfbc97413f497069 # v4.4.1
|
||||||
- name: Login to GitHub Container Registry
|
- name: Login to GitHub Container Registry
|
||||||
uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
|
uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
|
||||||
with:
|
with:
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
|
|||||||
@@ -174,13 +174,13 @@ jobs:
|
|||||||
run: cd src-ui && pnpm run test --max-workers=2 --shard=${{ matrix.shard-index }}/${{ matrix.shard-count }}
|
run: cd src-ui && pnpm run test --max-workers=2 --shard=${{ matrix.shard-index }}/${{ matrix.shard-count }}
|
||||||
- name: Upload test results to Codecov
|
- name: Upload test results to Codecov
|
||||||
if: always()
|
if: always()
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0
|
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
|
||||||
with:
|
with:
|
||||||
flags: frontend-node-${{ matrix.node-version }}
|
flags: frontend-node-${{ matrix.node-version }}
|
||||||
directory: src-ui/
|
directory: src-ui/
|
||||||
report_type: test_results
|
report_type: test_results
|
||||||
- name: Upload coverage to Codecov
|
- name: Upload coverage to Codecov
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0
|
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
|
||||||
with:
|
with:
|
||||||
flags: frontend-node-${{ matrix.node-version }}
|
flags: frontend-node-${{ matrix.node-version }}
|
||||||
directory: src-ui/coverage/
|
directory: src-ui/coverage/
|
||||||
@@ -216,7 +216,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: '3.12'
|
python-version: '3.12'
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: '0.12.x'
|
version: '0.12.x'
|
||||||
enable-cache: false
|
enable-cache: false
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: false
|
enable-cache: false
|
||||||
@@ -212,7 +212,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
python-version: ${{ env.DEFAULT_PYTHON_VERSION }}
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: false
|
enable-cache: false
|
||||||
|
|||||||
@@ -44,7 +44,7 @@ jobs:
|
|||||||
- name: Run Semgrep
|
- name: Run Semgrep
|
||||||
run: semgrep scan --config auto --sarif-output results.sarif
|
run: semgrep scan --config auto --sarif-output results.sarif
|
||||||
- name: Upload results to GitHub code scanning
|
- name: Upload results to GitHub code scanning
|
||||||
uses: github/codeql-action/upload-sarif@b96794f015dfd88f77b49b1c93e0fa7110f94c63 # v4.38.0
|
uses: github/codeql-action/upload-sarif@1c5b675653bb5c22dbe9b12b556ec555138e09fd # v4.38.1
|
||||||
if: always()
|
if: always()
|
||||||
with:
|
with:
|
||||||
sarif_file: results.sarif
|
sarif_file: results.sarif
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- name: Clean temporary images
|
- name: Clean temporary images
|
||||||
if: "${{ env.TOKEN != '' }}"
|
if: "${{ env.TOKEN != '' }}"
|
||||||
uses: stumpylog/image-cleaner-action/ephemeral@4fe057d991d63b8f6d5d22c40f17c1bca2226537 # v0.12.0
|
uses: stumpylog/image-cleaner-action/ephemeral@21f875bab2376314e0525c614e923c669434f0e4 # v0.13.0
|
||||||
with:
|
with:
|
||||||
token: "${{ env.TOKEN }}"
|
token: "${{ env.TOKEN }}"
|
||||||
owner: "${{ github.repository_owner }}"
|
owner: "${{ github.repository_owner }}"
|
||||||
@@ -56,7 +56,7 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- name: Clean untagged images
|
- name: Clean untagged images
|
||||||
if: "${{ env.TOKEN != '' }}"
|
if: "${{ env.TOKEN != '' }}"
|
||||||
uses: stumpylog/image-cleaner-action/untagged@4fe057d991d63b8f6d5d22c40f17c1bca2226537 # v0.12.0
|
uses: stumpylog/image-cleaner-action/untagged@21f875bab2376314e0525c614e923c669434f0e4 # v0.13.0
|
||||||
with:
|
with:
|
||||||
token: "${{ env.TOKEN }}"
|
token: "${{ env.TOKEN }}"
|
||||||
owner: "${{ github.repository_owner }}"
|
owner: "${{ github.repository_owner }}"
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ jobs:
|
|||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
# Initializes the CodeQL tools for scanning.
|
# Initializes the CodeQL tools for scanning.
|
||||||
- name: Initialize CodeQL
|
- name: Initialize CodeQL
|
||||||
uses: github/codeql-action/init@b96794f015dfd88f77b49b1c93e0fa7110f94c63 # v4.38.0
|
uses: github/codeql-action/init@1c5b675653bb5c22dbe9b12b556ec555138e09fd # v4.38.1
|
||||||
with:
|
with:
|
||||||
languages: ${{ matrix.language }}
|
languages: ${{ matrix.language }}
|
||||||
# If you wish to specify custom queries, you can do so here or in a config file.
|
# If you wish to specify custom queries, you can do so here or in a config file.
|
||||||
@@ -47,4 +47,4 @@ jobs:
|
|||||||
# Prefix the list here with "+" to use these queries and those in the config file.
|
# Prefix the list here with "+" to use these queries and those in the config file.
|
||||||
# queries: ./path/to/local/query, your-org/your-repo/queries@main
|
# queries: ./path/to/local/query, your-org/your-repo/queries@main
|
||||||
- name: Perform CodeQL Analysis
|
- name: Perform CodeQL Analysis
|
||||||
uses: github/codeql-action/analyze@b96794f015dfd88f77b49b1c93e0fa7110f94c63 # v4.38.0
|
uses: github/codeql-action/analyze@1c5b675653bb5c22dbe9b12b556ec555138e09fd # v4.38.1
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ jobs:
|
|||||||
token: ${{ secrets.PNGX_BOT_PAT }}
|
token: ${{ secrets.PNGX_BOT_PAT }}
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
- name: crowdin action
|
- name: crowdin action
|
||||||
uses: crowdin/github-action@0d5670f539973aea2f01abce61a8989934df0025 # v3.0.2
|
uses: crowdin/github-action@df474cdfb9f41d6ae777118749477c2cdf7cacc8 # v3.2.0
|
||||||
with:
|
with:
|
||||||
upload_translations: false
|
upload_translations: false
|
||||||
download_translations: true
|
download_translations: true
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ jobs:
|
|||||||
sudo apt-get update -qq
|
sudo apt-get update -qq
|
||||||
sudo apt-get install -qq --no-install-recommends gettext
|
sudo apt-get install -qq --no-install-recommends gettext
|
||||||
- name: Install uv
|
- name: Install uv
|
||||||
uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0
|
uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0
|
||||||
with:
|
with:
|
||||||
version: ${{ env.DEFAULT_UV_VERSION }}
|
version: ${{ env.DEFAULT_UV_VERSION }}
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
|
|||||||
+706
-456
File diff suppressed because it is too large
Load Diff
@@ -38,7 +38,7 @@ repos:
|
|||||||
- json
|
- json
|
||||||
# See https://github.com/prettier/prettier/issues/15742 for the fork reason
|
# See https://github.com/prettier/prettier/issues/15742 for the fork reason
|
||||||
- repo: https://github.com/rbubley/mirrors-prettier
|
- repo: https://github.com/rbubley/mirrors-prettier
|
||||||
rev: 'v3.9.6'
|
rev: 'v3.9.8'
|
||||||
hooks:
|
hooks:
|
||||||
- id: prettier
|
- id: prettier
|
||||||
types_or:
|
types_or:
|
||||||
@@ -46,11 +46,11 @@ repos:
|
|||||||
- ts
|
- ts
|
||||||
- markdown
|
- markdown
|
||||||
additional_dependencies:
|
additional_dependencies:
|
||||||
- prettier@3.9.6
|
- prettier@3.9.8
|
||||||
- 'prettier-plugin-organize-imports@4.3.0'
|
- 'prettier-plugin-organize-imports@4.3.0'
|
||||||
# Python hooks
|
# Python hooks
|
||||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||||
rev: v0.16.7
|
rev: v0.16.8
|
||||||
hooks:
|
hooks:
|
||||||
- id: ruff-check
|
- id: ruff-check
|
||||||
- id: ruff-format
|
- id: ruff-format
|
||||||
|
|||||||
+4804
-12941
File diff suppressed because one or more lines are too long
+4
-2
@@ -30,7 +30,7 @@ RUN set -eux \
|
|||||||
# Purpose: Installs s6-overlay and rootfs
|
# Purpose: Installs s6-overlay and rootfs
|
||||||
# Comments:
|
# Comments:
|
||||||
# - Don't leave anything extra in here either
|
# - Don't leave anything extra in here either
|
||||||
FROM ghcr.io/astral-sh/uv:0.12.16-python3.14-trixie-slim AS s6-overlay-base
|
FROM ghcr.io/astral-sh/uv:0.12.20-python3.14-trixie-slim AS s6-overlay-base
|
||||||
|
|
||||||
WORKDIR /usr/src/s6
|
WORKDIR /usr/src/s6
|
||||||
|
|
||||||
@@ -171,7 +171,9 @@ RUN set -eux \
|
|||||||
&& cp /etc/ImageMagick-6/paperless-policy.xml /etc/ImageMagick-6/policy.xml \
|
&& cp /etc/ImageMagick-6/paperless-policy.xml /etc/ImageMagick-6/policy.xml \
|
||||||
&& echo "Cleaning up image layer" \
|
&& echo "Cleaning up image layer" \
|
||||||
&& rm --force --verbose *.deb \
|
&& rm --force --verbose *.deb \
|
||||||
&& rm --recursive --force --verbose /var/lib/apt/lists/*
|
&& rm --recursive --force --verbose /var/lib/apt/lists/* \
|
||||||
|
&& echo "Configuring interactive shells to source the s6 container environment" \
|
||||||
|
&& echo '. /etc/profile.d/contenv.sh' >> /etc/bash.bashrc
|
||||||
|
|
||||||
WORKDIR /usr/src/paperless/src/
|
WORKDIR /usr/src/paperless/src/
|
||||||
|
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ services:
|
|||||||
network_mode: host
|
network_mode: host
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
greenmail:
|
greenmail:
|
||||||
image: docker.io/greenmail/standalone:2.1.13
|
image: docker.io/greenmail/standalone:2.1.14
|
||||||
hostname: greenmail
|
hostname: greenmail
|
||||||
container_name: greenmail
|
container_name: greenmail
|
||||||
environment:
|
environment:
|
||||||
|
|||||||
Executable
+18
@@ -0,0 +1,18 @@
|
|||||||
|
#!/bin/sh
|
||||||
|
# Source s6 container environment for interactive shells.
|
||||||
|
# Ensures variables resolved from *_FILE secret injection are visible
|
||||||
|
# when using 'docker exec bash'. Does not affect s6 services (those
|
||||||
|
# use with-contenv directly). Has no effect in non-container contexts
|
||||||
|
# because the directory will not exist.
|
||||||
|
# Note: sh/dash shells opened via 'docker exec sh' are not covered;
|
||||||
|
# only bash-based sessions benefit from this file.
|
||||||
|
_pngx_contenv="/run/s6/container_environment"
|
||||||
|
if [ -d "${_pngx_contenv}" ]; then
|
||||||
|
for _pngx_f in "${_pngx_contenv}"/*; do
|
||||||
|
[ -f "${_pngx_f}" ] || continue
|
||||||
|
_pngx_name=$(basename "${_pngx_f}")
|
||||||
|
_pngx_val=$(cat "${_pngx_f}")
|
||||||
|
export "${_pngx_name}=${_pngx_val}"
|
||||||
|
done
|
||||||
|
fi
|
||||||
|
unset _pngx_contenv _pngx_f _pngx_name _pngx_val
|
||||||
+14
-9
@@ -136,13 +136,15 @@ for suggested generation and embedding models.
|
|||||||
### AI-assisted suggestions
|
### AI-assisted suggestions
|
||||||
|
|
||||||
With AI enabled, Paperless-ngx can suggest a title, tags, correspondent, document type,
|
With AI enabled, Paperless-ngx can suggest a title, tags, correspondent, document type,
|
||||||
storage path and dates by sending the document to the LLM. This is **opt-in per request**
|
storage path and dates by sending the document to the LLM using "Suggest" button on the document
|
||||||
and surfaces through the "Suggest" control on the document detail page, alongside the
|
detail page. You can choose which type of suggestions are requested by default under Settings >
|
||||||
classic classifier-based suggestions — it does not disable them. Suggestions are requested
|
Documents, either ML (classifier-based) suggestions, AI suggestions, or both. When both are requested
|
||||||
automatically when you open a document that carries an inbox tag unless "Automatically request
|
the results are combined.
|
||||||
suggestions for inbox documents" under Settings > Documents is disabled. Suggestion output
|
|
||||||
language can be steered with
|
Suggestions are requested automatically when you open a document that carries an inbox tag
|
||||||
[`PAPERLESS_AI_LLM_OUTPUT_LANGUAGE`](configuration.md#PAPERLESS_AI_LLM_OUTPUT_LANGUAGE)
|
unless "Automatically request suggestions for inbox documents" under Settings > Documents is disabled.
|
||||||
|
|
||||||
|
Suggestion output language can be steered with [`PAPERLESS_AI_LLM_OUTPUT_LANGUAGE`](configuration.md#PAPERLESS_AI_LLM_OUTPUT_LANGUAGE)
|
||||||
(otherwise it follows the user's UI language).
|
(otherwise it follows the user's UI language).
|
||||||
|
|
||||||
### The LLM index (RAG) and similar documents
|
### The LLM index (RAG) and similar documents
|
||||||
@@ -153,8 +155,11 @@ in similar existing documents, and the document chat can retrieve relevant conte
|
|||||||
|
|
||||||
Enable it by setting
|
Enable it by setting
|
||||||
[`PAPERLESS_AI_LLM_EMBEDDING_BACKEND`](configuration.md#PAPERLESS_AI_LLM_EMBEDDING_BACKEND)
|
[`PAPERLESS_AI_LLM_EMBEDDING_BACKEND`](configuration.md#PAPERLESS_AI_LLM_EMBEDDING_BACKEND)
|
||||||
(`huggingface` for fully-local embeddings, or `ollama` / `openai-like`). The index is only
|
(`huggingface` for fully-local embeddings, or `ollama` / `openai-like`). By default, the main
|
||||||
built when AI is enabled **and** an embedding backend is set.
|
LLM API key and endpoint are used, but an optional embedding-specific[API key](configuration.md#PAPERLESS_AI_LLM_EMBEDDING_API_KEY)
|
||||||
|
and [endpoint](configuration.md#PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT) can be configured.
|
||||||
|
|
||||||
|
The index is only built when AI is enabled **and** an embedding backend is set.
|
||||||
|
|
||||||
The index is updated automatically on a schedule controlled by
|
The index is updated automatically on a schedule controlled by
|
||||||
[`PAPERLESS_LLM_INDEX_TASK_CRON`](configuration.md#PAPERLESS_LLM_INDEX_TASK_CRON) (daily by
|
[`PAPERLESS_LLM_INDEX_TASK_CRON`](configuration.md#PAPERLESS_LLM_INDEX_TASK_CRON) (daily by
|
||||||
|
|||||||
+21
-6
@@ -1576,9 +1576,6 @@ ports.
|
|||||||
#### [`PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS=<bool>`](#PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS) {#PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS}
|
#### [`PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS=<bool>`](#PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS) {#PAPERLESS_WEBHOOKS_ALLOW_INTERNAL_REQUESTS}
|
||||||
|
|
||||||
: If set to false, webhooks cannot be sent to internal URLs (e.g., localhost).
|
: If set to false, webhooks cannot be sent to internal URLs (e.g., localhost).
|
||||||
A hostname is blocked if any of the addresses it resolves to is non-public.
|
|
||||||
Webhook requests connect directly, without using the `HTTP_PROXY` or
|
|
||||||
`HTTPS_PROXY` environment variables, and never follow redirects.
|
|
||||||
|
|
||||||
Defaults to true, which allows internal requests.
|
Defaults to true, which allows internal requests.
|
||||||
|
|
||||||
@@ -1587,7 +1584,7 @@ Webhook requests connect directly, without using the `HTTP_PROXY` or
|
|||||||
#### [`PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS=<bool>`](#PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS) {#PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS}
|
#### [`PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS=<bool>`](#PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS) {#PAPERLESS_EMAIL_ALLOW_INTERNAL_HOSTS}
|
||||||
|
|
||||||
: If set to false, incoming mail account connections are blocked when the
|
: If set to false, incoming mail account connections are blocked when the
|
||||||
configured IMAP hostname resolves to any non-public address (for example,
|
configured IMAP hostname resolves to a non-public address (for example,
|
||||||
localhost, link-local, or RFC1918 private ranges).
|
localhost, link-local, or RFC1918 private ranges).
|
||||||
|
|
||||||
Defaults to true, which allows internal hosts.
|
Defaults to true, which allows internal hosts.
|
||||||
@@ -2136,6 +2133,13 @@ for language and resource considerations.
|
|||||||
|
|
||||||
Defaults to None.
|
Defaults to None.
|
||||||
|
|
||||||
|
#### [`PAPERLESS_AI_LLM_EMBEDDING_API_KEY=<str>`](#PAPERLESS_AI_LLM_EMBEDDING_API_KEY) {#PAPERLESS_AI_LLM_EMBEDDING_API_KEY}
|
||||||
|
|
||||||
|
: The API key to use for the embedding backend. If not supplied, embeddings use
|
||||||
|
`PAPERLESS_AI_LLM_API_KEY`.
|
||||||
|
|
||||||
|
Defaults to None.
|
||||||
|
|
||||||
#### [`PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT=<str>`](#PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT) {#PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT}
|
#### [`PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT=<str>`](#PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT) {#PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT}
|
||||||
|
|
||||||
: The endpoint / url to use for the embedding backend. If not supplied, embeddings use
|
: The endpoint / url to use for the embedding backend. If not supplied, embeddings use
|
||||||
@@ -2217,11 +2221,22 @@ used with the OpenAI-compatible backend to target a custom provider or local gat
|
|||||||
#### [`PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS=<bool>`](#PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS) {#PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS}
|
#### [`PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS=<bool>`](#PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS) {#PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS}
|
||||||
|
|
||||||
: If set to false, Paperless blocks AI endpoint URLs that resolve to non-public addresses (e.g., localhost, etc).
|
: If set to false, Paperless blocks AI endpoint URLs that resolve to non-public addresses (e.g., localhost, etc).
|
||||||
A hostname is blocked if any of the addresses it resolves to is non-public, and redirects are checked the same way.
|
|
||||||
Requests to a configured AI endpoint connect directly, without using the `HTTP_PROXY` or `HTTPS_PROXY` environment variables.
|
|
||||||
|
|
||||||
Defaults to true, which allows internal endpoints.
|
Defaults to true, which allows internal endpoints.
|
||||||
|
|
||||||
|
#### [`PAPERLESS_AI_LLM_EXTRA_PARAMS=<json>`](#PAPERLESS_AI_LLM_EXTRA_PARAMS) {#PAPERLESS_AI_LLM_EXTRA_PARAMS}
|
||||||
|
|
||||||
|
: A JSON object of extra parameters sent with every LLM request, for providers that require a parameter Paperless does not
|
||||||
|
set itself. Values here override Paperless' own, and no validation is performed. Whatever you put here is passed to the
|
||||||
|
backend as-is, so an invalid parameter will simply be rejected by your provider. For example, current OpenAI reasoning
|
||||||
|
models refuse tool calls on the chat completions API unless reasoning is off:
|
||||||
|
|
||||||
|
```
|
||||||
|
PAPERLESS_AI_LLM_EXTRA_PARAMS={"reasoning_effort": "none"}
|
||||||
|
```
|
||||||
|
|
||||||
|
Defaults to empty, which adds nothing to requests.
|
||||||
|
|
||||||
#### [`PAPERLESS_LLM_INDEX_TASK_CRON=<cron expression>`](#PAPERLESS_LLM_INDEX_TASK_CRON) {#PAPERLESS_LLM_INDEX_TASK_CRON}
|
#### [`PAPERLESS_LLM_INDEX_TASK_CRON=<cron expression>`](#PAPERLESS_LLM_INDEX_TASK_CRON) {#PAPERLESS_LLM_INDEX_TASK_CRON}
|
||||||
|
|
||||||
: Configures the schedule to update the AI embeddings of text content and metadata for all documents. Only performed if
|
: Configures the schedule to update the AI embeddings of text content and metadata for all documents. Only performed if
|
||||||
|
|||||||
+1
-1
@@ -613,7 +613,7 @@ The following workflow action types are available:
|
|||||||
- The request headers as key-value pairs
|
- The request headers as key-value pairs
|
||||||
|
|
||||||
For security reasons, webhooks can be limited to specific ports and disallowed from connecting to local URLs. See the relevant
|
For security reasons, webhooks can be limited to specific ports and disallowed from connecting to local URLs. See the relevant
|
||||||
[configuration settings](configuration.md#workflow-webhooks) to change this behavior. Webhook requests connect directly (proxy environment variables are not used) and do not follow redirects. If you are allowing non-admins to create workflows,
|
[configuration settings](configuration.md#workflow-webhooks) to change this behavior. If you are allowing non-admins to create workflows,
|
||||||
you may want to adjust these settings to prevent abuse.
|
you may want to adjust these settings to prevent abuse.
|
||||||
|
|
||||||
##### Move to Trash {#workflow-action-move-to-trash}
|
##### Move to Trash {#workflow-action-move-to-trash}
|
||||||
|
|||||||
+7
-10
@@ -17,10 +17,8 @@ classifiers = [
|
|||||||
# TODO: Move certain things to groups and then utilize that further
|
# TODO: Move certain things to groups and then utilize that further
|
||||||
# This will allow testing to not install a webserver, mysql, etc
|
# This will allow testing to not install a webserver, mysql, etc
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"anyio>=4.12",
|
|
||||||
"azure-ai-documentintelligence>=1.0.2",
|
"azure-ai-documentintelligence>=1.0.2",
|
||||||
"babel>=2.17",
|
"babel>=2.17",
|
||||||
"bleach~=6.4.0",
|
|
||||||
"celery[redis]~=5.6.2",
|
"celery[redis]~=5.6.2",
|
||||||
"channels~=4.2",
|
"channels~=4.2",
|
||||||
"channels-redis~=4.2",
|
"channels-redis~=4.2",
|
||||||
@@ -36,7 +34,7 @@ dependencies = [
|
|||||||
"django-cors-headers~=4.9.0",
|
"django-cors-headers~=4.9.0",
|
||||||
"django-extensions~=4.1",
|
"django-extensions~=4.1",
|
||||||
"django-filter~=25.1",
|
"django-filter~=25.1",
|
||||||
"django-guardian>=3.3.3,<3.5",
|
"django-guardian>=3.3.3,<3.6",
|
||||||
"django-multiselectfield~=1.0.1",
|
"django-multiselectfield~=1.0.1",
|
||||||
"django-rich~=2.2.0",
|
"django-rich~=2.2.0",
|
||||||
"django-soft-delete~=1.0.18",
|
"django-soft-delete~=1.0.18",
|
||||||
@@ -46,10 +44,8 @@ dependencies = [
|
|||||||
"drf-spectacular-sidecar>=2026.7.1,<2026.10",
|
"drf-spectacular-sidecar>=2026.7.1,<2026.10",
|
||||||
"drf-writable-nested~=0.7.1",
|
"drf-writable-nested~=0.7.1",
|
||||||
"filelock~=3.32.0",
|
"filelock~=3.32.0",
|
||||||
"flower>=2.0.1,<2.2",
|
"flower>=2.0.1,<2.3",
|
||||||
"gotenberg-client[httpx]~=1.0",
|
"gotenberg-client[httpx]~=1.0",
|
||||||
"httpcore~=1.0.9",
|
|
||||||
"httpx~=0.28.1",
|
|
||||||
"httpx-oauth~=0.17",
|
"httpx-oauth~=0.17",
|
||||||
"ijson>=3.5.1",
|
"ijson>=3.5.1",
|
||||||
"imap-tools>=1.14,<1.16",
|
"imap-tools>=1.14,<1.16",
|
||||||
@@ -80,6 +76,7 @@ dependencies = [
|
|||||||
"tantivy~=0.26.0",
|
"tantivy~=0.26.0",
|
||||||
"tika-client[httpx]~=1.0",
|
"tika-client[httpx]~=1.0",
|
||||||
"torch>=2.13,<2.15",
|
"torch>=2.13,<2.15",
|
||||||
|
"turbohtml~=1.10.0",
|
||||||
"watchfiles>=1.2",
|
"watchfiles>=1.2",
|
||||||
"whitenoise~=6.11",
|
"whitenoise~=6.11",
|
||||||
"whoosh-compat[tantivy]==0.3",
|
"whoosh-compat[tantivy]==0.3",
|
||||||
@@ -87,13 +84,13 @@ dependencies = [
|
|||||||
]
|
]
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
mariadb = [
|
mariadb = [
|
||||||
"mysqlclient~=2.2.7",
|
"mysqlclient>=2.2.7,<2.4",
|
||||||
]
|
]
|
||||||
postgres = [
|
postgres = [
|
||||||
"psycopg[c,pool]==3.3.4",
|
"psycopg[c,pool]==3.3.4",
|
||||||
# Direct dependency for proper resolution of the pre-built wheels
|
# Direct dependency for proper resolution of the pre-built wheels
|
||||||
"psycopg-c==3.3.4",
|
"psycopg-c==3.3.4",
|
||||||
"psycopg-pool==3.3.1",
|
"psycopg-pool==3.3.2",
|
||||||
]
|
]
|
||||||
webserver = [
|
webserver = [
|
||||||
"granian[uvloop]>=2.7,<2.9",
|
"granian[uvloop]>=2.7,<2.9",
|
||||||
@@ -115,7 +112,7 @@ lint = [
|
|||||||
testing = [
|
testing = [
|
||||||
"daphne",
|
"daphne",
|
||||||
"factory-boy~=3.3.1",
|
"factory-boy~=3.3.1",
|
||||||
"faker>=40.36,<40.39",
|
"faker>=40.36,<40.40",
|
||||||
"imagehash",
|
"imagehash",
|
||||||
"pytest~=9.1.1",
|
"pytest~=9.1.1",
|
||||||
"pytest-cov~=7.1.0",
|
"pytest-cov~=7.1.0",
|
||||||
@@ -139,7 +136,6 @@ typing = [
|
|||||||
"mypy",
|
"mypy",
|
||||||
"mypy-baseline",
|
"mypy-baseline",
|
||||||
"pyrefly",
|
"pyrefly",
|
||||||
"types-bleach",
|
|
||||||
"types-channels",
|
"types-channels",
|
||||||
"types-colorama",
|
"types-colorama",
|
||||||
"types-dateparser",
|
"types-dateparser",
|
||||||
@@ -154,6 +150,7 @@ typing = [
|
|||||||
|
|
||||||
[tool.uv]
|
[tool.uv]
|
||||||
required-version = ">=0.9.0"
|
required-version = ">=0.9.0"
|
||||||
|
prerelease = "disallow"
|
||||||
environments = [
|
environments = [
|
||||||
"sys_platform == 'darwin'",
|
"sys_platform == 'darwin'",
|
||||||
"sys_platform == 'linux'",
|
"sys_platform == 'linux'",
|
||||||
|
|||||||
+286
-209
File diff suppressed because it is too large
Load Diff
+15
-15
@@ -15,16 +15,16 @@
|
|||||||
},
|
},
|
||||||
"private": true,
|
"private": true,
|
||||||
"dependencies": {
|
"dependencies": {
|
||||||
"@angular/cdk": "^22.1.6",
|
"@angular/cdk": "^22.1.7",
|
||||||
"@angular/common": "~22.1.6",
|
"@angular/common": "~22.1.7",
|
||||||
"@angular/compiler": "~22.1.6",
|
"@angular/compiler": "~22.1.7",
|
||||||
"@angular/core": "~22.1.6",
|
"@angular/core": "~22.1.7",
|
||||||
"@angular/forms": "~22.1.6",
|
"@angular/forms": "~22.1.7",
|
||||||
"@angular/localize": "~22.1.6",
|
"@angular/localize": "~22.1.7",
|
||||||
"@angular/platform-browser": "~22.1.6",
|
"@angular/platform-browser": "~22.1.7",
|
||||||
"@angular/router": "~22.1.6",
|
"@angular/router": "~22.1.7",
|
||||||
"@ng-bootstrap/ng-bootstrap": "^21.0.0",
|
"@ng-bootstrap/ng-bootstrap": "^21.0.0",
|
||||||
"@ng-select/ng-select": "~24.1.1",
|
"@ng-select/ng-select": "~24.1.2",
|
||||||
"@ngneat/dirty-check-forms": "^3.0.3",
|
"@ngneat/dirty-check-forms": "^3.0.3",
|
||||||
"@popperjs/core": "^2.11.8",
|
"@popperjs/core": "^2.11.8",
|
||||||
"bootstrap": "^5.3.8",
|
"bootstrap": "^5.3.8",
|
||||||
@@ -54,20 +54,20 @@
|
|||||||
"@angular-eslint/template-parser": "22.5.0",
|
"@angular-eslint/template-parser": "22.5.0",
|
||||||
"@angular/build": "22.1.8",
|
"@angular/build": "22.1.8",
|
||||||
"@angular/cli": "22.1.8",
|
"@angular/cli": "22.1.8",
|
||||||
"@angular/compiler-cli": "~22.1.6",
|
"@angular/compiler-cli": "~22.1.7",
|
||||||
"@playwright/test": "^1.62.1",
|
"@playwright/test": "^1.62.1",
|
||||||
"@types/jest": "^30.0.0",
|
"@types/jest": "^30.0.0",
|
||||||
"@types/node": "^26.5.0",
|
"@types/node": "^26.6.2",
|
||||||
"@typescript-eslint/eslint-plugin": "^8.70.0",
|
"@typescript-eslint/eslint-plugin": "^8.70.0",
|
||||||
"@typescript-eslint/parser": "^8.70.0",
|
"@typescript-eslint/parser": "^8.70.0",
|
||||||
"@typescript-eslint/utils": "^8.70.0",
|
"@typescript-eslint/utils": "^8.70.0",
|
||||||
"eslint": "^10.10.0",
|
"eslint": "^10.11.0",
|
||||||
"jest": "30.5.1",
|
"jest": "30.5.2",
|
||||||
"jest-environment-jsdom": "^30.5.1",
|
"jest-environment-jsdom": "^30.5.2",
|
||||||
"jest-junit": "^17.0.0",
|
"jest-junit": "^17.0.0",
|
||||||
"jest-preset-angular": "^17.0.0",
|
"jest-preset-angular": "^17.0.0",
|
||||||
"jest-websocket-mock": "^2.5.0",
|
"jest-websocket-mock": "^2.5.0",
|
||||||
"prettier": "^3.9.6",
|
"prettier": "^3.9.8",
|
||||||
"prettier-plugin-organize-imports": "^4.3.0",
|
"prettier-plugin-organize-imports": "^4.3.0",
|
||||||
"ts-node": "~10.9.2",
|
"ts-node": "~10.9.2",
|
||||||
"typescript": "^6.0.3"
|
"typescript": "^6.0.3"
|
||||||
|
|||||||
Generated
+531
-518
File diff suppressed because it is too large
Load Diff
@@ -253,6 +253,24 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
@if (aiEnabled) {
|
||||||
|
<div class="row mb-3">
|
||||||
|
<div class="col-md-3 col-form-label pt-0">
|
||||||
|
<span i18n>Suggestions default to</span>
|
||||||
|
</div>
|
||||||
|
<div class="col">
|
||||||
|
<fieldset class="btn-group btn-group-sm">
|
||||||
|
<input type="radio" class="btn-check" id="suggestionSourceBoth" [value]="SuggestionSource.Both" formControlName="documentEditingSuggestionSource">
|
||||||
|
<label class="btn btn-outline-primary" for="suggestionSourceBoth"><ng-container i18n>Both</ng-container></label>
|
||||||
|
<input type="radio" class="btn-check" id="suggestionSourceML" [value]="SuggestionSource.ML" formControlName="documentEditingSuggestionSource">
|
||||||
|
<label class="btn btn-outline-primary" for="suggestionSourceML"><i-bs class="me-1" name="cpu"></i-bs><ng-container i18n>ML only</ng-container></label>
|
||||||
|
<input type="radio" class="btn-check" id="suggestionSourceAI" [value]="SuggestionSource.AI" formControlName="documentEditingSuggestionSource">
|
||||||
|
<label class="btn btn-outline-primary" for="suggestionSourceAI"><i-bs class="me-1" name="stars"></i-bs><ng-container i18n>AI only</ng-container></label>
|
||||||
|
</fieldset>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
}
|
||||||
|
|
||||||
<div class="row">
|
<div class="row">
|
||||||
<div class="col">
|
<div class="col">
|
||||||
<pngx-input-check i18n-title title="Automatically request suggestions for inbox documents" i18n-hint hint="If un-checked, suggestions must be requested via the Suggest button." formControlName="documentEditingAutoSuggest"></pngx-input-check>
|
<pngx-input-check i18n-title title="Automatically request suggestions for inbox documents" i18n-hint hint="If un-checked, suggestions must be requested via the Suggest button." formControlName="documentEditingAutoSuggest"></pngx-input-check>
|
||||||
|
|||||||
@@ -307,7 +307,7 @@ describe('SettingsComponent', () => {
|
|||||||
expect(toastErrorSpy).toHaveBeenCalled()
|
expect(toastErrorSpy).toHaveBeenCalled()
|
||||||
expect(storeSpy).toHaveBeenCalled()
|
expect(storeSpy).toHaveBeenCalled()
|
||||||
expect(appearanceSettingsSpy).not.toHaveBeenCalled()
|
expect(appearanceSettingsSpy).not.toHaveBeenCalled()
|
||||||
expect(setSpy).toHaveBeenCalledTimes(34)
|
expect(setSpy).toHaveBeenCalledTimes(35)
|
||||||
expect(setSpy).toHaveBeenCalledWith(SETTINGS_KEYS.SIDEBAR_HIDDEN_ITEMS, [
|
expect(setSpy).toHaveBeenCalledWith(SETTINGS_KEYS.SIDEBAR_HIDDEN_ITEMS, [
|
||||||
HideableSidebarItemID.Workflows,
|
HideableSidebarItemID.Workflows,
|
||||||
])
|
])
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ import {
|
|||||||
HIDEABLE_SIDEBAR_ITEM_IDS,
|
HIDEABLE_SIDEBAR_ITEM_IDS,
|
||||||
HideableSidebarItemID,
|
HideableSidebarItemID,
|
||||||
SETTINGS_KEYS,
|
SETTINGS_KEYS,
|
||||||
|
SuggestionSource,
|
||||||
} from 'src/app/data/ui-settings'
|
} from 'src/app/data/ui-settings'
|
||||||
import { User } from 'src/app/data/user'
|
import { User } from 'src/app/data/user'
|
||||||
import { IfPermissionsDirective } from 'src/app/directives/if-permissions.directive'
|
import { IfPermissionsDirective } from 'src/app/directives/if-permissions.directive'
|
||||||
@@ -184,6 +185,7 @@ export class SettingsComponent
|
|||||||
documentEditingRemoveInboxTags: new FormControl(null),
|
documentEditingRemoveInboxTags: new FormControl(null),
|
||||||
documentEditingOverlayThumbnail: new FormControl(null),
|
documentEditingOverlayThumbnail: new FormControl(null),
|
||||||
documentEditingAutoSuggest: new FormControl(null),
|
documentEditingAutoSuggest: new FormControl(null),
|
||||||
|
documentEditingSuggestionSource: new FormControl(null),
|
||||||
documentDetailsHiddenFields: new FormControl([]),
|
documentDetailsHiddenFields: new FormControl([]),
|
||||||
searchDbOnly: new FormControl(null),
|
searchDbOnly: new FormControl(null),
|
||||||
searchLink: new FormControl(null),
|
searchLink: new FormControl(null),
|
||||||
@@ -217,6 +219,11 @@ export class SettingsComponent
|
|||||||
public readonly PdfZoomScale = PdfZoomScale
|
public readonly PdfZoomScale = PdfZoomScale
|
||||||
|
|
||||||
public readonly PdfEditorEditMode = PdfEditorEditMode
|
public readonly PdfEditorEditMode = PdfEditorEditMode
|
||||||
|
public readonly SuggestionSource = SuggestionSource
|
||||||
|
|
||||||
|
get aiEnabled(): boolean {
|
||||||
|
return this.settings.get(SETTINGS_KEYS.AI_ENABLED)
|
||||||
|
}
|
||||||
|
|
||||||
public readonly documentDetailFieldOptions = documentDetailFieldOptions
|
public readonly documentDetailFieldOptions = documentDetailFieldOptions
|
||||||
public readonly sidebarItemOptions = HIDEABLE_SIDEBAR_ITEM_IDS.map((id) => ({
|
public readonly sidebarItemOptions = HIDEABLE_SIDEBAR_ITEM_IDS.map((id) => ({
|
||||||
@@ -404,6 +411,9 @@ export class SettingsComponent
|
|||||||
documentEditingAutoSuggest: this.settings.get(
|
documentEditingAutoSuggest: this.settings.get(
|
||||||
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST
|
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST
|
||||||
),
|
),
|
||||||
|
documentEditingSuggestionSource: this.settings.get(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE
|
||||||
|
),
|
||||||
documentDetailsHiddenFields: this.settings.get(
|
documentDetailsHiddenFields: this.settings.get(
|
||||||
SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS
|
SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS
|
||||||
),
|
),
|
||||||
@@ -625,6 +635,10 @@ export class SettingsComponent
|
|||||||
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST,
|
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST,
|
||||||
this.settingsForm.value.documentEditingAutoSuggest
|
this.settingsForm.value.documentEditingAutoSuggest
|
||||||
)
|
)
|
||||||
|
this.settings.set(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE,
|
||||||
|
this.settingsForm.value.documentEditingSuggestionSource
|
||||||
|
)
|
||||||
this.settings.set(
|
this.settings.set(
|
||||||
SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS,
|
SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS,
|
||||||
this.settingsForm.value.documentDetailsHiddenFields
|
this.settingsForm.value.documentDetailsHiddenFields
|
||||||
|
|||||||
@@ -154,11 +154,29 @@
|
|||||||
& section {
|
& section {
|
||||||
position: absolute;
|
position: absolute;
|
||||||
text-align: initial;
|
text-align: initial;
|
||||||
|
pointer-events: auto;
|
||||||
box-sizing: border-box;
|
box-sizing: border-box;
|
||||||
transform-origin: 0 0;
|
transform-origin: 0 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
& :is(.linkAnnotation, .buttonWidgetAnnotation.pushButton) > a {
|
||||||
|
position: absolute;
|
||||||
|
inset: 0;
|
||||||
|
font-size: 1em;
|
||||||
|
transition: none;
|
||||||
|
}
|
||||||
|
|
||||||
|
& :is(.linkAnnotation, .buttonWidgetAnnotation.pushButton):not(.hasBorder)
|
||||||
|
> a:hover {
|
||||||
|
opacity: 0.2;
|
||||||
|
background-color: rgb(255 255 0);
|
||||||
|
}
|
||||||
|
|
||||||
& .annotationTextContent {
|
& .annotationTextContent {
|
||||||
opacity: 0;
|
opacity: 0;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
:host ::ng-deep .textLayer.selecting ~ .annotationLayer section {
|
||||||
|
pointer-events: none;
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,7 +1,11 @@
|
|||||||
import { SimpleChange } from '@angular/core'
|
import { SimpleChange } from '@angular/core'
|
||||||
import { ComponentFixture, TestBed } from '@angular/core/testing'
|
import { ComponentFixture, TestBed } from '@angular/core/testing'
|
||||||
import * as pdfjs from 'pdfjs-dist/legacy/build/pdf.mjs'
|
import * as pdfjs from 'pdfjs-dist/legacy/build/pdf.mjs'
|
||||||
import { PDFSinglePageViewer, PDFViewer } from 'pdfjs-dist/web/pdf_viewer.mjs'
|
import {
|
||||||
|
LinkTarget,
|
||||||
|
PDFSinglePageViewer,
|
||||||
|
PDFViewer,
|
||||||
|
} from 'pdfjs-dist/web/pdf_viewer.mjs'
|
||||||
import { PngxPdfViewerComponent } from './pdf-viewer.component'
|
import { PngxPdfViewerComponent } from './pdf-viewer.component'
|
||||||
import { PdfRenderMode, PdfZoomLevel, PdfZoomScale } from './pdf-viewer.types'
|
import { PdfRenderMode, PdfZoomLevel, PdfZoomScale } from './pdf-viewer.types'
|
||||||
|
|
||||||
@@ -58,6 +62,16 @@ describe('PngxPdfViewerComponent', () => {
|
|||||||
expect((component as any).pdfViewer).toBeInstanceOf(PDFViewer)
|
expect((component as any).pdfViewer).toBeInstanceOf(PDFViewer)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('opens external links in a new tab', () => {
|
||||||
|
const linkService = (component as any).linkService
|
||||||
|
expect(linkService.options).toEqual(
|
||||||
|
expect.objectContaining({
|
||||||
|
externalLinkTarget: LinkTarget.BLANK,
|
||||||
|
externalLinkRel: 'noopener noreferrer nofollow',
|
||||||
|
})
|
||||||
|
)
|
||||||
|
})
|
||||||
|
|
||||||
it('resolves the worker source relative to the document base URI', async () => {
|
it('resolves the worker source relative to the document base URI', async () => {
|
||||||
setBaseHref('/paperless/')
|
setBaseHref('/paperless/')
|
||||||
const getDocumentSpy = jest.spyOn(pdfjs, 'getDocument')
|
const getDocumentSpy = jest.spyOn(pdfjs, 'getDocument')
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ import {
|
|||||||
} from 'pdfjs-dist/legacy/build/pdf.mjs'
|
} from 'pdfjs-dist/legacy/build/pdf.mjs'
|
||||||
import {
|
import {
|
||||||
EventBus,
|
EventBus,
|
||||||
|
LinkTarget,
|
||||||
PDFFindController,
|
PDFFindController,
|
||||||
PDFLinkService,
|
PDFLinkService,
|
||||||
PDFSinglePageViewer,
|
PDFSinglePageViewer,
|
||||||
@@ -75,7 +76,11 @@ export class PngxPdfViewerComponent
|
|||||||
private lastViewerPage?: number
|
private lastViewerPage?: number
|
||||||
|
|
||||||
private readonly eventBus = new EventBus()
|
private readonly eventBus = new EventBus()
|
||||||
private readonly linkService = new PDFLinkService({ eventBus: this.eventBus })
|
private readonly linkService = new PDFLinkService({
|
||||||
|
eventBus: this.eventBus,
|
||||||
|
externalLinkTarget: LinkTarget.BLANK,
|
||||||
|
externalLinkRel: 'noopener noreferrer nofollow',
|
||||||
|
})
|
||||||
private readonly findController = new PDFFindController({
|
private readonly findController = new PDFFindController({
|
||||||
eventBus: this.eventBus,
|
eventBus: this.eventBus,
|
||||||
linkService: this.linkService,
|
linkService: this.linkService,
|
||||||
|
|||||||
+77
-51
@@ -1,58 +1,84 @@
|
|||||||
<div class="btn-group">
|
<div class="d-flex align-items-center">
|
||||||
<button type="button" class="btn btn-sm btn-outline-primary" (click)="clickSuggest()" [disabled]="disabled() || loading() || (suggestions() && !aiEnabled())" [aria-label]="noSuggestions ? 'No suggestions' : 'Suggest'" i18n-aria-label>
|
<div class="btn-group">
|
||||||
@if (loading()) {
|
<button type="button" class="btn btn-sm btn-outline-primary" (click)="clickSuggest()" [disabled]="disabled() || loading() || (suggestions() && !aiEnabled())" [aria-label]="noSuggestions ? 'No suggestions' : 'Suggest'" i18n-aria-label>
|
||||||
<div class="spinner-border spinner-border-sm" role="status"></div>
|
@if (loading()) {
|
||||||
} @else if (noSuggestions) {
|
<div class="spinner-border spinner-border-sm" role="status"></div>
|
||||||
<i-bs width="1.2em" height="1.2em" name="check-circle"></i-bs>
|
} @else if (noSuggestions) {
|
||||||
} @else {
|
<i-bs width="1.2em" height="1.2em" name="check-circle"></i-bs>
|
||||||
<i-bs width="1.2em" height="1.2em" name="stars"></i-bs>
|
} @else {
|
||||||
|
<i-bs width="1.2em" height="1.2em" name="lightbulb"></i-bs>
|
||||||
|
}
|
||||||
|
@if (noSuggestions) {
|
||||||
|
<span class="d-none d-lg-inline ps-1" i18n>No suggestions</span>
|
||||||
|
} @else {
|
||||||
|
<span class="d-none d-lg-inline ps-1" i18n>Suggest</span>
|
||||||
|
}
|
||||||
|
@if (totalSuggestions > 0) {
|
||||||
|
<span class="badge bg-primary ms-2">{{ totalSuggestions }}</span>
|
||||||
|
}
|
||||||
|
</button>
|
||||||
|
|
||||||
|
@if (aiEnabled()) {
|
||||||
|
<div class="btn-group" ngbDropdown #dropdown="ngbDropdown" [popperOptions]="popperOptions">
|
||||||
|
<button type="button" class="btn btn-sm btn-outline-primary" ngbDropdownToggle [disabled]="disabled() || loading() || !suggestions()" aria-expanded="false" aria-controls="suggestionsDropdown" aria-label="Suggestions dropdown">
|
||||||
|
<span class="visually-hidden" i18n>Show suggestions</span>
|
||||||
|
</button>
|
||||||
|
|
||||||
|
<div ngbDropdownMenu aria-labelledby="suggestionsDropdown" class="shadow suggestions-dropdown">
|
||||||
|
<div class="list-group list-group-flush small pb-0">
|
||||||
|
@if (novelSuggestions === 0 && reusableSuggestions === 0) {
|
||||||
|
<div class="list-group-item text-muted fst-italic">
|
||||||
|
<small class="text-muted small fst-italic" i18n>No novel suggestions</small>
|
||||||
|
</div>
|
||||||
|
}
|
||||||
|
@if (suggestions()?.suggested_tags?.length > 0) {
|
||||||
|
<small class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="tags"></i-bs><ng-container i18n>Tags</ng-container></small>
|
||||||
|
@for (tag of suggestions().suggested_tags; track tag) {
|
||||||
|
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addTag.emit(tag)">{{ tag }}</button>
|
||||||
|
}
|
||||||
|
}
|
||||||
|
@if (suggestions()?.suggested_document_types?.length > 0) {
|
||||||
|
<div class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="hash"></i-bs><ng-container i18n>Document Types</ng-container></div>
|
||||||
|
@for (type of suggestions().suggested_document_types; track type) {
|
||||||
|
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addDocumentType.emit(type)">{{ type }}</button>
|
||||||
|
}
|
||||||
|
}
|
||||||
|
@if (suggestions()?.suggested_correspondents?.length > 0) {
|
||||||
|
<div class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="person"></i-bs><ng-container i18n>Correspondents</ng-container></div>
|
||||||
|
@for (correspondent of suggestions().suggested_correspondents; track correspondent) {
|
||||||
|
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addCorrespondent.emit(correspondent)">{{ correspondent }}</button>
|
||||||
|
}
|
||||||
|
}
|
||||||
|
@if (reusableSuggestions > 0) {
|
||||||
|
<div class="list-group-item text-muted fst-italic">
|
||||||
|
<small class="text-muted small fst-italic" i18n>{reusableSuggestions, plural, =1 {1 existing value suggested below} other {{{reusableSuggestions}} existing values suggested below}}</small>
|
||||||
|
</div>
|
||||||
|
}
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
}
|
}
|
||||||
@if (noSuggestions) {
|
</div>
|
||||||
<span class="d-none d-lg-inline ps-1" i18n>No suggestions</span>
|
|
||||||
} @else {
|
|
||||||
<span class="d-none d-lg-inline ps-1" i18n>Suggest</span>
|
|
||||||
}
|
|
||||||
@if (totalSuggestions > 0) {
|
|
||||||
<span class="badge bg-primary ms-2">{{ totalSuggestions }}</span>
|
|
||||||
}
|
|
||||||
</button>
|
|
||||||
|
|
||||||
@if (aiEnabled()) {
|
@if (aiEnabled()) {
|
||||||
<div class="btn-group" ngbDropdown #dropdown="ngbDropdown" [popperOptions]="popperOptions">
|
<div ngbDropdown autoClose="outside" placement="bottom-end" [popperOptions]="popperOptions">
|
||||||
<button type="button" class="btn btn-sm btn-outline-primary" ngbDropdownToggle [disabled]="disabled() || loading() || !suggestions()" aria-expanded="false" aria-controls="suggestionsDropdown" aria-label="Suggestions dropdown">
|
<button type="button" class="btn btn-sm btn-link position-relative" ngbDropdownToggle [disabled]="disabled() || loading()" i18n-title title="Suggestion options">
|
||||||
<span class="visually-hidden" i18n>Show suggestions</span>
|
<i-bs name="three-dots"></i-bs>
|
||||||
|
@if (source() !== defaultSource()) {
|
||||||
|
<span class="position-absolute top-0 start-100 translate-middle p-1 bg-primary border border-light rounded-circle">
|
||||||
|
<span class="visually-hidden" i18n>Not using default</span>
|
||||||
|
</span>
|
||||||
|
}
|
||||||
</button>
|
</button>
|
||||||
|
<div ngbDropdownMenu class="shadow p-3">
|
||||||
<div ngbDropdownMenu aria-labelledby="suggestionsDropdown" class="shadow suggestions-dropdown">
|
<div class="small text-muted mb-2" i18n>Suggest using:</div>
|
||||||
<div class="list-group list-group-flush small pb-0">
|
<div class="form-check small">
|
||||||
@if (novelSuggestions === 0 && reusableSuggestions === 0) {
|
<input class="form-check-input" type="checkbox" id="suggestionSourceML" [checked]="useML" [disabled]="useML && !useAI" (change)="setSources($event.target.checked, useAI)">
|
||||||
<div class="list-group-item text-muted fst-italic">
|
<label class="form-check-label d-inline-flex align-items-center gap-1" for="suggestionSourceML"><i-bs name="cpu"></i-bs><ng-container i18n>ML</ng-container></label>
|
||||||
<small class="text-muted small fst-italic" i18n>No novel suggestions</small>
|
</div>
|
||||||
</div>
|
<div class="form-check small">
|
||||||
}
|
<input class="form-check-input" type="checkbox" id="suggestionSourceAI" [checked]="useAI" [disabled]="useAI && !useML" (change)="setSources(useML, $event.target.checked)">
|
||||||
@if (suggestions()?.suggested_tags?.length > 0) {
|
<label class="form-check-label d-inline-flex align-items-center gap-1" for="suggestionSourceAI"><i-bs name="stars"></i-bs><ng-container i18n>AI</ng-container></label>
|
||||||
<small class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="tags"></i-bs><ng-container i18n>Tags</ng-container></small>
|
|
||||||
@for (tag of suggestions().suggested_tags; track tag) {
|
|
||||||
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addTag.emit(tag)">{{ tag }}</button>
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@if (suggestions()?.suggested_document_types?.length > 0) {
|
|
||||||
<div class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="hash"></i-bs><ng-container i18n>Document Types</ng-container></div>
|
|
||||||
@for (type of suggestions().suggested_document_types; track type) {
|
|
||||||
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addDocumentType.emit(type)">{{ type }}</button>
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@if (suggestions()?.suggested_correspondents?.length > 0) {
|
|
||||||
<div class="list-group-item text-uppercase text-muted small"><i-bs class="me-2" name="person"></i-bs><ng-container i18n>Correspondents</ng-container></div>
|
|
||||||
@for (correspondent of suggestions().suggested_correspondents; track correspondent) {
|
|
||||||
<button type="button" class="list-group-item list-group-item-action bg-light" (click)="addCorrespondent.emit(correspondent)">{{ correspondent }}</button>
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@if (reusableSuggestions > 0) {
|
|
||||||
<div class="list-group-item text-muted fst-italic">
|
|
||||||
<small class="text-muted small fst-italic" i18n>{reusableSuggestions, plural, =1 {1 existing value suggested below} other {{{reusableSuggestions}} existing values suggested below}}</small>
|
|
||||||
</div>
|
|
||||||
}
|
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|||||||
+4
@@ -1,3 +1,7 @@
|
|||||||
.suggestions-dropdown {
|
.suggestions-dropdown {
|
||||||
min-width: 250px;
|
min-width: 250px;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.btn-link.dropdown-toggle::after {
|
||||||
|
display: none;
|
||||||
|
}
|
||||||
|
|||||||
+59
-1
@@ -1,6 +1,7 @@
|
|||||||
import { ComponentFixture, TestBed } from '@angular/core/testing'
|
import { ComponentFixture, TestBed } from '@angular/core/testing'
|
||||||
import { NgbDropdownModule } from '@ng-bootstrap/ng-bootstrap'
|
import { NgbDropdownModule } from '@ng-bootstrap/ng-bootstrap'
|
||||||
import { NgxBootstrapIconsModule, allIcons } from 'ngx-bootstrap-icons'
|
import { NgxBootstrapIconsModule, allIcons } from 'ngx-bootstrap-icons'
|
||||||
|
import { SuggestionSource } from 'src/app/data/ui-settings'
|
||||||
import { SuggestionsDropdownComponent } from './suggestions-dropdown.component'
|
import { SuggestionsDropdownComponent } from './suggestions-dropdown.component'
|
||||||
|
|
||||||
describe('SuggestionsDropdownComponent', () => {
|
describe('SuggestionsDropdownComponent', () => {
|
||||||
@@ -179,14 +180,71 @@ describe('SuggestionsDropdownComponent', () => {
|
|||||||
|
|
||||||
it('should toggle dropdown when clickSuggest is called and suggestions are not null', () => {
|
it('should toggle dropdown when clickSuggest is called and suggestions are not null', () => {
|
||||||
fixture.componentRef.setInput('aiEnabled', true)
|
fixture.componentRef.setInput('aiEnabled', true)
|
||||||
|
fixture.componentRef.setInput('fetchedSources', [SuggestionSource.ML])
|
||||||
fixture.detectChanges()
|
fixture.detectChanges()
|
||||||
fixture.componentRef.setInput('suggestions', {
|
fixture.componentRef.setInput('suggestions', {
|
||||||
suggested_correspondents: [],
|
suggested_correspondents: [],
|
||||||
suggested_tags: [],
|
suggested_tags: [],
|
||||||
suggested_document_types: [],
|
suggested_document_types: [],
|
||||||
})
|
})
|
||||||
|
fixture.detectChanges()
|
||||||
component.clickSuggest()
|
component.clickSuggest()
|
||||||
expect(component.dropdown.open).toBeTruthy()
|
expect(component.dropdown.isOpen()).toBeTruthy()
|
||||||
expect(fixture.nativeElement.textContent).toContain('No novel suggestions')
|
expect(fixture.nativeElement.textContent).toContain('No novel suggestions')
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('should fetch unfetched sources and show existing suggestions', () => {
|
||||||
|
jest.spyOn(component.getSuggestions, 'emit')
|
||||||
|
fixture.componentRef.setInput('aiEnabled', true)
|
||||||
|
fixture.componentRef.setInput('source', SuggestionSource.Both)
|
||||||
|
fixture.componentRef.setInput('fetchedSources', [SuggestionSource.ML])
|
||||||
|
fixture.componentRef.setInput('suggestions', { tags: [1] })
|
||||||
|
fixture.detectChanges()
|
||||||
|
component.clickSuggest()
|
||||||
|
expect(component.getSuggestions.emit).toHaveBeenCalledWith(
|
||||||
|
SuggestionSource.Both
|
||||||
|
)
|
||||||
|
expect(component.dropdown.isOpen()).toBeTruthy()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should only show source options when AI is enabled', () => {
|
||||||
|
expect(
|
||||||
|
fixture.nativeElement.querySelector('#suggestionSourceML')
|
||||||
|
).toBeNull()
|
||||||
|
fixture.componentRef.setInput('aiEnabled', true)
|
||||||
|
fixture.detectChanges()
|
||||||
|
fixture.nativeElement
|
||||||
|
.querySelector('button[title="Suggestion options"]')
|
||||||
|
.click()
|
||||||
|
fixture.detectChanges()
|
||||||
|
expect(
|
||||||
|
fixture.nativeElement.querySelector('#suggestionSourceML')
|
||||||
|
).not.toBeNull()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should emit source changes and never allow no source', () => {
|
||||||
|
const emitSpy = jest.spyOn(component.sourceChange, 'emit')
|
||||||
|
component.setSources(true, true)
|
||||||
|
expect(emitSpy).toHaveBeenCalledWith(SuggestionSource.Both)
|
||||||
|
component.setSources(true, false)
|
||||||
|
expect(emitSpy).toHaveBeenCalledWith(SuggestionSource.ML)
|
||||||
|
component.setSources(false, true)
|
||||||
|
expect(emitSpy).toHaveBeenCalledWith(SuggestionSource.AI)
|
||||||
|
|
||||||
|
emitSpy.mockClear()
|
||||||
|
component.setSources(false, false)
|
||||||
|
expect(emitSpy).not.toHaveBeenCalled()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should indicate a non-default source', () => {
|
||||||
|
fixture.componentRef.setInput('aiEnabled', true)
|
||||||
|
fixture.componentRef.setInput('source', SuggestionSource.AI)
|
||||||
|
fixture.componentRef.setInput('defaultSource', SuggestionSource.AI)
|
||||||
|
fixture.detectChanges()
|
||||||
|
expect(fixture.nativeElement.textContent).not.toContain('Not using default')
|
||||||
|
|
||||||
|
fixture.componentRef.setInput('source', SuggestionSource.Both)
|
||||||
|
fixture.detectChanges()
|
||||||
|
expect(fixture.nativeElement.textContent).toContain('Not using default')
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
+40
-3
@@ -8,6 +8,7 @@ import {
|
|||||||
import { NgbDropdown, NgbDropdownModule } from '@ng-bootstrap/ng-bootstrap'
|
import { NgbDropdown, NgbDropdownModule } from '@ng-bootstrap/ng-bootstrap'
|
||||||
import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons'
|
import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons'
|
||||||
import { DocumentSuggestions } from 'src/app/data/document-suggestions'
|
import { DocumentSuggestions } from 'src/app/data/document-suggestions'
|
||||||
|
import { SuggestionSource } from 'src/app/data/ui-settings'
|
||||||
import { pngxPopperOptions } from 'src/app/utils/popper-options'
|
import { pngxPopperOptions } from 'src/app/utils/popper-options'
|
||||||
|
|
||||||
@Component({
|
@Component({
|
||||||
@@ -18,12 +19,16 @@ import { pngxPopperOptions } from 'src/app/utils/popper-options'
|
|||||||
})
|
})
|
||||||
export class SuggestionsDropdownComponent {
|
export class SuggestionsDropdownComponent {
|
||||||
public popperOptions = pngxPopperOptions
|
public popperOptions = pngxPopperOptions
|
||||||
|
public readonly SuggestionSource = SuggestionSource
|
||||||
|
|
||||||
@ViewChild('dropdown') dropdown: NgbDropdown
|
@ViewChild('dropdown') dropdown: NgbDropdown
|
||||||
readonly suggestions = input<DocumentSuggestions>(null)
|
readonly suggestions = input<DocumentSuggestions>(null)
|
||||||
readonly aiEnabled = input(false)
|
readonly aiEnabled = input(false)
|
||||||
readonly loading = input(false)
|
readonly loading = input(false)
|
||||||
readonly disabled = input(false)
|
readonly disabled = input(false)
|
||||||
|
readonly source = input<SuggestionSource>(SuggestionSource.ML)
|
||||||
|
readonly defaultSource = input<SuggestionSource>(SuggestionSource.ML)
|
||||||
|
readonly fetchedSources = input<SuggestionSource[]>([])
|
||||||
|
|
||||||
readonly appliedTags = input<number[]>([])
|
readonly appliedTags = input<number[]>([])
|
||||||
readonly appliedCorrespondent = input<number>(null)
|
readonly appliedCorrespondent = input<number>(null)
|
||||||
@@ -31,8 +36,10 @@ export class SuggestionsDropdownComponent {
|
|||||||
readonly appliedStoragePath = input<number>(null)
|
readonly appliedStoragePath = input<number>(null)
|
||||||
|
|
||||||
@Output()
|
@Output()
|
||||||
getSuggestions: EventEmitter<SuggestionsDropdownComponent> =
|
getSuggestions: EventEmitter<SuggestionSource> = new EventEmitter()
|
||||||
new EventEmitter()
|
|
||||||
|
@Output()
|
||||||
|
sourceChange: EventEmitter<SuggestionSource> = new EventEmitter()
|
||||||
|
|
||||||
@Output()
|
@Output()
|
||||||
addTag: EventEmitter<string> = new EventEmitter()
|
addTag: EventEmitter<string> = new EventEmitter()
|
||||||
@@ -53,12 +60,42 @@ export class SuggestionsDropdownComponent {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if (!this.suggestions()) {
|
if (!this.suggestions()) {
|
||||||
this.getSuggestions.emit(this)
|
this.getSuggestions.emit(this.source())
|
||||||
|
} else if (this.hasUnfetchedSources) {
|
||||||
|
// sources changed, fetch the rest and show what we have meanwhile
|
||||||
|
this.getSuggestions.emit(this.source())
|
||||||
|
this.dropdown?.open()
|
||||||
} else {
|
} else {
|
||||||
this.dropdown?.toggle()
|
this.dropdown?.toggle()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
get useML(): boolean {
|
||||||
|
return this.source() !== SuggestionSource.AI
|
||||||
|
}
|
||||||
|
|
||||||
|
get useAI(): boolean {
|
||||||
|
return this.source() !== SuggestionSource.ML
|
||||||
|
}
|
||||||
|
|
||||||
|
get hasUnfetchedSources(): boolean {
|
||||||
|
const fetched = this.fetchedSources()
|
||||||
|
return (
|
||||||
|
(this.useML && !fetched.includes(SuggestionSource.ML)) ||
|
||||||
|
(this.useAI && !fetched.includes(SuggestionSource.AI))
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
public setSources(ml: boolean, ai: boolean) {
|
||||||
|
if (ml && ai) {
|
||||||
|
this.sourceChange.emit(SuggestionSource.Both)
|
||||||
|
} else if (ml) {
|
||||||
|
this.sourceChange.emit(SuggestionSource.ML)
|
||||||
|
} else if (ai) {
|
||||||
|
this.sourceChange.emit(SuggestionSource.AI)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
get novelSuggestions(): number {
|
get novelSuggestions(): number {
|
||||||
return (
|
return (
|
||||||
(this.suggestions()?.suggested_correspondents?.length ?? 0) +
|
(this.suggestions()?.suggested_correspondents?.length ?? 0) +
|
||||||
|
|||||||
@@ -134,11 +134,15 @@
|
|||||||
[loading]="suggestionsLoading()"
|
[loading]="suggestionsLoading()"
|
||||||
[suggestions]="suggestions()"
|
[suggestions]="suggestions()"
|
||||||
[aiEnabled]="aiEnabled"
|
[aiEnabled]="aiEnabled"
|
||||||
|
[source]="suggestionSource"
|
||||||
|
[defaultSource]="defaultSuggestionSource"
|
||||||
|
[fetchedSources]="fetchedSuggestionSources()"
|
||||||
[appliedTags]="documentForm.value.tags"
|
[appliedTags]="documentForm.value.tags"
|
||||||
[appliedCorrespondent]="documentForm.value.correspondent"
|
[appliedCorrespondent]="documentForm.value.correspondent"
|
||||||
[appliedDocumentType]="documentForm.value.document_type"
|
[appliedDocumentType]="documentForm.value.document_type"
|
||||||
[appliedStoragePath]="documentForm.value.storage_path"
|
[appliedStoragePath]="documentForm.value.storage_path"
|
||||||
(getSuggestions)="getSuggestions()"
|
(getSuggestions)="getSuggestions($event)"
|
||||||
|
(sourceChange)="suggestionSourceOverride.set($event)"
|
||||||
(addTag)="createTag($event)"
|
(addTag)="createTag($event)"
|
||||||
(addDocumentType)="createDocumentType($event)"
|
(addDocumentType)="createDocumentType($event)"
|
||||||
(addCorrespondent)="createCorrespondent($event)">
|
(addCorrespondent)="createCorrespondent($event)">
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ import {
|
|||||||
} from 'src/app/data/filter-rule-type'
|
} from 'src/app/data/filter-rule-type'
|
||||||
import { StoragePath } from 'src/app/data/storage-path'
|
import { StoragePath } from 'src/app/data/storage-path'
|
||||||
import { Tag } from 'src/app/data/tag'
|
import { Tag } from 'src/app/data/tag'
|
||||||
import { SETTINGS_KEYS } from 'src/app/data/ui-settings'
|
import { SETTINGS_KEYS, SuggestionSource } from 'src/app/data/ui-settings'
|
||||||
import { PermissionsGuard } from 'src/app/guards/permissions.guard'
|
import { PermissionsGuard } from 'src/app/guards/permissions.guard'
|
||||||
import { CustomDatePipe } from 'src/app/pipes/custom-date.pipe'
|
import { CustomDatePipe } from 'src/app/pipes/custom-date.pipe'
|
||||||
import { DocumentTitlePipe } from 'src/app/pipes/document-title.pipe'
|
import { DocumentTitlePipe } from 'src/app/pipes/document-title.pipe'
|
||||||
@@ -1528,6 +1528,113 @@ describe('DocumentDetailComponent', () => {
|
|||||||
expect(component.suggestionsLoading()).toBeFalsy()
|
expect(component.suggestionsLoading()).toBeFalsy()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('should get and merge ML and AI suggestions when source is both', () => {
|
||||||
|
settingsService.set(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE,
|
||||||
|
SuggestionSource.Both
|
||||||
|
)
|
||||||
|
const getSetting = settingsService.get.bind(settingsService)
|
||||||
|
jest
|
||||||
|
.spyOn(settingsService, 'get')
|
||||||
|
.mockImplementation((key) =>
|
||||||
|
key === SETTINGS_KEYS.AI_ENABLED ? true : getSetting(key)
|
||||||
|
)
|
||||||
|
const suggestionsSpy = jest
|
||||||
|
.spyOn(documentService, 'getSuggestions')
|
||||||
|
.mockReturnValue(of({ tags: [42], dates: ['2024-01-01'] }))
|
||||||
|
const aiSuggestionsSpy = jest
|
||||||
|
.spyOn(documentService, 'getAiSuggestions')
|
||||||
|
.mockReturnValue(
|
||||||
|
of({ title: 'AI title', tags: [42, 43], suggested_tags: ['New'] })
|
||||||
|
)
|
||||||
|
initNormally()
|
||||||
|
expect(suggestionsSpy).toHaveBeenCalled()
|
||||||
|
expect(aiSuggestionsSpy).toHaveBeenCalled()
|
||||||
|
expect(component.suggestions().title).toEqual('AI title')
|
||||||
|
expect(component.suggestions().tags).toEqual([42, 43])
|
||||||
|
expect(component.suggestions().suggested_tags).toEqual(['New'])
|
||||||
|
expect(component.suggestions().dates).toEqual(['2024-01-01'])
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should only fetch sources not yet fetched for the document', () => {
|
||||||
|
settingsService.set(SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST, false)
|
||||||
|
settingsService.set(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE,
|
||||||
|
SuggestionSource.ML
|
||||||
|
)
|
||||||
|
const getSetting = settingsService.get.bind(settingsService)
|
||||||
|
jest
|
||||||
|
.spyOn(settingsService, 'get')
|
||||||
|
.mockImplementation((key) =>
|
||||||
|
key === SETTINGS_KEYS.AI_ENABLED ? true : getSetting(key)
|
||||||
|
)
|
||||||
|
const suggestionsSpy = jest
|
||||||
|
.spyOn(documentService, 'getSuggestions')
|
||||||
|
.mockReturnValue(of({ tags: [42] }))
|
||||||
|
const aiSuggestionsSpy = jest
|
||||||
|
.spyOn(documentService, 'getAiSuggestions')
|
||||||
|
.mockReturnValue(of({ tags: [43] }))
|
||||||
|
initNormally()
|
||||||
|
|
||||||
|
component.getSuggestions()
|
||||||
|
expect(suggestionsSpy).toHaveBeenCalledTimes(1)
|
||||||
|
expect(aiSuggestionsSpy).not.toHaveBeenCalled()
|
||||||
|
|
||||||
|
component.getSuggestions(SuggestionSource.Both)
|
||||||
|
expect(suggestionsSpy).toHaveBeenCalledTimes(1)
|
||||||
|
expect(aiSuggestionsSpy).toHaveBeenCalledTimes(1)
|
||||||
|
expect(component.suggestions().tags).toEqual([42, 43])
|
||||||
|
|
||||||
|
component.getSuggestions(SuggestionSource.Both)
|
||||||
|
expect(suggestionsSpy).toHaveBeenCalledTimes(1)
|
||||||
|
expect(aiSuggestionsSpy).toHaveBeenCalledTimes(1)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should use the per-document source override and reset it on document change', () => {
|
||||||
|
settingsService.set(SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST, false)
|
||||||
|
const getSetting = settingsService.get.bind(settingsService)
|
||||||
|
jest
|
||||||
|
.spyOn(settingsService, 'get')
|
||||||
|
.mockImplementation((key) =>
|
||||||
|
key === SETTINGS_KEYS.AI_ENABLED ? true : getSetting(key)
|
||||||
|
)
|
||||||
|
initNormally()
|
||||||
|
expect(component.suggestionSource).toEqual(SuggestionSource.AI)
|
||||||
|
component.suggestionSourceOverride.set(SuggestionSource.ML)
|
||||||
|
expect(component.suggestionSource).toEqual(SuggestionSource.ML)
|
||||||
|
|
||||||
|
jest
|
||||||
|
.spyOn(documentService, 'get')
|
||||||
|
.mockReturnValueOnce(of(Object.assign({}, doc)))
|
||||||
|
;(component as any).loadDocument(doc.id, true)
|
||||||
|
expect(component.suggestionSourceOverride()).toBeNull()
|
||||||
|
expect(component.fetchedSuggestionSources()).toEqual([])
|
||||||
|
})
|
||||||
|
|
||||||
|
it('should keep suggestions from one source if the other fails', () => {
|
||||||
|
settingsService.set(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE,
|
||||||
|
SuggestionSource.Both
|
||||||
|
)
|
||||||
|
const getSetting = settingsService.get.bind(settingsService)
|
||||||
|
jest
|
||||||
|
.spyOn(settingsService, 'get')
|
||||||
|
.mockImplementation((key) =>
|
||||||
|
key === SETTINGS_KEYS.AI_ENABLED ? true : getSetting(key)
|
||||||
|
)
|
||||||
|
const errorSpy = jest.spyOn(toastService, 'showError')
|
||||||
|
jest
|
||||||
|
.spyOn(documentService, 'getSuggestions')
|
||||||
|
.mockReturnValue(of({ tags: [42] }))
|
||||||
|
jest
|
||||||
|
.spyOn(documentService, 'getAiSuggestions')
|
||||||
|
.mockReturnValue(throwError(() => new Error('failed')))
|
||||||
|
initNormally()
|
||||||
|
expect(errorSpy).toHaveBeenCalled()
|
||||||
|
expect(component.suggestions().tags).toEqual([42])
|
||||||
|
expect(component.fetchedSuggestionSources()).toEqual([SuggestionSource.ML])
|
||||||
|
})
|
||||||
|
|
||||||
it('should show error if needed for get suggestions', () => {
|
it('should show error if needed for get suggestions', () => {
|
||||||
const suggestionsSpy = jest.spyOn(documentService, 'getSuggestions')
|
const suggestionsSpy = jest.spyOn(documentService, 'getSuggestions')
|
||||||
const errorSpy = jest.spyOn(toastService, 'showError')
|
const errorSpy = jest.spyOn(toastService, 'showError')
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ import {
|
|||||||
import { dirtyCheck, DirtyComponent } from '@ngneat/dirty-check-forms'
|
import { dirtyCheck, DirtyComponent } from '@ngneat/dirty-check-forms'
|
||||||
import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons'
|
import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons'
|
||||||
import { DeviceDetectorService } from 'ngx-device-detector'
|
import { DeviceDetectorService } from 'ngx-device-detector'
|
||||||
import { BehaviorSubject, Observable, of, Subject, timer } from 'rxjs'
|
import { BehaviorSubject, merge, Observable, of, Subject, timer } from 'rxjs'
|
||||||
import {
|
import {
|
||||||
catchError,
|
catchError,
|
||||||
debounceTime,
|
debounceTime,
|
||||||
@@ -48,7 +48,10 @@ import { DataType } from 'src/app/data/datatype'
|
|||||||
import { Document, DocumentVersionInfo } from 'src/app/data/document'
|
import { Document, DocumentVersionInfo } from 'src/app/data/document'
|
||||||
import { DocumentMetadata } from 'src/app/data/document-metadata'
|
import { DocumentMetadata } from 'src/app/data/document-metadata'
|
||||||
import { DocumentNote } from 'src/app/data/document-note'
|
import { DocumentNote } from 'src/app/data/document-note'
|
||||||
import { DocumentSuggestions } from 'src/app/data/document-suggestions'
|
import {
|
||||||
|
DocumentSuggestions,
|
||||||
|
mergeSuggestions,
|
||||||
|
} from 'src/app/data/document-suggestions'
|
||||||
import { DocumentType } from 'src/app/data/document-type'
|
import { DocumentType } from 'src/app/data/document-type'
|
||||||
import { FilterRule } from 'src/app/data/filter-rule'
|
import { FilterRule } from 'src/app/data/filter-rule'
|
||||||
import {
|
import {
|
||||||
@@ -63,7 +66,7 @@ import {
|
|||||||
import { ObjectWithId } from 'src/app/data/object-with-id'
|
import { ObjectWithId } from 'src/app/data/object-with-id'
|
||||||
import { StoragePath } from 'src/app/data/storage-path'
|
import { StoragePath } from 'src/app/data/storage-path'
|
||||||
import { Tag } from 'src/app/data/tag'
|
import { Tag } from 'src/app/data/tag'
|
||||||
import { SETTINGS_KEYS } from 'src/app/data/ui-settings'
|
import { SETTINGS_KEYS, SuggestionSource } from 'src/app/data/ui-settings'
|
||||||
import { User } from 'src/app/data/user'
|
import { User } from 'src/app/data/user'
|
||||||
import { IfPermissionsDirective } from 'src/app/directives/if-permissions.directive'
|
import { IfPermissionsDirective } from 'src/app/directives/if-permissions.directive'
|
||||||
import { CustomDatePipe } from 'src/app/pipes/custom-date.pipe'
|
import { CustomDatePipe } from 'src/app/pipes/custom-date.pipe'
|
||||||
@@ -240,6 +243,10 @@ export class DocumentDetailComponent
|
|||||||
private readonly autoSuggestSetting = this.settings.getSignal<boolean>(
|
private readonly autoSuggestSetting = this.settings.getSignal<boolean>(
|
||||||
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST
|
SETTINGS_KEYS.DOCUMENT_EDITING_AUTO_SUGGEST
|
||||||
)
|
)
|
||||||
|
private readonly suggestionSourceSetting =
|
||||||
|
this.settings.getSignal<SuggestionSource>(
|
||||||
|
SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE
|
||||||
|
)
|
||||||
private readonly hiddenFieldsSetting = this.settings.getSignal<
|
private readonly hiddenFieldsSetting = this.settings.getSignal<
|
||||||
DocumentDetailFieldID[]
|
DocumentDetailFieldID[]
|
||||||
>(SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS)
|
>(SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS)
|
||||||
@@ -261,6 +268,9 @@ export class DocumentDetailComponent
|
|||||||
readonly metadata = signal<DocumentMetadata>(undefined)
|
readonly metadata = signal<DocumentMetadata>(undefined)
|
||||||
readonly suggestions = signal<DocumentSuggestions>(undefined)
|
readonly suggestions = signal<DocumentSuggestions>(undefined)
|
||||||
readonly suggestionsLoading = signal(false)
|
readonly suggestionsLoading = signal(false)
|
||||||
|
// per-document, resets on navigation
|
||||||
|
readonly suggestionSourceOverride = signal<SuggestionSource>(null)
|
||||||
|
readonly fetchedSuggestionSources = signal<SuggestionSource[]>([])
|
||||||
readonly users = signal<User[]>(undefined)
|
readonly users = signal<User[]>(undefined)
|
||||||
|
|
||||||
readonly title = signal<string>(undefined)
|
readonly title = signal<string>(undefined)
|
||||||
@@ -365,6 +375,15 @@ export class DocumentDetailComponent
|
|||||||
return this.autoSuggestSetting()
|
return this.autoSuggestSetting()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
get defaultSuggestionSource(): SuggestionSource {
|
||||||
|
return this.aiEnabled ? this.suggestionSourceSetting() : SuggestionSource.ML
|
||||||
|
}
|
||||||
|
|
||||||
|
get suggestionSource(): SuggestionSource {
|
||||||
|
if (!this.aiEnabled) return SuggestionSource.ML
|
||||||
|
return this.suggestionSourceOverride() ?? this.defaultSuggestionSource
|
||||||
|
}
|
||||||
|
|
||||||
get archiveContentRenderType(): ContentRenderType {
|
get archiveContentRenderType(): ContentRenderType {
|
||||||
const hasArchiveVersion =
|
const hasArchiveVersion =
|
||||||
this.metadata()?.has_archive_version ??
|
this.metadata()?.has_archive_version ??
|
||||||
@@ -590,6 +609,8 @@ export class DocumentDetailComponent
|
|||||||
}
|
}
|
||||||
this.documentId.set(doc.id)
|
this.documentId.set(doc.id)
|
||||||
this.suggestions.set(null)
|
this.suggestions.set(null)
|
||||||
|
this.suggestionSourceOverride.set(null)
|
||||||
|
this.fetchedSuggestionSources.set([])
|
||||||
const openDocument = this.openDocumentService.getOpenDocument(
|
const openDocument = this.openDocumentService.getOpenDocument(
|
||||||
this.documentId()
|
this.documentId()
|
||||||
)
|
)
|
||||||
@@ -1077,29 +1098,44 @@ export class DocumentDetailComponent
|
|||||||
return this.documentForm.get('custom_fields') as FormArray
|
return this.documentForm.get('custom_fields') as FormArray
|
||||||
}
|
}
|
||||||
|
|
||||||
getSuggestions() {
|
getSuggestions(source: SuggestionSource = this.suggestionSource) {
|
||||||
|
const sources = (
|
||||||
|
source === SuggestionSource.Both
|
||||||
|
? [SuggestionSource.ML, SuggestionSource.AI]
|
||||||
|
: [source]
|
||||||
|
).filter((s) => !this.fetchedSuggestionSources().includes(s))
|
||||||
|
if (!sources.length) return
|
||||||
|
|
||||||
this.suggestionsLoading.set(true)
|
this.suggestionsLoading.set(true)
|
||||||
const suggestionsObservable = this.aiEnabled
|
merge(
|
||||||
? this.documentsService.getAiSuggestions(this.documentId())
|
...sources.map((s) =>
|
||||||
: this.documentsService.getSuggestions(this.documentId())
|
(s === SuggestionSource.AI
|
||||||
suggestionsObservable
|
? this.documentsService.getAiSuggestions(this.documentId())
|
||||||
|
: this.documentsService.getSuggestions(this.documentId())
|
||||||
|
).pipe(
|
||||||
|
first(),
|
||||||
|
map((result) => ({ source: s, result })),
|
||||||
|
catchError((error) => {
|
||||||
|
this.toastService.showError(
|
||||||
|
$localize`Error retrieving suggestions.`,
|
||||||
|
error
|
||||||
|
)
|
||||||
|
return of(null)
|
||||||
|
})
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
.pipe(
|
.pipe(
|
||||||
first(),
|
|
||||||
takeUntil(this.unsubscribeNotifier),
|
takeUntil(this.unsubscribeNotifier),
|
||||||
takeUntil(this.docChangeNotifier),
|
takeUntil(this.docChangeNotifier),
|
||||||
finalize(() => this.suggestionsLoading.set(false))
|
finalize(() => this.suggestionsLoading.set(false))
|
||||||
)
|
)
|
||||||
.subscribe({
|
.subscribe((response) => {
|
||||||
next: (result) => {
|
if (!response) return
|
||||||
this.suggestions.set(result)
|
this.fetchedSuggestionSources.update((f) => [...f, response.source])
|
||||||
},
|
this.suggestions.set(
|
||||||
error: (error) => {
|
mergeSuggestions(this.suggestions(), response.result)
|
||||||
this.suggestions.set(null)
|
)
|
||||||
this.toastService.showError(
|
|
||||||
$localize`Error retrieving suggestions.`,
|
|
||||||
error
|
|
||||||
)
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -146,6 +146,19 @@ describe('DocumentListComponent', () => {
|
|||||||
expect(reloadSpy).toHaveBeenCalled()
|
expect(reloadSpy).toHaveBeenCalled()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('should stop reloading on document deleted after destroy', () => {
|
||||||
|
const reloadSpy = jest.spyOn(documentListService, 'reload')
|
||||||
|
const documentDeletedSubject = new Subject<boolean>()
|
||||||
|
jest
|
||||||
|
.spyOn(websocketStatusService, 'onDocumentDeleted')
|
||||||
|
.mockReturnValue(documentDeletedSubject)
|
||||||
|
fixture.detectChanges()
|
||||||
|
fixture.destroy()
|
||||||
|
reloadSpy.mockClear()
|
||||||
|
documentDeletedSubject.next(true)
|
||||||
|
expect(reloadSpy).not.toHaveBeenCalled()
|
||||||
|
})
|
||||||
|
|
||||||
it('should show score sort fields on fulltext queries', () => {
|
it('should show score sort fields on fulltext queries', () => {
|
||||||
documentListService.setFilterRules([
|
documentListService.setFilterRules([
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -270,9 +270,12 @@ export class DocumentListComponent
|
|||||||
this.list.reload()
|
this.list.reload()
|
||||||
})
|
})
|
||||||
|
|
||||||
this.websocketStatusService.onDocumentDeleted().subscribe(() => {
|
this.websocketStatusService
|
||||||
this.list.reload()
|
.onDocumentDeleted()
|
||||||
})
|
.pipe(takeUntil(this.unsubscribeNotifier))
|
||||||
|
.subscribe(() => {
|
||||||
|
this.list.reload()
|
||||||
|
})
|
||||||
|
|
||||||
this.route.paramMap
|
this.route.paramMap
|
||||||
.pipe(
|
.pipe(
|
||||||
|
|||||||
@@ -15,3 +15,33 @@ export interface DocumentSuggestions {
|
|||||||
|
|
||||||
dates?: string[] // ISO-formatted date string e.g. 2022-11-03
|
dates?: string[] // ISO-formatted date string e.g. 2022-11-03
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const union = <T>(a: T[] = [], b: T[] = []): T[] => [...new Set([...a, ...b])]
|
||||||
|
|
||||||
|
export function mergeSuggestions(
|
||||||
|
a: DocumentSuggestions,
|
||||||
|
b: DocumentSuggestions
|
||||||
|
): DocumentSuggestions {
|
||||||
|
if (!a) return b
|
||||||
|
return {
|
||||||
|
title: a.title || b.title,
|
||||||
|
tags: union(a.tags, b.tags),
|
||||||
|
suggested_tags: union(a.suggested_tags, b.suggested_tags),
|
||||||
|
correspondents: union(a.correspondents, b.correspondents),
|
||||||
|
suggested_correspondents: union(
|
||||||
|
a.suggested_correspondents,
|
||||||
|
b.suggested_correspondents
|
||||||
|
),
|
||||||
|
document_types: union(a.document_types, b.document_types),
|
||||||
|
suggested_document_types: union(
|
||||||
|
a.suggested_document_types,
|
||||||
|
b.suggested_document_types
|
||||||
|
),
|
||||||
|
storage_paths: union(a.storage_paths, b.storage_paths),
|
||||||
|
suggested_storage_paths: union(
|
||||||
|
a.suggested_storage_paths,
|
||||||
|
b.suggested_storage_paths
|
||||||
|
),
|
||||||
|
dates: union(a.dates, b.dates),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -353,6 +353,14 @@ export const PaperlessConfigOptions: ConfigOption[] = [
|
|||||||
config_key: 'PAPERLESS_AI_LLM_EMBEDDING_MODEL',
|
config_key: 'PAPERLESS_AI_LLM_EMBEDDING_MODEL',
|
||||||
category: ConfigCategory.AI,
|
category: ConfigCategory.AI,
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
key: 'llm_embedding_api_key',
|
||||||
|
title: $localize`LLM Embedding API Key`,
|
||||||
|
type: ConfigOptionType.Password,
|
||||||
|
note: $localize`Used for embeddings when set, otherwise LLM API key is used.`,
|
||||||
|
config_key: 'PAPERLESS_AI_LLM_EMBEDDING_API_KEY',
|
||||||
|
category: ConfigCategory.AI,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
key: 'llm_embedding_endpoint',
|
key: 'llm_embedding_endpoint',
|
||||||
title: $localize`LLM Embedding Endpoint`,
|
title: $localize`LLM Embedding Endpoint`,
|
||||||
@@ -457,6 +465,7 @@ export interface PaperlessConfig extends ObjectWithId {
|
|||||||
ai_enabled: boolean
|
ai_enabled: boolean
|
||||||
llm_embedding_backend: string
|
llm_embedding_backend: string
|
||||||
llm_embedding_model: string
|
llm_embedding_model: string
|
||||||
|
llm_embedding_api_key: string
|
||||||
llm_embedding_endpoint: string
|
llm_embedding_endpoint: string
|
||||||
llm_embedding_chunk_size: number
|
llm_embedding_chunk_size: number
|
||||||
llm_context_size: number
|
llm_context_size: number
|
||||||
|
|||||||
@@ -20,6 +20,12 @@ export enum GlobalSearchType {
|
|||||||
TITLE_CONTENT = 'title-content',
|
TITLE_CONTENT = 'title-content',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export enum SuggestionSource {
|
||||||
|
ML = 'ml',
|
||||||
|
AI = 'ai',
|
||||||
|
Both = 'both',
|
||||||
|
}
|
||||||
|
|
||||||
export enum CollapsibleSection {
|
export enum CollapsibleSection {
|
||||||
ATTRIBUTES = 'attributes',
|
ATTRIBUTES = 'attributes',
|
||||||
}
|
}
|
||||||
@@ -98,6 +104,8 @@ export const SETTINGS_KEYS = {
|
|||||||
'general-settings:document-editing:overlay-thumbnail',
|
'general-settings:document-editing:overlay-thumbnail',
|
||||||
DOCUMENT_EDITING_AUTO_SUGGEST:
|
DOCUMENT_EDITING_AUTO_SUGGEST:
|
||||||
'general-settings:document-editing:auto-suggest',
|
'general-settings:document-editing:auto-suggest',
|
||||||
|
DOCUMENT_EDITING_SUGGESTION_SOURCE:
|
||||||
|
'general-settings:document-editing:suggestion-source',
|
||||||
DOCUMENT_DETAILS_HIDDEN_FIELDS:
|
DOCUMENT_DETAILS_HIDDEN_FIELDS:
|
||||||
'general-settings:document-details:hidden-fields',
|
'general-settings:document-details:hidden-fields',
|
||||||
SEARCH_DB_ONLY: 'general-settings:search:db-only',
|
SEARCH_DB_ONLY: 'general-settings:search:db-only',
|
||||||
@@ -326,6 +334,11 @@ export const SETTINGS: UiSetting[] = [
|
|||||||
type: 'boolean',
|
type: 'boolean',
|
||||||
default: true,
|
default: true,
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
key: SETTINGS_KEYS.DOCUMENT_EDITING_SUGGESTION_SOURCE,
|
||||||
|
type: 'string',
|
||||||
|
default: SuggestionSource.AI,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
key: SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS,
|
key: SETTINGS_KEYS.DOCUMENT_DETAILS_HIDDEN_FIELDS,
|
||||||
type: 'array',
|
type: 'array',
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ import {
|
|||||||
clipboardCheckFill,
|
clipboardCheckFill,
|
||||||
clipboardFill,
|
clipboardFill,
|
||||||
clockHistory,
|
clockHistory,
|
||||||
|
cpu,
|
||||||
creditCard,
|
creditCard,
|
||||||
dash,
|
dash,
|
||||||
dashCircle,
|
dashCircle,
|
||||||
@@ -118,6 +119,7 @@ import {
|
|||||||
infoCircle,
|
infoCircle,
|
||||||
journalBookmarkFill,
|
journalBookmarkFill,
|
||||||
journals,
|
journals,
|
||||||
|
lightbulb,
|
||||||
link,
|
link,
|
||||||
list,
|
list,
|
||||||
listNested,
|
listNested,
|
||||||
@@ -322,6 +324,7 @@ const icons = {
|
|||||||
clipboardCheckFill,
|
clipboardCheckFill,
|
||||||
clipboardFill,
|
clipboardFill,
|
||||||
clockHistory,
|
clockHistory,
|
||||||
|
cpu,
|
||||||
cash,
|
cash,
|
||||||
creditCard,
|
creditCard,
|
||||||
dash,
|
dash,
|
||||||
@@ -367,6 +370,7 @@ const icons = {
|
|||||||
infoCircle,
|
infoCircle,
|
||||||
journalBookmarkFill,
|
journalBookmarkFill,
|
||||||
journals,
|
journals,
|
||||||
|
lightbulb,
|
||||||
link,
|
link,
|
||||||
list,
|
list,
|
||||||
listNested,
|
listNested,
|
||||||
|
|||||||
@@ -292,6 +292,7 @@ a.btn-link:active,
|
|||||||
a.btn-link:focus-visible,
|
a.btn-link:focus-visible,
|
||||||
.btn-link:hover,
|
.btn-link:hover,
|
||||||
.btn-link:active,
|
.btn-link:active,
|
||||||
|
.btn-link.show,
|
||||||
.btn-link:focus-visible {
|
.btn-link:focus-visible {
|
||||||
color: var(--pngx-primary-lighten-10) !important;
|
color: var(--pngx-primary-lighten-10) !important;
|
||||||
.primary-light & {
|
.primary-light & {
|
||||||
|
|||||||
@@ -25,10 +25,20 @@ export class PDFFindController {
|
|||||||
onIsPageVisible?: () => boolean
|
onIsPageVisible?: () => boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export const LinkTarget = {
|
||||||
|
NONE: 0,
|
||||||
|
SELF: 1,
|
||||||
|
BLANK: 2,
|
||||||
|
PARENT: 3,
|
||||||
|
TOP: 4,
|
||||||
|
}
|
||||||
|
|
||||||
export class PDFLinkService {
|
export class PDFLinkService {
|
||||||
private document?: unknown
|
private document?: unknown
|
||||||
private viewer?: unknown
|
private viewer?: unknown
|
||||||
|
|
||||||
|
constructor(readonly options: Record<string, unknown> = {}) {}
|
||||||
|
|
||||||
setDocument(document: unknown): void {
|
setDocument(document: unknown): void {
|
||||||
this.document = document
|
this.document = document
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,14 +17,10 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
from django.contrib.auth.models import User
|
from django.contrib.auth.models import User
|
||||||
from pytest_django.fixtures import Settings
|
from pytest_django.fixtures import Settings
|
||||||
from pytest_mock import MockerFixture
|
|
||||||
from rest_framework.test import APIClient
|
from rest_framework.test import APIClient
|
||||||
|
|
||||||
from paperless_testing.dirs import PaperlessDirs
|
from paperless_testing.dirs import PaperlessDirs
|
||||||
from paperless_testing.fakes.progress import FakeProgressManager
|
from paperless_testing.fakes.progress import FakeProgressManager
|
||||||
from paperless_testing.outbound import DialRecorder
|
|
||||||
from paperless_testing.outbound import FakeDNS
|
|
||||||
from paperless_testing.outbound import LocalHTTPServer
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session", autouse=True)
|
@pytest.fixture(scope="session", autouse=True)
|
||||||
@@ -153,41 +149,3 @@ def fake_progress_manager(
|
|||||||
|
|
||||||
monkeypatch.setattr("documents.tasks.ProgressManager", FakeProgressManager)
|
monkeypatch.setattr("documents.tasks.ProgressManager", FakeProgressManager)
|
||||||
return FakeProgressManager
|
return FakeProgressManager
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def local_http_server() -> Generator[LocalHTTPServer, None, None]:
|
|
||||||
"""A recording HTTP server on 127.0.0.1, for outbound connection tests."""
|
|
||||||
from paperless_testing.outbound import running_http_server
|
|
||||||
|
|
||||||
with running_http_server() as server:
|
|
||||||
yield server
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def fake_dns(mocker: MockerFixture) -> FakeDNS:
|
|
||||||
"""Per-hostname answers for the outbound guard's resolver hooks."""
|
|
||||||
from paperless_testing.outbound import install_fake_dns
|
|
||||||
|
|
||||||
return install_fake_dns(mocker)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def dial_recorder(mocker: MockerFixture) -> DialRecorder:
|
|
||||||
"""Records which addresses the outbound guard actually dialled."""
|
|
||||||
from paperless_testing.outbound import install_dial_recorder
|
|
||||||
|
|
||||||
return install_dial_recorder(mocker)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def every_address_is_public(mocker: MockerFixture) -> None:
|
|
||||||
"""Disable the outbound guard's address policy: every address passes.
|
|
||||||
|
|
||||||
For tests that are not themselves exercising which addresses the guard
|
|
||||||
accepts, so loopback and other private addresses dial just like a
|
|
||||||
public one.
|
|
||||||
"""
|
|
||||||
from paperless_testing.outbound import allow_all_addresses
|
|
||||||
|
|
||||||
allow_all_addresses(mocker)
|
|
||||||
|
|||||||
@@ -26,7 +26,6 @@ class DocumentsConfig(AppConfig):
|
|||||||
document_consumption_finished.connect(set_document_type)
|
document_consumption_finished.connect(set_document_type)
|
||||||
document_consumption_finished.connect(set_tags)
|
document_consumption_finished.connect(set_tags)
|
||||||
document_consumption_finished.connect(set_storage_path)
|
document_consumption_finished.connect(set_storage_path)
|
||||||
document_consumption_finished.connect(add_to_index)
|
|
||||||
document_consumption_finished.connect(run_workflows_added)
|
document_consumption_finished.connect(run_workflows_added)
|
||||||
document_consumption_finished.connect(add_to_index)
|
document_consumption_finished.connect(add_to_index)
|
||||||
document_consumption_finished.connect(add_or_update_document_in_llm_index)
|
document_consumption_finished.connect(add_or_update_document_in_llm_index)
|
||||||
|
|||||||
@@ -857,8 +857,9 @@ class ConsumerPlugin(
|
|||||||
self.log.debug(f"Creation date from parse_date: {create_date}")
|
self.log.debug(f"Creation date from parse_date: {create_date}")
|
||||||
else:
|
else:
|
||||||
stats = Path(self.input_doc.original_file).stat()
|
stats = Path(self.input_doc.original_file).stat()
|
||||||
create_date = timezone.make_aware(
|
create_date = datetime.datetime.fromtimestamp(
|
||||||
datetime.datetime.fromtimestamp(stats.st_mtime),
|
stats.st_mtime,
|
||||||
|
tz=timezone.get_current_timezone(),
|
||||||
)
|
)
|
||||||
self.log.debug(f"Creation date from st_mtime: {create_date}")
|
self.log.debug(f"Creation date from st_mtime: {create_date}")
|
||||||
|
|
||||||
|
|||||||
@@ -1079,6 +1079,16 @@ class DocumentSerializer(
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_page_count(self, obj) -> int | None:
|
def get_page_count(self, obj) -> int | None:
|
||||||
|
# Like content versions get their own page count from the newest version,
|
||||||
|
# use the prefetched versions cache to avoid an extra query
|
||||||
|
prefetched_cache = getattr(obj, "_prefetched_objects_cache", None)
|
||||||
|
prefetched_versions = (
|
||||||
|
prefetched_cache.get("versions")
|
||||||
|
if isinstance(prefetched_cache, dict)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if obj.root_document_id is None and prefetched_versions:
|
||||||
|
return sort_versions_newest_first(prefetched_versions)[0].page_count
|
||||||
return obj.page_count
|
return obj.page_count
|
||||||
|
|
||||||
@extend_schema_field(DuplicateDocumentSummarySerializer(many=True))
|
@extend_schema_field(DuplicateDocumentSummarySerializer(many=True))
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ from documents.permissions import get_objects_for_user_owner_aware
|
|||||||
from documents.plugins.helpers import DocumentsStatusManager
|
from documents.plugins.helpers import DocumentsStatusManager
|
||||||
from documents.templating.utils import convert_format_str_to_template_format
|
from documents.templating.utils import convert_format_str_to_template_format
|
||||||
from documents.utils import compute_checksum
|
from documents.utils import compute_checksum
|
||||||
|
from documents.utils import copy_file_with_basic_stats
|
||||||
from documents.workflows.actions import build_workflow_action_context
|
from documents.workflows.actions import build_workflow_action_context
|
||||||
from documents.workflows.actions import execute_email_action
|
from documents.workflows.actions import execute_email_action
|
||||||
from documents.workflows.actions import execute_move_to_trash_action
|
from documents.workflows.actions import execute_move_to_trash_action
|
||||||
@@ -363,7 +364,11 @@ def cleanup_document_deletion(sender, instance, **kwargs) -> None:
|
|||||||
|
|
||||||
logger.debug(f"Moving {instance.source_path} to trash at {new_file_path}")
|
logger.debug(f"Moving {instance.source_path} to trash at {new_file_path}")
|
||||||
try:
|
try:
|
||||||
shutil.move(instance.source_path, new_file_path)
|
shutil.move(
|
||||||
|
instance.source_path,
|
||||||
|
new_file_path,
|
||||||
|
copy_function=copy_file_with_basic_stats,
|
||||||
|
)
|
||||||
except OSError as e:
|
except OSError as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Failed to move {instance.source_path} to trash at "
|
f"Failed to move {instance.source_path} to trash at "
|
||||||
|
|||||||
@@ -12,7 +12,7 @@
|
|||||||
<meta name="robots" content="noindex,nofollow">
|
<meta name="robots" content="noindex,nofollow">
|
||||||
<meta name="author" content="The Paperless-ngx Team">
|
<meta name="author" content="The Paperless-ngx Team">
|
||||||
<link rel="icon" type="image/x-icon" href="favicon.ico">
|
<link rel="icon" type="image/x-icon" href="favicon.ico">
|
||||||
<link rel="manifest" href="{% static webmanifest %}">
|
<link rel="manifest" href="{% static webmanifest %}" crossorigin="use-credentials">
|
||||||
<link rel="stylesheet" href="{% static styles_css %}">
|
<link rel="stylesheet" href="{% static styles_css %}">
|
||||||
<link rel="apple-touch-icon" href="{% static apple_touch_icon %}">
|
<link rel="apple-touch-icon" href="{% static apple_touch_icon %}">
|
||||||
</head>
|
</head>
|
||||||
|
|||||||
@@ -81,6 +81,7 @@ class TestApiAppConfig(DirectoriesMixin, APITestCase):
|
|||||||
"ai_enabled": None,
|
"ai_enabled": None,
|
||||||
"llm_embedding_backend": None,
|
"llm_embedding_backend": None,
|
||||||
"llm_embedding_model": None,
|
"llm_embedding_model": None,
|
||||||
|
"llm_embedding_api_key": None,
|
||||||
"llm_embedding_endpoint": None,
|
"llm_embedding_endpoint": None,
|
||||||
"llm_embedding_chunk_size": None,
|
"llm_embedding_chunk_size": None,
|
||||||
"llm_context_size": None,
|
"llm_context_size": None,
|
||||||
@@ -922,6 +923,49 @@ class TestApiAppConfig(DirectoriesMixin, APITestCase):
|
|||||||
self.assertEqual(response.status_code, status.HTTP_405_METHOD_NOT_ALLOWED)
|
self.assertEqual(response.status_code, status.HTTP_405_METHOD_NOT_ALLOWED)
|
||||||
self.assertEqual(ApplicationConfiguration.objects.count(), 1)
|
self.assertEqual(ApplicationConfiguration.objects.count(), 1)
|
||||||
|
|
||||||
|
def test_update_llm_embedding_api_key(self) -> None:
|
||||||
|
"""
|
||||||
|
GIVEN:
|
||||||
|
- Existing config with llm_embedding_api_key specified
|
||||||
|
WHEN:
|
||||||
|
- API to update llm_embedding_api_key is called with all *s
|
||||||
|
- API to update llm_embedding_api_key is called with empty string
|
||||||
|
THEN:
|
||||||
|
- llm_embedding_api_key is unchanged
|
||||||
|
- llm_embedding_api_key is set to None
|
||||||
|
"""
|
||||||
|
config = ApplicationConfiguration.objects.first()
|
||||||
|
assert config is not None
|
||||||
|
config.llm_embedding_api_key = "1234567890"
|
||||||
|
config.save()
|
||||||
|
|
||||||
|
# Test with all *
|
||||||
|
response = self.client.patch(
|
||||||
|
f"{self.ENDPOINT}1/",
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"llm_embedding_api_key": "*" * 32,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
config.refresh_from_db()
|
||||||
|
self.assertEqual(config.llm_embedding_api_key, "1234567890")
|
||||||
|
# Test with empty string
|
||||||
|
response = self.client.patch(
|
||||||
|
f"{self.ENDPOINT}1/",
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"llm_embedding_api_key": "",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
config.refresh_from_db()
|
||||||
|
self.assertEqual(config.llm_embedding_api_key, None)
|
||||||
|
|
||||||
def test_update_llm_api_key(self) -> None:
|
def test_update_llm_api_key(self) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
|
|||||||
@@ -166,7 +166,15 @@ class TestBulkDownload(DirectoriesMixin, SampleDirMixin, APITestCase):
|
|||||||
),
|
),
|
||||||
content_type="application/json",
|
content_type="application/json",
|
||||||
)
|
)
|
||||||
response.close()
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response["Content-Type"], "application/zip")
|
||||||
|
|
||||||
|
with zipfile.ZipFile(io.BytesIO(read_streaming_response(response))) as zipf:
|
||||||
|
self.assertEqual(zipf.infolist()[0].compress_type, zipfile.ZIP_LZMA)
|
||||||
|
|
||||||
|
with self.doc2.source_file as f:
|
||||||
|
self.assertEqual(f.read(), zipf.read("2021-01-01 document A.pdf"))
|
||||||
|
|
||||||
@override_settings(FILENAME_FORMAT="{correspondent}/{title}")
|
@override_settings(FILENAME_FORMAT="{correspondent}/{title}")
|
||||||
def test_formatted_download_originals(self) -> None:
|
def test_formatted_download_originals(self) -> None:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from rest_framework.test import APITestCase
|
|||||||
|
|
||||||
from documents.models import Correspondent
|
from documents.models import Correspondent
|
||||||
from documents.models import CustomField
|
from documents.models import CustomField
|
||||||
|
from documents.models import CustomFieldInstance
|
||||||
from documents.models import Document
|
from documents.models import Document
|
||||||
from documents.models import DocumentType
|
from documents.models import DocumentType
|
||||||
from documents.models import StoragePath
|
from documents.models import StoragePath
|
||||||
@@ -2525,7 +2526,7 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
|
|||||||
WHEN:
|
WHEN:
|
||||||
- API to bulk edit documents is called
|
- API to bulk edit documents is called
|
||||||
THEN:
|
THEN:
|
||||||
- Audit log is created
|
- Audit log is created with the old and new correspondent
|
||||||
"""
|
"""
|
||||||
LogEntry.objects.all().delete()
|
LogEntry.objects.all().delete()
|
||||||
response = self.client.post(
|
response = self.client.post(
|
||||||
@@ -2541,7 +2542,8 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
self.assertEqual(LogEntry.objects.filter(object_pk=self.doc1.id).count(), 1)
|
entry = LogEntry.objects.get_for_object(self.doc1).get()
|
||||||
|
self.assertEqual(entry.changes, {"correspondent": [None, self.c2.id]})
|
||||||
|
|
||||||
@override_settings(AUDIT_LOG_ENABLED=True)
|
@override_settings(AUDIT_LOG_ENABLED=True)
|
||||||
def test_bulk_edit_audit_log_enabled_tags(self) -> None:
|
def test_bulk_edit_audit_log_enabled_tags(self) -> None:
|
||||||
@@ -2549,16 +2551,18 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
|
|||||||
GIVEN:
|
GIVEN:
|
||||||
- Audit log is enabled
|
- Audit log is enabled
|
||||||
WHEN:
|
WHEN:
|
||||||
- API to bulk edit tags is called
|
- API to bulk edit tags is called on an untagged document and a
|
||||||
|
document with several tags
|
||||||
THEN:
|
THEN:
|
||||||
- Audit log is created
|
- Audit log is created for each document with its full tag list
|
||||||
|
before and after the edit
|
||||||
"""
|
"""
|
||||||
LogEntry.objects.all().delete()
|
LogEntry.objects.all().delete()
|
||||||
response = self.client.post(
|
response = self.client.post(
|
||||||
"/api/documents/bulk_edit/",
|
"/api/documents/bulk_edit/",
|
||||||
json.dumps(
|
json.dumps(
|
||||||
{
|
{
|
||||||
"documents": [self.doc1.id],
|
"documents": [self.doc1.id, self.doc4.id],
|
||||||
"method": "modify_tags",
|
"method": "modify_tags",
|
||||||
"parameters": {
|
"parameters": {
|
||||||
"add_tags": [self.t1.id],
|
"add_tags": [self.t1.id],
|
||||||
@@ -2570,18 +2574,32 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
self.assertEqual(LogEntry.objects.filter(object_pk=self.doc1.id).count(), 1)
|
entry = LogEntry.objects.get_for_object(self.doc1).get()
|
||||||
|
self.assertEqual(entry.changes, {"tags": [[], [self.t1.id]]})
|
||||||
|
entry = LogEntry.objects.get_for_object(self.doc4).get()
|
||||||
|
self.assertEqual(
|
||||||
|
entry.changes,
|
||||||
|
{"tags": [[self.t1.id, self.t2.id], [self.t1.id]]},
|
||||||
|
)
|
||||||
|
|
||||||
@override_settings(AUDIT_LOG_ENABLED=True)
|
@override_settings(AUDIT_LOG_ENABLED=True)
|
||||||
def test_bulk_edit_audit_log_enabled_custom_fields(self) -> None:
|
def test_bulk_edit_audit_log_enabled_custom_fields(self) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
- Audit log is enabled
|
- Audit log is enabled
|
||||||
|
- A document with two custom fields
|
||||||
WHEN:
|
WHEN:
|
||||||
- API to bulk edit custom fields is called
|
- API to bulk edit custom fields is called to add a third
|
||||||
THEN:
|
THEN:
|
||||||
- Audit log is created
|
- Audit log is created with every custom field instance before and
|
||||||
|
after the edit
|
||||||
|
- Audit log is created for the new custom field instance
|
||||||
"""
|
"""
|
||||||
|
cf3 = CustomField.objects.create(name="cf3", data_type="string")
|
||||||
|
existing = [
|
||||||
|
CustomFieldInstance.objects.create(document=self.doc1, field=field)
|
||||||
|
for field in (self.cf2, cf3)
|
||||||
|
]
|
||||||
LogEntry.objects.all().delete()
|
LogEntry.objects.all().delete()
|
||||||
response = self.client.post(
|
response = self.client.post(
|
||||||
"/api/documents/bulk_edit/",
|
"/api/documents/bulk_edit/",
|
||||||
@@ -2599,7 +2617,14 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
self.assertEqual(LogEntry.objects.filter(object_pk=self.doc1.id).count(), 2)
|
added = CustomFieldInstance.objects.get(document=self.doc1, field=self.cf1)
|
||||||
|
existing_ids = [instance.id for instance in existing]
|
||||||
|
entry = LogEntry.objects.get_for_object(self.doc1).get()
|
||||||
|
self.assertEqual(
|
||||||
|
entry.changes,
|
||||||
|
{"custom_fields": [existing_ids, [*existing_ids, added.id]]},
|
||||||
|
)
|
||||||
|
self.assertEqual(LogEntry.objects.get_for_object(added).count(), 1)
|
||||||
|
|
||||||
def test_api_bulk_edit_with_bad_search_query_returns_400(self) -> None:
|
def test_api_bulk_edit_with_bad_search_query_returns_400(self) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from documents.models import Document
|
|||||||
from documents.versioning import annotate_effective_content
|
from documents.versioning import annotate_effective_content
|
||||||
from documents.views import DocumentSelectionMixin
|
from documents.views import DocumentSelectionMixin
|
||||||
from paperless_testing.dirs import DirectoriesMixin
|
from paperless_testing.dirs import DirectoriesMixin
|
||||||
|
from paperless_testing.factories import DocumentFactory
|
||||||
from paperless_testing.factories import UserFactory
|
from paperless_testing.factories import UserFactory
|
||||||
from paperless_testing.http import read_streaming_response
|
from paperless_testing.http import read_streaming_response
|
||||||
from paperless_testing.permissions import grant_global
|
from paperless_testing.permissions import grant_global
|
||||||
@@ -821,6 +822,26 @@ class TestDocumentVersioningApi(DirectoriesMixin, APITestCase):
|
|||||||
self.assertEqual(resp.status_code, status.HTTP_200_OK)
|
self.assertEqual(resp.status_code, status.HTTP_200_OK)
|
||||||
self.assertEqual(resp.data["content"], "v1-content")
|
self.assertEqual(resp.data["content"], "v1-content")
|
||||||
|
|
||||||
|
def test_page_count_resolves_to_latest_version(self) -> None:
|
||||||
|
root = DocumentFactory(page_count=2)
|
||||||
|
DocumentFactory(root_document=root, version_index=1, page_count=1)
|
||||||
|
unversioned = DocumentFactory(page_count=5)
|
||||||
|
|
||||||
|
resp = self.client.get("/api/documents/?fields=id,page_count")
|
||||||
|
self.assertEqual(resp.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(
|
||||||
|
{doc["id"]: doc["page_count"] for doc in resp.data["results"]},
|
||||||
|
{root.id: 1, unversioned.id: 5},
|
||||||
|
)
|
||||||
|
|
||||||
|
resp = self.client.get(f"/api/documents/{root.id}/")
|
||||||
|
self.assertEqual(resp.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(resp.data["page_count"], 1)
|
||||||
|
|
||||||
|
resp = self.client.get(f"/api/documents/{root.id}/?version={root.id}")
|
||||||
|
self.assertEqual(resp.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(resp.data["page_count"], 2)
|
||||||
|
|
||||||
def _make_root_with_out_of_order_versions(self) -> tuple[Document, ...]:
|
def _make_root_with_out_of_order_versions(self) -> tuple[Document, ...]:
|
||||||
"""
|
"""
|
||||||
A root whose newest version has a *lower* id than an older one, which is
|
A root whose newest version has a *lower* id than an older one, which is
|
||||||
|
|||||||
@@ -64,16 +64,15 @@ class TestApiProfile(DirectoriesMixin, APITestCase):
|
|||||||
)
|
)
|
||||||
self.client.force_authenticate(user=self.user)
|
self.client.force_authenticate(user=self.user)
|
||||||
|
|
||||||
def setupSocialAccount(self) -> None:
|
def setupSocialAccount(self) -> SocialAccount:
|
||||||
SocialApp.objects.create(
|
SocialApp.objects.create(
|
||||||
name="Keycloak",
|
name="Keycloak",
|
||||||
provider="openid_connect",
|
provider="openid_connect",
|
||||||
provider_id="keycloak-test",
|
provider_id="keycloak-test",
|
||||||
)
|
)
|
||||||
self.user.socialaccount_set.add(
|
social_account = SocialAccount(uid="123456789", provider="keycloak-test")
|
||||||
SocialAccount(uid="123456789", provider="keycloak-test"),
|
self.user.socialaccount_set.add(social_account, bulk=False)
|
||||||
bulk=False,
|
return social_account
|
||||||
)
|
|
||||||
|
|
||||||
def test_get_profile(self) -> None:
|
def test_get_profile(self) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -111,19 +110,17 @@ class TestApiProfile(DirectoriesMixin, APITestCase):
|
|||||||
THEN:
|
THEN:
|
||||||
- Profile is returned with social accounts
|
- Profile is returned with social accounts
|
||||||
"""
|
"""
|
||||||
self.setupSocialAccount()
|
social_account = self.setupSocialAccount()
|
||||||
|
|
||||||
openid_provider = (
|
openid_provider = MockOpenIDConnectProvider(
|
||||||
MockOpenIDConnectProvider(
|
app=SocialApp.objects.get(provider_id="keycloak-test"),
|
||||||
app=SocialApp.objects.get(provider_id="keycloak-test"),
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
mock_list_providers.return_value = [
|
mock_list_providers.return_value = [
|
||||||
openid_provider,
|
openid_provider,
|
||||||
]
|
]
|
||||||
mock_get_provider_account.return_value = MockOpenIDConnectProviderAccount(
|
mock_get_provider_account.return_value = MockOpenIDConnectProviderAccount(
|
||||||
mock_social_account_dict={
|
mock_social_account_dict={
|
||||||
"name": openid_provider[0].name,
|
"name": openid_provider.name,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -135,7 +132,7 @@ class TestApiProfile(DirectoriesMixin, APITestCase):
|
|||||||
response.data["social_accounts"],
|
response.data["social_accounts"],
|
||||||
[
|
[
|
||||||
{
|
{
|
||||||
"id": 1,
|
"id": social_account.pk,
|
||||||
"provider": "keycloak-test",
|
"provider": "keycloak-test",
|
||||||
"name": "Keycloak",
|
"name": "Keycloak",
|
||||||
},
|
},
|
||||||
@@ -152,7 +149,7 @@ class TestApiProfile(DirectoriesMixin, APITestCase):
|
|||||||
THEN:
|
THEN:
|
||||||
- Profile is returned with "Unknown App" as name
|
- Profile is returned with "Unknown App" as name
|
||||||
"""
|
"""
|
||||||
self.setupSocialAccount()
|
social_account = self.setupSocialAccount()
|
||||||
|
|
||||||
# Remove the social app
|
# Remove the social app
|
||||||
SocialApp.objects.get(provider_id="keycloak-test").delete()
|
SocialApp.objects.get(provider_id="keycloak-test").delete()
|
||||||
@@ -165,7 +162,7 @@ class TestApiProfile(DirectoriesMixin, APITestCase):
|
|||||||
response.data["social_accounts"],
|
response.data["social_accounts"],
|
||||||
[
|
[
|
||||||
{
|
{
|
||||||
"id": 1,
|
"id": social_account.pk,
|
||||||
"provider": "keycloak-test",
|
"provider": "keycloak-test",
|
||||||
"name": "Unknown App",
|
"name": "Unknown App",
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -677,12 +677,13 @@ class TestExportImport(
|
|||||||
THEN:
|
THEN:
|
||||||
- Error is raised
|
- Error is raised
|
||||||
"""
|
"""
|
||||||
args = ["document_exporter", "/tmp/foo/bar"]
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
|
args = ["document_exporter", str(Path(tmp_dir) / "does-not-exist")]
|
||||||
|
|
||||||
with self.assertRaises(CommandError) as e:
|
with self.assertRaises(CommandError) as e:
|
||||||
call_command(*args, skip_checks=True)
|
call_command(*args, skip_checks=True)
|
||||||
|
|
||||||
self.assertEqual("That path doesn't exist", str(e.exception))
|
self.assertEqual("That path doesn't exist", str(e.exception))
|
||||||
|
|
||||||
def test_export_target_exists_but_is_file(self) -> None:
|
def test_export_target_exists_but_is_file(self) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -123,14 +123,14 @@ class TestFuzzyMatchCommand(TestCase):
|
|||||||
- Output contains clickable links to the documents instead of titles
|
- Output contains clickable links to the documents instead of titles
|
||||||
"""
|
"""
|
||||||
# Content similarity is 86.667
|
# Content similarity is 86.667
|
||||||
Document.objects.create(
|
doc1 = Document.objects.create(
|
||||||
checksum="BEEFCAFE",
|
checksum="BEEFCAFE",
|
||||||
title="A",
|
title="A",
|
||||||
content="first document scanned by bob",
|
content="first document scanned by bob",
|
||||||
mime_type="application/pdf",
|
mime_type="application/pdf",
|
||||||
filename="test.pdf",
|
filename="test.pdf",
|
||||||
)
|
)
|
||||||
Document.objects.create(
|
doc2 = Document.objects.create(
|
||||||
checksum="DEADBEAF",
|
checksum="DEADBEAF",
|
||||||
title="A",
|
title="A",
|
||||||
content="first document scanned by alice",
|
content="first document scanned by alice",
|
||||||
@@ -145,8 +145,8 @@ class TestFuzzyMatchCommand(TestCase):
|
|||||||
"http://localhost:8000",
|
"http://localhost:8000",
|
||||||
)
|
)
|
||||||
self.assertIn("Found 1 matching pair(s)", stdout)
|
self.assertIn("Found 1 matching pair(s)", stdout)
|
||||||
self.assertIn("http://localhost:8000/documents/1/details", stdout)
|
self.assertIn(f"http://localhost:8000/documents/{doc1.pk}/details", stdout)
|
||||||
self.assertIn("http://localhost:8000/documents/2/details", stdout)
|
self.assertIn(f"http://localhost:8000/documents/{doc2.pk}/details", stdout)
|
||||||
|
|
||||||
def test_with_3_matches(self) -> None:
|
def test_with_3_matches(self) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -198,14 +198,14 @@ class TestFuzzyMatchCommand(TestCase):
|
|||||||
- Documents 1 and 2 remain
|
- Documents 1 and 2 remain
|
||||||
"""
|
"""
|
||||||
# Content similarity is 86.667
|
# Content similarity is 86.667
|
||||||
Document.objects.create(
|
doc1 = Document.objects.create(
|
||||||
checksum="BEEFCAFE",
|
checksum="BEEFCAFE",
|
||||||
title="A",
|
title="A",
|
||||||
content="first document scanned by bob",
|
content="first document scanned by bob",
|
||||||
mime_type="application/pdf",
|
mime_type="application/pdf",
|
||||||
filename="test.pdf",
|
filename="test.pdf",
|
||||||
)
|
)
|
||||||
Document.objects.create(
|
doc2 = Document.objects.create(
|
||||||
checksum="DEADBEAF",
|
checksum="DEADBEAF",
|
||||||
title="A",
|
title="A",
|
||||||
content="second document scanned by alice",
|
content="second document scanned by alice",
|
||||||
@@ -235,8 +235,8 @@ class TestFuzzyMatchCommand(TestCase):
|
|||||||
self.assertIn("Deleting 1 document(s)", stdout)
|
self.assertIn("Deleting 1 document(s)", stdout)
|
||||||
|
|
||||||
self.assertEqual(Document.objects.count(), 2)
|
self.assertEqual(Document.objects.count(), 2)
|
||||||
self.assertIsNotNone(Document.objects.get(pk=1))
|
self.assertIsNotNone(Document.objects.get(pk=doc1.pk))
|
||||||
self.assertIsNotNone(Document.objects.get(pk=2))
|
self.assertIsNotNone(Document.objects.get(pk=doc2.pk))
|
||||||
|
|
||||||
def test_document_deletion_cancelled(self) -> None:
|
def test_document_deletion_cancelled(self) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from paperless_testing.dirs import DirectoriesMixin
|
|||||||
class TestManageSuperUser(DirectoriesMixin, TestCase):
|
class TestManageSuperUser(DirectoriesMixin, TestCase):
|
||||||
def call_command(self, environ):
|
def call_command(self, environ):
|
||||||
out = StringIO()
|
out = StringIO()
|
||||||
with mock.patch.dict(os.environ, environ):
|
with mock.patch.dict(os.environ, environ, clear=True):
|
||||||
call_command(
|
call_command(
|
||||||
"manage_superuser",
|
"manage_superuser",
|
||||||
"--no-color",
|
"--no-color",
|
||||||
|
|||||||
@@ -430,6 +430,53 @@ class TestBulkDownloadPermissionChecksRootDocument:
|
|||||||
) # version-only grant must not substitute for root permission
|
) # version-only grant must not substitute for root permission
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.django_db
|
||||||
|
class TestDocumentOperationPermissionChecksRootDocument:
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("endpoint", "payload"),
|
||||||
|
[
|
||||||
|
pytest.param("/api/documents/merge/", {}, id="merge"),
|
||||||
|
pytest.param("/api/documents/rotate/", {"degrees": 90}, id="rotate"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
@pytest.mark.parametrize("version_owner", ["none", "requester"])
|
||||||
|
def test_version_operation_acts_on_root(
|
||||||
|
self,
|
||||||
|
rest_api_client: APIClient,
|
||||||
|
endpoint: str,
|
||||||
|
payload: dict,
|
||||||
|
version_owner: str,
|
||||||
|
) -> None:
|
||||||
|
owner = UserFactory(username="owner")
|
||||||
|
requester = UserFactory(username="requester")
|
||||||
|
grant_global(requester, "change_document")
|
||||||
|
grant_global(requester, "add_document")
|
||||||
|
rest_api_client.force_authenticate(user=requester)
|
||||||
|
root = DocumentFactory(owner=owner)
|
||||||
|
# A version whose owner went stale, e.g. created before the root changed hands
|
||||||
|
version = DocumentFactory(
|
||||||
|
owner=requester if version_owner == "requester" else None,
|
||||||
|
root_document=root,
|
||||||
|
version_index=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("documents.views.bulk_edit.merge") as mock_merge,
|
||||||
|
patch("documents.views.bulk_edit.rotate") as mock_rotate,
|
||||||
|
):
|
||||||
|
mock_merge.__name__ = "merge"
|
||||||
|
mock_rotate.__name__ = "rotate"
|
||||||
|
response = rest_api_client.post(
|
||||||
|
endpoint,
|
||||||
|
{"documents": [version.pk], **payload},
|
||||||
|
format="json",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == HTTPStatus.FORBIDDEN
|
||||||
|
mock_merge.assert_not_called()
|
||||||
|
mock_rotate.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.django_db
|
@pytest.mark.django_db
|
||||||
@pytest.mark.usefixtures("_search_index")
|
@pytest.mark.usefixtures("_search_index")
|
||||||
class TestTrashRestorePermissionBoundary:
|
class TestTrashRestorePermissionBoundary:
|
||||||
|
|||||||
@@ -339,15 +339,6 @@ class ShareLinkBundleBuildTaskTests(DirectoriesMixin, APITestCase):
|
|||||||
)
|
)
|
||||||
self.document.archive_checksum = ""
|
self.document.archive_checksum = ""
|
||||||
self.document.save()
|
self.document.save()
|
||||||
self.addCleanup(
|
|
||||||
setattr,
|
|
||||||
settings,
|
|
||||||
"SHARE_LINK_BUNDLE_DIR",
|
|
||||||
settings.SHARE_LINK_BUNDLE_DIR,
|
|
||||||
)
|
|
||||||
settings.SHARE_LINK_BUNDLE_DIR = (
|
|
||||||
Path(settings.MEDIA_ROOT) / "documents" / "share_link_bundles"
|
|
||||||
)
|
|
||||||
|
|
||||||
def _write_document_file(self, *, archive: bool, content: bytes) -> Path:
|
def _write_document_file(self, *, archive: bool, content: bytes) -> Path:
|
||||||
if archive:
|
if archive:
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ from documents.models import WorkflowAction
|
|||||||
from documents.sanity_checker import SanityCheckFailedException
|
from documents.sanity_checker import SanityCheckFailedException
|
||||||
from documents.sanity_checker import SanityCheckMessages
|
from documents.sanity_checker import SanityCheckMessages
|
||||||
from documents.tests.helpers import dummy_preprocess
|
from documents.tests.helpers import dummy_preprocess
|
||||||
from paperless_ai.exceptions import LLMBlockedError
|
|
||||||
from paperless_testing.assertions import FileSystemAssertsMixin
|
from paperless_testing.assertions import FileSystemAssertsMixin
|
||||||
from paperless_testing.dirs import DirectoriesMixin
|
from paperless_testing.dirs import DirectoriesMixin
|
||||||
|
|
||||||
@@ -556,37 +555,3 @@ class TestApplyAISuggestionsTask(DirectoriesMixin, TestCase):
|
|||||||
|
|
||||||
apply_suggestions.assert_not_called()
|
apply_suggestions.assert_not_called()
|
||||||
self.assertIn("no longer exists", "".join(cm.output))
|
self.assertIn("no longer exists", "".join(cm.output))
|
||||||
|
|
||||||
@override_settings(AI_ENABLED=True)
|
|
||||||
def test_blocked_request_fails_without_retry(self) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- AI enabled and a document with content
|
|
||||||
- The AI classification call blocked by the outbound request policy
|
|
||||||
WHEN:
|
|
||||||
- The task runs through Celery
|
|
||||||
THEN:
|
|
||||||
- The workflow code does not swallow the block
|
|
||||||
- The task fails with LLMBlockedError and is never retried
|
|
||||||
"""
|
|
||||||
with (
|
|
||||||
mock.patch(
|
|
||||||
"documents.workflows.ai.get_ai_document_classification",
|
|
||||||
side_effect=LLMBlockedError(
|
|
||||||
"AI backend request was blocked by the outbound request "
|
|
||||||
"policy: detail",
|
|
||||||
),
|
|
||||||
),
|
|
||||||
mock.patch.object(
|
|
||||||
tasks.apply_ai_suggestions,
|
|
||||||
"retry",
|
|
||||||
wraps=tasks.apply_ai_suggestions.retry,
|
|
||||||
) as retry,
|
|
||||||
):
|
|
||||||
result = tasks.apply_ai_suggestions.apply(
|
|
||||||
args=(self.action.pk, self.doc.pk),
|
|
||||||
)
|
|
||||||
|
|
||||||
self.assertTrue(result.failed())
|
|
||||||
self.assertIsInstance(result.result, LLMBlockedError)
|
|
||||||
retry.assert_not_called()
|
|
||||||
|
|||||||
@@ -29,7 +29,6 @@ from documents.models import Tag
|
|||||||
from documents.models import UiSettings
|
from documents.models import UiSettings
|
||||||
from documents.signals.handlers import update_llm_suggestions_cache
|
from documents.signals.handlers import update_llm_suggestions_cache
|
||||||
from paperless.models import ApplicationConfiguration
|
from paperless.models import ApplicationConfiguration
|
||||||
from paperless_ai.exceptions import LLMBlockedError
|
|
||||||
from paperless_ai.exceptions import LLMProviderError
|
from paperless_ai.exceptions import LLMProviderError
|
||||||
from paperless_ai.exceptions import LLMTimeoutError
|
from paperless_ai.exceptions import LLMTimeoutError
|
||||||
from paperless_testing.dirs import DirectoriesMixin
|
from paperless_testing.dirs import DirectoriesMixin
|
||||||
@@ -771,48 +770,6 @@ class TestAISuggestions(DirectoriesMixin, TestCase):
|
|||||||
get_llm_suggestion_cache(self.document.pk, backend="openai-like"),
|
get_llm_suggestion_cache(self.document.pk, backend="openai-like"),
|
||||||
)
|
)
|
||||||
|
|
||||||
@patch("documents.views.get_ai_document_classification")
|
|
||||||
@override_settings(
|
|
||||||
AI_ENABLED=True,
|
|
||||||
LLM_BACKEND="openai-like",
|
|
||||||
)
|
|
||||||
def test_ai_suggestions_with_blocked_llm_request(
|
|
||||||
self,
|
|
||||||
mock_get_ai_classification,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- An AI backend request blocked by the outbound request policy
|
|
||||||
WHEN:
|
|
||||||
- AI suggestions are requested
|
|
||||||
THEN:
|
|
||||||
- 502 is returned with a generic message and nothing is cached
|
|
||||||
"""
|
|
||||||
mock_get_ai_classification.side_effect = LLMBlockedError(
|
|
||||||
"AI backend request was blocked by the outbound request policy: detail",
|
|
||||||
)
|
|
||||||
|
|
||||||
self.client.force_login(user=self.user)
|
|
||||||
response = self.client.get(
|
|
||||||
f"/api/documents/{self.document.pk}/ai_suggestions/",
|
|
||||||
)
|
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_502_BAD_GATEWAY)
|
|
||||||
self.assertEqual(
|
|
||||||
response.json(),
|
|
||||||
{
|
|
||||||
"ai": [
|
|
||||||
(
|
|
||||||
"AI backend request was blocked by the outbound request "
|
|
||||||
"policy. Check logs for details."
|
|
||||||
),
|
|
||||||
],
|
|
||||||
},
|
|
||||||
)
|
|
||||||
self.assertIsNone(
|
|
||||||
get_llm_suggestion_cache(self.document.pk, backend="openai-like"),
|
|
||||||
)
|
|
||||||
|
|
||||||
@patch("documents.views.get_ai_document_classification")
|
@patch("documents.views.get_ai_document_classification")
|
||||||
@override_settings(
|
@override_settings(
|
||||||
AI_ENABLED=True,
|
AI_ENABLED=True,
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
import datetime
|
import datetime
|
||||||
import json
|
import json
|
||||||
import shutil
|
import shutil
|
||||||
|
import socket
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from collections.abc import Callable
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
@@ -17,11 +19,11 @@ from django.test import override_settings
|
|||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from guardian.shortcuts import get_groups_with_perms
|
from guardian.shortcuts import get_groups_with_perms
|
||||||
from guardian.shortcuts import get_users_with_perms
|
from guardian.shortcuts import get_users_with_perms
|
||||||
|
from httpx import ConnectError
|
||||||
from httpx import HTTPError
|
from httpx import HTTPError
|
||||||
from httpx import HTTPStatusError
|
from httpx import HTTPStatusError
|
||||||
from pytest_django.fixtures import Settings
|
from pytest_django.fixtures import Settings
|
||||||
from pytest_httpx import HTTPXMock
|
from pytest_httpx import HTTPXMock
|
||||||
from pytest_mock import MockerFixture
|
|
||||||
from rest_framework.test import APIClient
|
from rest_framework.test import APIClient
|
||||||
from rest_framework.test import APITestCase
|
from rest_framework.test import APITestCase
|
||||||
|
|
||||||
@@ -31,12 +33,8 @@ from documents.file_handling import generate_unique_filename
|
|||||||
from documents.signals.handlers import run_workflows
|
from documents.signals.handlers import run_workflows
|
||||||
from documents.workflows.ai import apply_ai_suggestions_to_document
|
from documents.workflows.ai import apply_ai_suggestions_to_document
|
||||||
from documents.workflows.webhooks import send_webhook
|
from documents.workflows.webhooks import send_webhook
|
||||||
from paperless.network import OutboundRequestBlockedError
|
|
||||||
from paperless_ai.base_model import ClassificationSuggestions
|
from paperless_ai.base_model import ClassificationSuggestions
|
||||||
from paperless_ai.exceptions import LLMTimeoutError
|
from paperless_ai.exceptions import LLMTimeoutError
|
||||||
from paperless_testing.outbound import DialRecorder
|
|
||||||
from paperless_testing.outbound import FakeDNS
|
|
||||||
from paperless_testing.outbound import LocalHTTPServer
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from django.db.models import QuerySet
|
from django.db.models import QuerySet
|
||||||
@@ -71,7 +69,9 @@ from paperless_mail.models import MailAccount
|
|||||||
from paperless_mail.models import MailRule
|
from paperless_mail.models import MailRule
|
||||||
from paperless_testing.assertions import FileSystemAssertsMixin
|
from paperless_testing.assertions import FileSystemAssertsMixin
|
||||||
from paperless_testing.dirs import DirectoriesMixin
|
from paperless_testing.dirs import DirectoriesMixin
|
||||||
|
from paperless_testing.factories import DocumentFactory
|
||||||
from paperless_testing.factories import UserFactory
|
from paperless_testing.factories import UserFactory
|
||||||
|
from paperless_testing.permissions import grant_global
|
||||||
from paperless_testing.permissions import grant_object
|
from paperless_testing.permissions import grant_object
|
||||||
|
|
||||||
|
|
||||||
@@ -1061,6 +1061,41 @@ class TestWorkflows(
|
|||||||
self.assertEqual(doc.correspondent, self.c2)
|
self.assertEqual(doc.correspondent, self.c2)
|
||||||
self.assertEqual(doc.title, f"Doc created in {created.year}")
|
self.assertEqual(doc.title, f"Doc created in {created.year}")
|
||||||
|
|
||||||
|
@pytest.mark.usefixtures("_search_index")
|
||||||
|
def test_document_added_workflow_indexes_final_title(self) -> None:
|
||||||
|
trigger = WorkflowTrigger.objects.create(
|
||||||
|
type=WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
|
||||||
|
filter_filename="*sample*",
|
||||||
|
)
|
||||||
|
action = WorkflowAction.objects.create(
|
||||||
|
assign_title="Linked document",
|
||||||
|
assign_owner=self.user2,
|
||||||
|
)
|
||||||
|
link_field = CustomField.objects.create(
|
||||||
|
name="Related documents",
|
||||||
|
data_type=CustomField.FieldDataType.DOCUMENTLINK,
|
||||||
|
)
|
||||||
|
action.assign_custom_fields.add(link_field)
|
||||||
|
workflow = Workflow.objects.create(name="Link workflow", order=0)
|
||||||
|
workflow.triggers.add(trigger)
|
||||||
|
workflow.actions.add(action)
|
||||||
|
|
||||||
|
doc = DocumentFactory.create()
|
||||||
|
document_consumption_finished.send(sender=self.__class__, document=doc)
|
||||||
|
|
||||||
|
self.assertTrue(doc.custom_fields.filter(field=link_field).exists())
|
||||||
|
doc.refresh_from_db()
|
||||||
|
self.assertEqual(doc.title, "Linked document")
|
||||||
|
|
||||||
|
grant_global(self.user2, "view_document")
|
||||||
|
self.client.force_authenticate(user=self.user2)
|
||||||
|
response = self.client.get("/api/documents/?title_search=linked")
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertEqual(
|
||||||
|
[result["id"] for result in response.data["results"]],
|
||||||
|
[doc.pk],
|
||||||
|
)
|
||||||
|
|
||||||
def test_document_added_no_match_filename(self) -> None:
|
def test_document_added_no_match_filename(self) -> None:
|
||||||
trigger = WorkflowTrigger.objects.create(
|
trigger = WorkflowTrigger.objects.create(
|
||||||
type=WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
|
type=WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
|
||||||
@@ -4424,18 +4459,18 @@ class TestWorkflows(
|
|||||||
)
|
)
|
||||||
|
|
||||||
@mock.patch("documents.bulk_edit.remove_password")
|
@mock.patch("documents.bulk_edit.remove_password")
|
||||||
def test_password_removal_action_fails_without_correct_password(
|
def test_password_removal_action_skips_blank_and_whitespace_passwords(
|
||||||
self,
|
self,
|
||||||
mock_remove_password,
|
mock_remove_password,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
- Workflow password removal action
|
- Workflow password removal action
|
||||||
- No correct password provided
|
- Only blank and whitespace-only passwords configured
|
||||||
WHEN:
|
WHEN:
|
||||||
- Document updated triggering the workflow
|
- Document updated triggering the workflow
|
||||||
THEN:
|
THEN:
|
||||||
- Password removal is attempted for all passwords and fails
|
- Password removal is not attempted
|
||||||
"""
|
"""
|
||||||
doc = Document.objects.create(
|
doc = Document.objects.create(
|
||||||
title="Protected",
|
title="Protected",
|
||||||
@@ -4456,6 +4491,60 @@ class TestWorkflows(
|
|||||||
|
|
||||||
mock_remove_password.assert_not_called()
|
mock_remove_password.assert_not_called()
|
||||||
|
|
||||||
|
@mock.patch("documents.bulk_edit.remove_password")
|
||||||
|
def test_password_removal_action_fails_without_correct_password(
|
||||||
|
self,
|
||||||
|
mock_remove_password,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
GIVEN:
|
||||||
|
- Workflow password removal action
|
||||||
|
- No configured password is correct
|
||||||
|
WHEN:
|
||||||
|
- Document updated triggering the workflow
|
||||||
|
THEN:
|
||||||
|
- Password removal is attempted for every configured password and fails
|
||||||
|
"""
|
||||||
|
doc = Document.objects.create(
|
||||||
|
title="Protected",
|
||||||
|
checksum="pw-checksum-3",
|
||||||
|
)
|
||||||
|
trigger = WorkflowTrigger.objects.create(
|
||||||
|
type=WorkflowTrigger.WorkflowTriggerType.DOCUMENT_UPDATED,
|
||||||
|
)
|
||||||
|
action = WorkflowAction.objects.create(
|
||||||
|
type=WorkflowAction.WorkflowActionType.PASSWORD_REMOVAL,
|
||||||
|
passwords=["wrong", "also-wrong"],
|
||||||
|
)
|
||||||
|
workflow = Workflow.objects.create(name="Password workflow wrong passwords")
|
||||||
|
workflow.triggers.add(trigger)
|
||||||
|
workflow.actions.add(action)
|
||||||
|
|
||||||
|
mock_remove_password.side_effect = ValueError("wrong password")
|
||||||
|
|
||||||
|
with self.assertLogs("paperless.workflows.actions", level="ERROR"):
|
||||||
|
run_workflows(trigger.type, doc)
|
||||||
|
|
||||||
|
assert mock_remove_password.call_count == 2
|
||||||
|
mock_remove_password.assert_has_calls(
|
||||||
|
[
|
||||||
|
mock.call(
|
||||||
|
[doc.id],
|
||||||
|
password="wrong",
|
||||||
|
update_document=True,
|
||||||
|
user=doc.owner,
|
||||||
|
source_paths_by_id=None,
|
||||||
|
),
|
||||||
|
mock.call(
|
||||||
|
[doc.id],
|
||||||
|
password="also-wrong",
|
||||||
|
update_document=True,
|
||||||
|
user=doc.owner,
|
||||||
|
source_paths_by_id=None,
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
@mock.patch("documents.bulk_edit.remove_password")
|
@mock.patch("documents.bulk_edit.remove_password")
|
||||||
def test_password_removal_action_skips_without_passwords(
|
def test_password_removal_action_skips_without_passwords(
|
||||||
self,
|
self,
|
||||||
@@ -5071,6 +5160,25 @@ class TestWebhookSend:
|
|||||||
assert httpx_mock.get_request().headers["Content-Type"] == "application/json"
|
assert httpx_mock.get_request().headers["Content-Type"] == "application/json"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def resolve_to(monkeypatch: pytest.MonkeyPatch) -> Callable[[str], None]:
|
||||||
|
"""
|
||||||
|
Force DNS resolution to a specific IP for any hostname.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _set(ip: str) -> None:
|
||||||
|
def fake_getaddrinfo(
|
||||||
|
host: str,
|
||||||
|
*_args: object,
|
||||||
|
**_kwargs: object,
|
||||||
|
) -> list[tuple[Any, ...]]:
|
||||||
|
return [(socket.AF_INET, None, None, "", (ip, 0))]
|
||||||
|
|
||||||
|
monkeypatch.setattr(socket, "getaddrinfo", fake_getaddrinfo)
|
||||||
|
|
||||||
|
return _set
|
||||||
|
|
||||||
|
|
||||||
class TestWebhookSecurity:
|
class TestWebhookSecurity:
|
||||||
def test_blocks_invalid_scheme_or_hostname(self, httpx_mock: HTTPXMock) -> None:
|
def test_blocks_invalid_scheme_or_hostname(self, httpx_mock: HTTPXMock) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -5120,145 +5228,60 @@ class TestWebhookSecurity:
|
|||||||
|
|
||||||
assert httpx_mock.get_request() is None
|
assert httpx_mock.get_request() is None
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"address",
|
|
||||||
[
|
|
||||||
pytest.param("127.0.0.1", id="loopback"),
|
|
||||||
pytest.param("10.0.0.1", id="private"),
|
|
||||||
pytest.param("169.254.169.254", id="link-local-metadata"),
|
|
||||||
pytest.param("::ffff:127.0.0.1", id="ipv4-mapped-loopback"),
|
|
||||||
pytest.param("64:ff9b::7f00:1", id="nat64-wrapping-loopback"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
@override_settings(WEBHOOKS_ALLOW_INTERNAL_REQUESTS=False)
|
@override_settings(WEBHOOKS_ALLOW_INTERNAL_REQUESTS=False)
|
||||||
def test_blocks_private_loopback_linklocal(
|
def test_blocks_private_loopback_linklocal(
|
||||||
self,
|
self,
|
||||||
local_http_server: LocalHTTPServer,
|
httpx_mock: HTTPXMock,
|
||||||
fake_dns: FakeDNS,
|
resolve_to,
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
address: str,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
- A webhook host resolving to a non-public address
|
- URL with a private, loopback, or link-local IP address
|
||||||
- WEBHOOKS_ALLOW_INTERNAL_REQUESTS is False
|
- WEBHOOKS_ALLOW_INTERNAL_REQUESTS is False
|
||||||
WHEN:
|
WHEN:
|
||||||
- send_webhook is called
|
- send_webhook is called with such URL
|
||||||
THEN:
|
THEN:
|
||||||
- The request is blocked before any connection is opened
|
- ValueError is raised
|
||||||
"""
|
"""
|
||||||
fake_dns.add("webhook.test", address)
|
resolve_to("127.0.0.1")
|
||||||
|
with pytest.raises(ConnectError):
|
||||||
with pytest.raises(OutboundRequestBlockedError):
|
|
||||||
send_webhook(
|
send_webhook(
|
||||||
f"http://webhook.test:{local_http_server.port}",
|
"http://paperless-ngx.com",
|
||||||
data="",
|
data="",
|
||||||
headers={},
|
headers={},
|
||||||
files=None,
|
files=None,
|
||||||
as_json=False,
|
as_json=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert local_http_server.connections == 0
|
def test_allows_public_ip_and_sends(
|
||||||
assert dial_recorder.hosts() == []
|
|
||||||
|
|
||||||
@override_settings(WEBHOOKS_ALLOW_INTERNAL_REQUESTS=False)
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_sends_to_validated_address(
|
|
||||||
self,
|
self,
|
||||||
local_http_server: LocalHTTPServer,
|
httpx_mock: HTTPXMock,
|
||||||
fake_dns: FakeDNS,
|
resolve_to,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
- A webhook host resolving to an address the policy accepts
|
- URL with a public IP address
|
||||||
- WEBHOOKS_ALLOW_INTERNAL_REQUESTS is False
|
|
||||||
WHEN:
|
WHEN:
|
||||||
- send_webhook is called
|
- send_webhook is called with such URL
|
||||||
THEN:
|
THEN:
|
||||||
- The payload arrives with the webhook hostname in the Host header
|
- Request is sent successfully
|
||||||
"""
|
"""
|
||||||
fake_dns.add("webhook.test", "127.0.0.1")
|
resolve_to("52.207.186.75")
|
||||||
|
httpx_mock.add_response(content=b"ok")
|
||||||
|
|
||||||
send_webhook(
|
send_webhook(
|
||||||
url=f"http://webhook.test:{local_http_server.port}",
|
url="http://paperless-ngx.com",
|
||||||
data="hi",
|
data="hi",
|
||||||
headers={},
|
headers={},
|
||||||
files=None,
|
files=None,
|
||||||
as_json=False,
|
as_json=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
received = local_http_server.requests[0]
|
req = httpx_mock.get_request()
|
||||||
assert received.body == b"hi"
|
assert req.url.host == "52.207.186.75"
|
||||||
assert received.headers["host"] == f"webhook.test:{local_http_server.port}"
|
assert req.headers["host"] == "paperless-ngx.com"
|
||||||
|
|
||||||
@override_settings(WEBHOOKS_ALLOW_INTERNAL_REQUESTS=True)
|
def test_follow_redirects_disabled(self, httpx_mock: HTTPXMock, resolve_to) -> None:
|
||||||
def test_allow_internal_sends_to_internal_address(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A webhook to localhost
|
|
||||||
- WEBHOOKS_ALLOW_INTERNAL_REQUESTS is True
|
|
||||||
WHEN:
|
|
||||||
- send_webhook is called
|
|
||||||
THEN:
|
|
||||||
- The payload arrives at the internal address
|
|
||||||
- The guard does not resolve the host, leaving it to the stock
|
|
||||||
connection path
|
|
||||||
"""
|
|
||||||
send_webhook(
|
|
||||||
url=f"http://localhost:{local_http_server.port}",
|
|
||||||
data="hi",
|
|
||||||
headers={},
|
|
||||||
files=None,
|
|
||||||
as_json=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
received = local_http_server.requests[0]
|
|
||||||
assert received.body == b"hi"
|
|
||||||
assert fake_dns.lookups == []
|
|
||||||
|
|
||||||
@override_settings(WEBHOOKS_ALLOW_INTERNAL_REQUESTS=False)
|
|
||||||
def test_block_is_an_expected_task_failure(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A webhook host resolving to a loopback address
|
|
||||||
- WEBHOOKS_ALLOW_INTERNAL_REQUESTS is False
|
|
||||||
WHEN:
|
|
||||||
- The webhook task runs through Celery
|
|
||||||
THEN:
|
|
||||||
- The task fails with the original block error, not a wrapper,
|
|
||||||
so it matches the task's expected errors, and is not retried
|
|
||||||
"""
|
|
||||||
fake_dns.add("webhook.test", "127.0.0.1")
|
|
||||||
retry = mocker.spy(send_webhook, "retry")
|
|
||||||
|
|
||||||
result = send_webhook.apply(
|
|
||||||
kwargs={
|
|
||||||
"url": f"http://webhook.test:{local_http_server.port}",
|
|
||||||
"data": "",
|
|
||||||
"headers": {},
|
|
||||||
"files": None,
|
|
||||||
"as_json": False,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result.failed()
|
|
||||||
assert isinstance(result.result, OutboundRequestBlockedError)
|
|
||||||
assert isinstance(result.result, send_webhook.throws)
|
|
||||||
retry.assert_not_called()
|
|
||||||
assert local_http_server.connections == 0
|
|
||||||
assert dial_recorder.hosts() == []
|
|
||||||
|
|
||||||
def test_follow_redirects_disabled(self, httpx_mock: HTTPXMock) -> None:
|
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
- A URL that redirects
|
- A URL that redirects
|
||||||
@@ -5267,6 +5290,7 @@ class TestWebhookSecurity:
|
|||||||
THEN:
|
THEN:
|
||||||
- Request is made to the original URL and does not follow the redirect
|
- Request is made to the original URL and does not follow the redirect
|
||||||
"""
|
"""
|
||||||
|
resolve_to("52.207.186.75")
|
||||||
# Return a redirect and ensure we don't follow it (only one request recorded)
|
# Return a redirect and ensure we don't follow it (only one request recorded)
|
||||||
httpx_mock.add_response(
|
httpx_mock.add_response(
|
||||||
status_code=302,
|
status_code=302,
|
||||||
@@ -5288,6 +5312,7 @@ class TestWebhookSecurity:
|
|||||||
def test_strips_user_supplied_host_header(
|
def test_strips_user_supplied_host_header(
|
||||||
self,
|
self,
|
||||||
httpx_mock: HTTPXMock,
|
httpx_mock: HTTPXMock,
|
||||||
|
resolve_to: Callable[[str], None],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
@@ -5295,8 +5320,9 @@ class TestWebhookSecurity:
|
|||||||
WHEN:
|
WHEN:
|
||||||
- send_webhook is called with a malicious Host header
|
- send_webhook is called with a malicious Host header
|
||||||
THEN:
|
THEN:
|
||||||
- The Host header is stripped and set from the URL hostname
|
- The Host header is stripped and replaced with the resolved hostname
|
||||||
"""
|
"""
|
||||||
|
resolve_to("52.207.186.75")
|
||||||
httpx_mock.add_response(content=b"ok")
|
httpx_mock.add_response(content=b"ok")
|
||||||
|
|
||||||
send_webhook(
|
send_webhook(
|
||||||
|
|||||||
+57
-54
@@ -49,7 +49,6 @@ from django.db.models import Sum
|
|||||||
from django.db.models import When
|
from django.db.models import When
|
||||||
from django.db.models.functions import Coalesce
|
from django.db.models.functions import Coalesce
|
||||||
from django.db.models.functions import Lower
|
from django.db.models.functions import Lower
|
||||||
from django.db.models.manager import Manager
|
|
||||||
from django.http import FileResponse
|
from django.http import FileResponse
|
||||||
from django.http import Http404
|
from django.http import Http404
|
||||||
from django.http import HttpRequest
|
from django.http import HttpRequest
|
||||||
@@ -256,7 +255,6 @@ from paperless.views import StandardPagination
|
|||||||
from paperless_ai.ai_classifier import get_ai_document_classification
|
from paperless_ai.ai_classifier import get_ai_document_classification
|
||||||
from paperless_ai.ai_classifier import get_llm_output_language
|
from paperless_ai.ai_classifier import get_llm_output_language
|
||||||
from paperless_ai.chat import stream_chat_with_documents
|
from paperless_ai.chat import stream_chat_with_documents
|
||||||
from paperless_ai.exceptions import LLMBlockedError
|
|
||||||
from paperless_ai.exceptions import LLMProviderError
|
from paperless_ai.exceptions import LLMProviderError
|
||||||
from paperless_ai.exceptions import LLMTimeoutError
|
from paperless_ai.exceptions import LLMTimeoutError
|
||||||
from paperless_ai.matching import extract_unmatched_names
|
from paperless_ai.matching import extract_unmatched_names
|
||||||
@@ -1194,6 +1192,7 @@ class DocumentViewSet(
|
|||||||
"version_label",
|
"version_label",
|
||||||
"root_document_id",
|
"root_document_id",
|
||||||
"version_index",
|
"version_index",
|
||||||
|
"page_count",
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
"tags",
|
"tags",
|
||||||
@@ -1273,13 +1272,16 @@ class DocumentViewSet(
|
|||||||
if (
|
if (
|
||||||
"version" not in request.query_params
|
"version" not in request.query_params
|
||||||
or not isinstance(response.data, dict)
|
or not isinstance(response.data, dict)
|
||||||
or "content" not in response.data
|
or not ({"content", "page_count"} & response.data.keys())
|
||||||
):
|
):
|
||||||
return response
|
return response
|
||||||
|
|
||||||
root_doc = self.get_object()
|
root_doc = self.get_object()
|
||||||
content_doc = self._resolve_file_doc(root_doc, request)
|
content_doc = self._resolve_file_doc(root_doc, request)
|
||||||
response.data["content"] = content_doc.content or ""
|
if "content" in response.data:
|
||||||
|
response.data["content"] = content_doc.content or ""
|
||||||
|
if "page_count" in response.data:
|
||||||
|
response.data["page_count"] = content_doc.page_count
|
||||||
return response
|
return response
|
||||||
|
|
||||||
def update(self, request, *args, **kwargs):
|
def update(self, request, *args, **kwargs):
|
||||||
@@ -1698,23 +1700,6 @@ class DocumentViewSet(
|
|||||||
},
|
},
|
||||||
status=status.HTTP_502_BAD_GATEWAY,
|
status=status.HTTP_502_BAD_GATEWAY,
|
||||||
)
|
)
|
||||||
except LLMBlockedError as exc:
|
|
||||||
logger.warning(
|
|
||||||
"AI backend request for document %s was blocked: %s",
|
|
||||||
doc.pk,
|
|
||||||
exc,
|
|
||||||
)
|
|
||||||
return Response(
|
|
||||||
{
|
|
||||||
"ai": [
|
|
||||||
_(
|
|
||||||
"AI backend request was blocked by the outbound "
|
|
||||||
"request policy. Check logs for details.",
|
|
||||||
),
|
|
||||||
],
|
|
||||||
},
|
|
||||||
status=status.HTTP_502_BAD_GATEWAY,
|
|
||||||
)
|
|
||||||
set_llm_suggestions_cache(
|
set_llm_suggestions_cache(
|
||||||
doc.pk,
|
doc.pk,
|
||||||
llm_suggestions,
|
llm_suggestions,
|
||||||
@@ -2986,11 +2971,15 @@ class DocumentOperationPermissionMixin(PassUserMixin, DocumentSelectionMixin):
|
|||||||
if user.is_superuser:
|
if user.is_superuser:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
document_objs = Document.objects.select_related("owner").filter(
|
root_docs = {
|
||||||
pk__in=documents,
|
get_root_document(doc)
|
||||||
)
|
for doc in Document.objects.select_related(
|
||||||
|
"owner",
|
||||||
|
"root_document__owner",
|
||||||
|
).filter(pk__in=documents)
|
||||||
|
}
|
||||||
user_is_owner_of_all_documents = all(
|
user_is_owner_of_all_documents = all(
|
||||||
(doc.owner == user or doc.owner is None) for doc in document_objs
|
(doc.owner == user or doc.owner is None) for doc in root_docs
|
||||||
)
|
)
|
||||||
|
|
||||||
# check global and object permissions for all documents
|
# check global and object permissions for all documents
|
||||||
@@ -2998,9 +2987,13 @@ class DocumentOperationPermissionMixin(PassUserMixin, DocumentSelectionMixin):
|
|||||||
user.has_perm(
|
user.has_perm(
|
||||||
"documents.change_document",
|
"documents.change_document",
|
||||||
)
|
)
|
||||||
and not document_objs.exclude(
|
and not Document.global_objects.filter(
|
||||||
|
pk__in=[doc.pk for doc in root_docs],
|
||||||
|
)
|
||||||
|
.exclude(
|
||||||
pk__in=permitted_document_ids(user, perm="change_document"),
|
pk__in=permitted_document_ids(user, perm="change_document"),
|
||||||
).exists()
|
)
|
||||||
|
.exists()
|
||||||
)
|
)
|
||||||
|
|
||||||
# check ownership for methods that change original document
|
# check ownership for methods that change original document
|
||||||
@@ -3159,6 +3152,38 @@ class BulkEditView(DocumentOperationPermissionMixin):
|
|||||||
|
|
||||||
serializer_class = BulkEditSerializer
|
serializer_class = BulkEditSerializer
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _snapshot_field(doc_ids: list[int], field: str) -> dict[int, Any]:
|
||||||
|
"""
|
||||||
|
Returns each document's current value of field, for the audit log.
|
||||||
|
|
||||||
|
Tags and custom fields are one row per value, so they are gathered
|
||||||
|
into a sorted list of pks per document (empty when there are none).
|
||||||
|
Reading them through Document.values() instead would join those rows
|
||||||
|
and return one arbitrary value per document.
|
||||||
|
"""
|
||||||
|
if field == "tags":
|
||||||
|
rows = (
|
||||||
|
Document.tags.through.objects.filter(document_id__in=doc_ids)
|
||||||
|
.order_by("tag_id")
|
||||||
|
.values_list("document_id", "tag_id")
|
||||||
|
)
|
||||||
|
elif field == "custom_fields":
|
||||||
|
rows = (
|
||||||
|
CustomFieldInstance.objects.filter(document_id__in=doc_ids)
|
||||||
|
.order_by("pk")
|
||||||
|
.values_list("document_id", "pk")
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return dict(
|
||||||
|
Document.objects.filter(pk__in=doc_ids).values_list("pk", field),
|
||||||
|
)
|
||||||
|
|
||||||
|
values: dict[int, list[int]] = {doc_id: [] for doc_id in doc_ids}
|
||||||
|
for doc_id, pk in rows:
|
||||||
|
values[doc_id].append(pk)
|
||||||
|
return values
|
||||||
|
|
||||||
def post(self, request, *args, **kwargs):
|
def post(self, request, *args, **kwargs):
|
||||||
request_method = request.data.get("method")
|
request_method = request.data.get("method")
|
||||||
api_version = int(request.version or settings.REST_FRAMEWORK["DEFAULT_VERSION"])
|
api_version = int(request.version or settings.REST_FRAMEWORK["DEFAULT_VERSION"])
|
||||||
@@ -3205,41 +3230,19 @@ class BulkEditView(DocumentOperationPermissionMixin):
|
|||||||
try:
|
try:
|
||||||
modified_field = self.MODIFIED_FIELD_BY_METHOD.get(method.__name__, None)
|
modified_field = self.MODIFIED_FIELD_BY_METHOD.get(method.__name__, None)
|
||||||
if settings.AUDIT_LOG_ENABLED and modified_field:
|
if settings.AUDIT_LOG_ENABLED and modified_field:
|
||||||
old_documents = {
|
old_values = self._snapshot_field(documents, modified_field)
|
||||||
obj["pk"]: obj
|
|
||||||
for obj in Document.objects.filter(pk__in=documents).values(
|
|
||||||
"pk",
|
|
||||||
"correspondent",
|
|
||||||
"document_type",
|
|
||||||
"storage_path",
|
|
||||||
"tags",
|
|
||||||
"custom_fields",
|
|
||||||
"deleted_at",
|
|
||||||
"checksum",
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
result = method(documents, **parameters)
|
result = method(documents, **parameters)
|
||||||
|
|
||||||
if settings.AUDIT_LOG_ENABLED and modified_field:
|
if settings.AUDIT_LOG_ENABLED and modified_field:
|
||||||
new_documents = Document.objects.filter(pk__in=documents)
|
new_values = self._snapshot_field(documents, modified_field)
|
||||||
for doc in new_documents:
|
for doc in Document.objects.filter(pk__in=documents):
|
||||||
old_value = old_documents[doc.pk][modified_field]
|
|
||||||
new_value = getattr(doc, modified_field)
|
|
||||||
|
|
||||||
if isinstance(new_value, Model):
|
|
||||||
# correspondent, document type, etc.
|
|
||||||
new_value = new_value.pk
|
|
||||||
elif isinstance(new_value, Manager):
|
|
||||||
# tags, custom fields
|
|
||||||
new_value = list(new_value.values_list("pk", flat=True))
|
|
||||||
|
|
||||||
LogEntry.objects.log_create(
|
LogEntry.objects.log_create(
|
||||||
instance=doc,
|
instance=doc,
|
||||||
changes={
|
changes={
|
||||||
modified_field: [
|
modified_field: [
|
||||||
old_value,
|
old_values[doc.pk],
|
||||||
new_value,
|
new_values[doc.pk],
|
||||||
],
|
],
|
||||||
},
|
},
|
||||||
action=LogEntry.Action.UPDATE,
|
action=LogEntry.Action.UPDATE,
|
||||||
|
|||||||
@@ -4,8 +4,7 @@ import httpx
|
|||||||
from celery import shared_task
|
from celery import shared_task
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
|
|
||||||
from paperless.network import GuardedHTTPTransport
|
from paperless.network import PinnedHostHTTPTransport
|
||||||
from paperless.network import OutboundRequestBlockedError
|
|
||||||
from paperless.network import validate_outbound_http_url
|
from paperless.network import validate_outbound_http_url
|
||||||
|
|
||||||
logger = logging.getLogger("paperless.workflows.webhooks")
|
logger = logging.getLogger("paperless.workflows.webhooks")
|
||||||
@@ -15,7 +14,7 @@ logger = logging.getLogger("paperless.workflows.webhooks")
|
|||||||
retry_backoff=True,
|
retry_backoff=True,
|
||||||
autoretry_for=(httpx.HTTPStatusError,),
|
autoretry_for=(httpx.HTTPStatusError,),
|
||||||
max_retries=3,
|
max_retries=3,
|
||||||
throws=(httpx.HTTPError, OutboundRequestBlockedError),
|
throws=(httpx.HTTPError,),
|
||||||
)
|
)
|
||||||
def send_webhook(
|
def send_webhook(
|
||||||
url: str,
|
url: str,
|
||||||
@@ -30,15 +29,14 @@ def send_webhook(
|
|||||||
url,
|
url,
|
||||||
allowed_schemes=settings.WEBHOOKS_ALLOWED_SCHEMES,
|
allowed_schemes=settings.WEBHOOKS_ALLOWED_SCHEMES,
|
||||||
allowed_ports=settings.WEBHOOKS_ALLOWED_PORTS,
|
allowed_ports=settings.WEBHOOKS_ALLOWED_PORTS,
|
||||||
# Scheme and port only; the transport enforces the internal-address
|
# Internal-address checks happen in transport to preserve ConnectError behavior.
|
||||||
# policy at connect time, on the address actually dialled.
|
|
||||||
allow_internal=True,
|
allow_internal=True,
|
||||||
)
|
)
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
logger.warning("Webhook blocked: %s", e)
|
logger.warning("Webhook blocked: %s", e)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
transport = GuardedHTTPTransport(
|
transport = PinnedHostHTTPTransport(
|
||||||
allow_internal=settings.WEBHOOKS_ALLOW_INTERNAL_REQUESTS,
|
allow_internal=settings.WEBHOOKS_ALLOW_INTERNAL_REQUESTS,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ msgid ""
|
|||||||
msgstr ""
|
msgstr ""
|
||||||
"Project-Id-Version: paperless-ngx\n"
|
"Project-Id-Version: paperless-ngx\n"
|
||||||
"Report-Msgid-Bugs-To: \n"
|
"Report-Msgid-Bugs-To: \n"
|
||||||
"POT-Creation-Date: 2026-09-21 19:00+0000\n"
|
"POT-Creation-Date: 2026-09-28 06:26+0000\n"
|
||||||
"PO-Revision-Date: 2022-02-17 04:17\n"
|
"PO-Revision-Date: 2022-02-17 04:17\n"
|
||||||
"Last-Translator: \n"
|
"Last-Translator: \n"
|
||||||
"Language-Team: English\n"
|
"Language-Team: English\n"
|
||||||
@@ -1632,7 +1632,7 @@ msgid "workflow runs"
|
|||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:514 documents/serialisers.py:871
|
#: documents/serialisers.py:514 documents/serialisers.py:871
|
||||||
#: documents/serialisers.py:2885 documents/views.py:343 documents/views.py:2726
|
#: documents/serialisers.py:2895 documents/views.py:342 documents/views.py:2729
|
||||||
#: paperless_mail/serialisers.py:156
|
#: paperless_mail/serialisers.py:156
|
||||||
msgid "Insufficient permissions."
|
msgid "Insufficient permissions."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
@@ -1641,39 +1641,39 @@ msgstr ""
|
|||||||
msgid "Invalid color."
|
msgid "Invalid color."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2352
|
#: documents/serialisers.py:2362
|
||||||
#, python-format
|
#, python-format
|
||||||
msgid "File type %(type)s not supported"
|
msgid "File type %(type)s not supported"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2396
|
#: documents/serialisers.py:2406
|
||||||
#, python-format
|
#, python-format
|
||||||
msgid "Custom field id must be an integer: %(id)s"
|
msgid "Custom field id must be an integer: %(id)s"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2403
|
#: documents/serialisers.py:2413
|
||||||
#, python-format
|
#, python-format
|
||||||
msgid "Custom field with id %(id)s does not exist"
|
msgid "Custom field with id %(id)s does not exist"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2420 documents/serialisers.py:2430
|
#: documents/serialisers.py:2430 documents/serialisers.py:2440
|
||||||
msgid ""
|
msgid ""
|
||||||
"Custom fields must be a list of integers or an object mapping ids to values."
|
"Custom fields must be a list of integers or an object mapping ids to values."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2425
|
#: documents/serialisers.py:2435
|
||||||
msgid "Some custom fields don't exist or were specified twice."
|
msgid "Some custom fields don't exist or were specified twice."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2572
|
#: documents/serialisers.py:2582
|
||||||
msgid "Invalid variable detected."
|
msgid "Invalid variable detected."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2941
|
#: documents/serialisers.py:2951
|
||||||
msgid "Duplicate document identifiers are not allowed."
|
msgid "Duplicate document identifiers are not allowed."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/serialisers.py:2971 documents/views.py:4763
|
#: documents/serialisers.py:2981 documents/views.py:4784
|
||||||
#, python-format
|
#, python-format
|
||||||
msgid "Documents not found: %(ids)s"
|
msgid "Documents not found: %(ids)s"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
@@ -1941,40 +1941,40 @@ msgstr ""
|
|||||||
msgid "Unable to parse URI {value}"
|
msgid "Unable to parse URI {value}"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:336 documents/views.py:2723
|
#: documents/views.py:335 documents/views.py:2726
|
||||||
msgid "Invalid more_like_id"
|
msgid "Invalid more_like_id"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:1670
|
#: documents/views.py:1673
|
||||||
msgid "Invalid AI configuration."
|
msgid "Invalid AI configuration."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:1681
|
#: documents/views.py:1684
|
||||||
msgid "AI backend request timed out."
|
msgid "AI backend request timed out."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:1693
|
#: documents/views.py:1696
|
||||||
msgid "AI backend rejected the request. Check logs for details."
|
msgid "AI backend rejected the request. Check logs for details."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:2548 documents/views.py:2864
|
#: documents/views.py:2551 documents/views.py:2867
|
||||||
msgid "Specify only one of text, title_search, query, or more_like_id."
|
msgid "Specify only one of text, title_search, query, or more_like_id."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:4776
|
#: documents/views.py:4797
|
||||||
#, python-format
|
#, python-format
|
||||||
msgid "Insufficient permissions to share document %(id)s."
|
msgid "Insufficient permissions to share document %(id)s."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:4822
|
#: documents/views.py:4843
|
||||||
msgid "Bundle is already being processed."
|
msgid "Bundle is already being processed."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:4886
|
#: documents/views.py:4907
|
||||||
msgid "The share link bundle is still being prepared. Please try again later."
|
msgid "The share link bundle is still being prepared. Please try again later."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: documents/views.py:4900
|
#: documents/views.py:4921
|
||||||
msgid "The share link bundle is unavailable."
|
msgid "The share link bundle is unavailable."
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
@@ -2219,194 +2219,198 @@ msgid "Sets the LLM embedding model"
|
|||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:369
|
#: paperless/models.py:369
|
||||||
msgid "Sets the LLM embedding endpoint, optional"
|
msgid "Sets the LLM embedding API key"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:376
|
#: paperless/models.py:376
|
||||||
|
msgid "Sets the LLM embedding endpoint, optional"
|
||||||
|
msgstr ""
|
||||||
|
|
||||||
|
#: paperless/models.py:383
|
||||||
msgid "Sets the LLM embedding chunk size"
|
msgid "Sets the LLM embedding chunk size"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:382
|
#: paperless/models.py:389
|
||||||
msgid "Sets the LLM context size"
|
msgid "Sets the LLM context size"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:388
|
#: paperless/models.py:395
|
||||||
msgid "Sets the LLM backend"
|
msgid "Sets the LLM backend"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:396
|
#: paperless/models.py:403
|
||||||
msgid "Sets the LLM model"
|
msgid "Sets the LLM model"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:403
|
#: paperless/models.py:410
|
||||||
msgid "Sets the LLM API key"
|
msgid "Sets the LLM API key"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:410
|
#: paperless/models.py:417
|
||||||
msgid "Sets the LLM endpoint, optional"
|
msgid "Sets the LLM endpoint, optional"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:417
|
#: paperless/models.py:424
|
||||||
msgid "Sets the LLM output language"
|
msgid "Sets the LLM output language"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:424
|
#: paperless/models.py:431
|
||||||
msgid "Sets the LLM timeout in seconds"
|
msgid "Sets the LLM timeout in seconds"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/models.py:430
|
#: paperless/models.py:437
|
||||||
msgid "paperless application settings"
|
msgid "paperless application settings"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:558
|
#: paperless/settings/__init__.py:559
|
||||||
msgid "English (US)"
|
msgid "English (US)"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:559
|
#: paperless/settings/__init__.py:560
|
||||||
msgid "Arabic"
|
msgid "Arabic"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:560
|
#: paperless/settings/__init__.py:561
|
||||||
msgid "Afrikaans"
|
msgid "Afrikaans"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:561
|
#: paperless/settings/__init__.py:562
|
||||||
msgid "Belarusian"
|
msgid "Belarusian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:562
|
#: paperless/settings/__init__.py:563
|
||||||
msgid "Bulgarian"
|
msgid "Bulgarian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:563
|
#: paperless/settings/__init__.py:564
|
||||||
msgid "Catalan"
|
msgid "Catalan"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:564
|
#: paperless/settings/__init__.py:565
|
||||||
msgid "Czech"
|
msgid "Czech"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:565
|
#: paperless/settings/__init__.py:566
|
||||||
msgid "Danish"
|
msgid "Danish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:566
|
#: paperless/settings/__init__.py:567
|
||||||
msgid "German"
|
msgid "German"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:567
|
#: paperless/settings/__init__.py:568
|
||||||
msgid "Greek"
|
msgid "Greek"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:568
|
#: paperless/settings/__init__.py:569
|
||||||
msgid "English (GB)"
|
msgid "English (GB)"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:569
|
#: paperless/settings/__init__.py:570
|
||||||
msgid "Spanish"
|
msgid "Spanish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:570
|
#: paperless/settings/__init__.py:571
|
||||||
msgid "Persian"
|
msgid "Persian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:571
|
#: paperless/settings/__init__.py:572
|
||||||
msgid "Finnish"
|
msgid "Finnish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:572
|
#: paperless/settings/__init__.py:573
|
||||||
msgid "French"
|
msgid "French"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:573
|
#: paperless/settings/__init__.py:574
|
||||||
msgid "Hungarian"
|
msgid "Hungarian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:574
|
#: paperless/settings/__init__.py:575
|
||||||
msgid "Indonesian"
|
msgid "Indonesian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:575
|
#: paperless/settings/__init__.py:576
|
||||||
msgid "Italian"
|
msgid "Italian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:576
|
#: paperless/settings/__init__.py:577
|
||||||
msgid "Japanese"
|
msgid "Japanese"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:577
|
#: paperless/settings/__init__.py:578
|
||||||
msgid "Korean"
|
msgid "Korean"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:578
|
#: paperless/settings/__init__.py:579
|
||||||
msgid "Luxembourgish"
|
msgid "Luxembourgish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:579
|
#: paperless/settings/__init__.py:580
|
||||||
msgid "Norwegian"
|
msgid "Norwegian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:580
|
#: paperless/settings/__init__.py:581
|
||||||
msgid "Dutch"
|
msgid "Dutch"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:581
|
#: paperless/settings/__init__.py:582
|
||||||
msgid "Polish"
|
msgid "Polish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:582
|
#: paperless/settings/__init__.py:583
|
||||||
msgid "Portuguese (Brazil)"
|
msgid "Portuguese (Brazil)"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:583
|
#: paperless/settings/__init__.py:584
|
||||||
msgid "Portuguese"
|
msgid "Portuguese"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:584
|
#: paperless/settings/__init__.py:585
|
||||||
msgid "Romanian"
|
msgid "Romanian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:585
|
#: paperless/settings/__init__.py:586
|
||||||
msgid "Russian"
|
msgid "Russian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:586
|
#: paperless/settings/__init__.py:587
|
||||||
msgid "Slovak"
|
msgid "Slovak"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:587
|
#: paperless/settings/__init__.py:588
|
||||||
msgid "Slovenian"
|
msgid "Slovenian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:588
|
#: paperless/settings/__init__.py:589
|
||||||
msgid "Serbian"
|
msgid "Serbian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:589
|
#: paperless/settings/__init__.py:590
|
||||||
msgid "Swedish"
|
msgid "Swedish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:590
|
#: paperless/settings/__init__.py:591
|
||||||
msgid "Turkish"
|
msgid "Turkish"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:591
|
#: paperless/settings/__init__.py:592
|
||||||
msgid "Ukrainian"
|
msgid "Ukrainian"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:592
|
#: paperless/settings/__init__.py:593
|
||||||
msgid "Vietnamese"
|
msgid "Vietnamese"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:593
|
#: paperless/settings/__init__.py:594
|
||||||
msgid "Chinese Simplified"
|
msgid "Chinese Simplified"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/settings/__init__.py:594
|
#: paperless/settings/__init__.py:595
|
||||||
msgid "Chinese Traditional"
|
msgid "Chinese Traditional"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
#: paperless/urls.py:435
|
#: paperless/urls.py:438
|
||||||
msgid "Paperless-ngx administration"
|
msgid "Paperless-ngx administration"
|
||||||
msgstr ""
|
msgstr ""
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import json
|
import json
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
|
|
||||||
@@ -244,6 +245,7 @@ class AIConfig(BaseConfig):
|
|||||||
ai_enabled: bool = dataclasses.field(init=False)
|
ai_enabled: bool = dataclasses.field(init=False)
|
||||||
llm_embedding_backend: str = dataclasses.field(init=False)
|
llm_embedding_backend: str = dataclasses.field(init=False)
|
||||||
llm_embedding_model: str = dataclasses.field(init=False)
|
llm_embedding_model: str = dataclasses.field(init=False)
|
||||||
|
llm_embedding_api_key: str = dataclasses.field(init=False)
|
||||||
llm_embedding_endpoint: str = dataclasses.field(init=False)
|
llm_embedding_endpoint: str = dataclasses.field(init=False)
|
||||||
llm_embedding_chunk_size: int = dataclasses.field(init=False)
|
llm_embedding_chunk_size: int = dataclasses.field(init=False)
|
||||||
llm_context_size: int = dataclasses.field(init=False)
|
llm_context_size: int = dataclasses.field(init=False)
|
||||||
@@ -254,6 +256,7 @@ class AIConfig(BaseConfig):
|
|||||||
llm_endpoint: str = dataclasses.field(init=False)
|
llm_endpoint: str = dataclasses.field(init=False)
|
||||||
llm_output_language: str = dataclasses.field(init=False)
|
llm_output_language: str = dataclasses.field(init=False)
|
||||||
llm_allow_internal_endpoints: bool = dataclasses.field(init=False)
|
llm_allow_internal_endpoints: bool = dataclasses.field(init=False)
|
||||||
|
llm_extra_params: dict[str, Any] = dataclasses.field(init=False)
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
app_config = self._get_config_instance()
|
app_config = self._get_config_instance()
|
||||||
@@ -269,6 +272,9 @@ class AIConfig(BaseConfig):
|
|||||||
self.llm_embedding_model = (
|
self.llm_embedding_model = (
|
||||||
app_config.llm_embedding_model or settings.LLM_EMBEDDING_MODEL
|
app_config.llm_embedding_model or settings.LLM_EMBEDDING_MODEL
|
||||||
)
|
)
|
||||||
|
self.llm_embedding_api_key = (
|
||||||
|
app_config.llm_embedding_api_key or settings.LLM_EMBEDDING_API_KEY
|
||||||
|
)
|
||||||
self.llm_embedding_endpoint = (
|
self.llm_embedding_endpoint = (
|
||||||
app_config.llm_embedding_endpoint or settings.LLM_EMBEDDING_ENDPOINT
|
app_config.llm_embedding_endpoint or settings.LLM_EMBEDDING_ENDPOINT
|
||||||
)
|
)
|
||||||
@@ -287,6 +293,7 @@ class AIConfig(BaseConfig):
|
|||||||
app_config.llm_output_language or settings.LLM_OUTPUT_LANGUAGE
|
app_config.llm_output_language or settings.LLM_OUTPUT_LANGUAGE
|
||||||
)
|
)
|
||||||
self.llm_allow_internal_endpoints = settings.LLM_ALLOW_INTERNAL_ENDPOINTS
|
self.llm_allow_internal_endpoints = settings.LLM_ALLOW_INTERNAL_ENDPOINTS
|
||||||
|
self.llm_extra_params = settings.LLM_EXTRA_PARAMS
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def llm_index_enabled(self) -> bool:
|
def llm_index_enabled(self) -> bool:
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
# Generated by Django 5.2.16 on 2026-09-11 09:32
|
||||||
|
|
||||||
|
from django.db import migrations
|
||||||
|
from django.db import models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
dependencies = [
|
||||||
|
("paperless", "0016_alter_applicationconfiguration_ai_enabled"),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.AddField(
|
||||||
|
model_name="applicationconfiguration",
|
||||||
|
name="llm_embedding_api_key",
|
||||||
|
field=models.CharField(
|
||||||
|
blank=True,
|
||||||
|
max_length=1024,
|
||||||
|
null=True,
|
||||||
|
verbose_name="Sets the LLM embedding API key",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -365,6 +365,13 @@ class ApplicationConfiguration(AbstractSingletonModel):
|
|||||||
max_length=128,
|
max_length=128,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
llm_embedding_api_key = models.CharField(
|
||||||
|
verbose_name=_("Sets the LLM embedding API key"),
|
||||||
|
blank=True,
|
||||||
|
null=True,
|
||||||
|
max_length=1024,
|
||||||
|
)
|
||||||
|
|
||||||
llm_embedding_endpoint = models.CharField(
|
llm_embedding_endpoint = models.CharField(
|
||||||
verbose_name=_("Sets the LLM embedding endpoint, optional"),
|
verbose_name=_("Sets the LLM embedding endpoint, optional"),
|
||||||
blank=True,
|
blank=True,
|
||||||
|
|||||||
+158
-519
@@ -1,533 +1,61 @@
|
|||||||
import functools
|
|
||||||
import ipaddress
|
import ipaddress
|
||||||
import logging
|
|
||||||
import math
|
|
||||||
import re
|
|
||||||
import socket
|
import socket
|
||||||
import time
|
|
||||||
from collections.abc import Callable
|
|
||||||
from collections.abc import Collection
|
from collections.abc import Collection
|
||||||
from collections.abc import Iterable
|
|
||||||
from enum import StrEnum
|
|
||||||
from typing import Any
|
|
||||||
from typing import Final
|
|
||||||
from typing import Self
|
|
||||||
from typing import TypeAlias
|
|
||||||
from urllib.parse import ParseResult
|
from urllib.parse import ParseResult
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
import anyio
|
|
||||||
import httpcore
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
# Not exported by httpcore; the guard asserts it is still the async default.
|
# Ranges ipaddress does not report as private, but which routinely front
|
||||||
from httpcore._backends.auto import AutoBackend
|
# internal infrastructure.
|
||||||
|
|
||||||
logger = logging.getLogger("paperless.network")
|
|
||||||
|
|
||||||
# requires-python is >=3.11, so no PEP 695 `type` statement.
|
|
||||||
IPAddress: TypeAlias = ipaddress.IPv4Address | ipaddress.IPv6Address
|
|
||||||
|
|
||||||
# Ranges that ipaddress reports as global but which still reach internal hosts.
|
|
||||||
_NON_PUBLIC_NETWORKS = (
|
_NON_PUBLIC_NETWORKS = (
|
||||||
|
# RFC 6598 shared address space: ISP CGNAT, and the default pod/service
|
||||||
|
# CIDR on several managed Kubernetes offerings.
|
||||||
|
ipaddress.ip_network("100.64.0.0/10"),
|
||||||
# RFC 6052 NAT64 well-known prefix: 64:ff9b::7f00:1 is 127.0.0.1 wherever
|
# RFC 6052 NAT64 well-known prefix: 64:ff9b::7f00:1 is 127.0.0.1 wherever
|
||||||
# a NAT64 gateway exists, yet ipaddress classifies the prefix as global.
|
# a NAT64 gateway exists.
|
||||||
ipaddress.ip_network("64:ff9b::/96"),
|
ipaddress.ip_network("64:ff9b::/96"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class BlockReason(StrEnum):
|
def is_public_ip(ip: str | int) -> bool:
|
||||||
NON_PUBLIC_ADDRESS = "non_public_address"
|
try:
|
||||||
UNIX_SOCKET = "unix_socket"
|
obj = ipaddress.ip_address(ip)
|
||||||
|
return not (
|
||||||
|
obj.is_private
|
||||||
class OutboundRequestBlockedError(Exception):
|
or obj.is_loopback
|
||||||
"""
|
or obj.is_link_local
|
||||||
An outbound connection was refused by policy before any socket was opened.
|
or obj.is_multicast
|
||||||
|
or obj.is_unspecified
|
||||||
For NON_PUBLIC_ADDRESS, ``host`` is the name or literal being connected to
|
or any(obj in network for network in _NON_PUBLIC_NETWORKS)
|
||||||
and ``address`` the first offending address. For UNIX_SOCKET, ``host`` is
|
|
||||||
the socket path and ``port`` and ``address`` are None.
|
|
||||||
|
|
||||||
``address`` is deliberately left out of the message: the message is logged
|
|
||||||
and stored on failed tasks, and must not disclose internal addresses.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
reason: BlockReason,
|
|
||||||
address: IPAddress | None = None,
|
|
||||||
) -> None:
|
|
||||||
self.host = host
|
|
||||||
self.port = port
|
|
||||||
self.reason = reason
|
|
||||||
self.address = address
|
|
||||||
target = host if port is None else f"{host}:{port}"
|
|
||||||
super().__init__(f"Outbound connection to {target} blocked ({reason})")
|
|
||||||
|
|
||||||
def __reduce__(self) -> tuple[Callable[..., Self], tuple[object, ...]]:
|
|
||||||
# Celery rebuilds failed-task exceptions by pickling; keyword-only
|
|
||||||
# fields cannot be recovered from ``args`` alone.
|
|
||||||
return (
|
|
||||||
functools.partial(
|
|
||||||
type(self),
|
|
||||||
host=self.host,
|
|
||||||
port=self.port,
|
|
||||||
reason=self.reason,
|
|
||||||
address=self.address,
|
|
||||||
),
|
|
||||||
(),
|
|
||||||
)
|
)
|
||||||
|
except ValueError: # pragma: no cover
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
class HostResolutionError(Exception):
|
def resolve_hostname_ips(hostname: str) -> list[str]:
|
||||||
"""The resolver returned no usable addresses for a host."""
|
try:
|
||||||
|
addr_info = socket.getaddrinfo(hostname, None)
|
||||||
|
except socket.gaierror as e:
|
||||||
|
raise ValueError(f"Could not resolve hostname: {hostname}") from e
|
||||||
|
|
||||||
def __init__(self, *, host: str, detail: str) -> None:
|
ips = [info[4][0] for info in addr_info if info and info[4]]
|
||||||
self.host = host
|
if not ips:
|
||||||
self.detail = detail
|
raise ValueError(f"Could not resolve hostname: {hostname}")
|
||||||
super().__init__(f"Could not resolve {host}: {detail}")
|
return ips
|
||||||
|
|
||||||
def __reduce__(self) -> tuple[Callable[..., Self], tuple[object, ...]]:
|
|
||||||
return (
|
|
||||||
functools.partial(type(self), host=self.host, detail=self.detail),
|
|
||||||
(),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def blocked_message(exc: OutboundRequestBlockedError | HostResolutionError) -> str:
|
def format_host_for_url(host: str) -> str:
|
||||||
"""User-facing text for validation errors, kept stable for existing callers."""
|
|
||||||
if isinstance(exc, HostResolutionError):
|
|
||||||
return f"Could not resolve hostname: {exc.host}"
|
|
||||||
if exc.reason is BlockReason.UNIX_SOCKET:
|
|
||||||
return "Connection blocked: unix sockets are not permitted"
|
|
||||||
return f"Connection blocked: {exc.host} resolves to a non-public address"
|
|
||||||
|
|
||||||
|
|
||||||
def is_public_ip(ip: IPAddress) -> bool:
|
|
||||||
"""
|
"""
|
||||||
True when ``ip`` is globally routable unicast and not in a range that
|
Format IP address for URL use (wrap IPv6 in brackets).
|
||||||
ipaddress reports as global but which still reaches internal hosts.
|
|
||||||
"""
|
|
||||||
return (
|
|
||||||
ip.is_global
|
|
||||||
and not ip.is_multicast
|
|
||||||
and not any(ip in network for network in _NON_PUBLIC_NETWORKS)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Resolver and clock indirection so tests can fake DNS and time for this module
|
|
||||||
# without changing how the stock httpcore backends resolve the literals the
|
|
||||||
# guard dials.
|
|
||||||
_getaddrinfo = socket.getaddrinfo
|
|
||||||
_agetaddrinfo = anyio.getaddrinfo
|
|
||||||
# The clock is a seam because time-machine does not mock monotonic clocks, and
|
|
||||||
# patching time.monotonic globally would also replace the asyncio event loop's
|
|
||||||
# own clock, hanging or misfiring its timers for the rest of the test.
|
|
||||||
_monotonic = time.monotonic
|
|
||||||
|
|
||||||
|
|
||||||
def _collect_addresses(
|
|
||||||
host: str,
|
|
||||||
infos: Iterable[tuple[Any, ...]],
|
|
||||||
) -> tuple[IPAddress, ...]:
|
|
||||||
# Resolver output is always an address, but a scoped IPv6 answer carries a
|
|
||||||
# zone id ("fe80::1%1"), which is dropped before classification.
|
|
||||||
# dict keys keep the first occurrence and resolver order
|
|
||||||
addresses: dict[IPAddress, None] = {}
|
|
||||||
for info in infos:
|
|
||||||
address = ipaddress.ip_address(str(info[4][0]).split("%", 1)[0])
|
|
||||||
addresses.setdefault(address, None)
|
|
||||||
if not addresses:
|
|
||||||
raise HostResolutionError(host=host, detail="no addresses returned")
|
|
||||||
return tuple(addresses)
|
|
||||||
|
|
||||||
|
|
||||||
def _require_public(
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
addresses: tuple[IPAddress, ...],
|
|
||||||
) -> tuple[IPAddress, ...]:
|
|
||||||
for address in addresses:
|
|
||||||
if not is_public_ip(address):
|
|
||||||
raise OutboundRequestBlockedError(
|
|
||||||
host=host,
|
|
||||||
port=port,
|
|
||||||
reason=BlockReason.NON_PUBLIC_ADDRESS,
|
|
||||||
address=address,
|
|
||||||
)
|
|
||||||
return addresses
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_public_addresses(host: str, port: int | None) -> tuple[IPAddress, ...]:
|
|
||||||
"""
|
|
||||||
Resolve ``host`` and return its addresses in resolver order, or raise if
|
|
||||||
any of them is non-public. A name is rejected as a whole; offending
|
|
||||||
addresses are never filtered out.
|
|
||||||
|
|
||||||
IP literals go through the resolver too: getaddrinfo answers them without
|
|
||||||
a lookup, and validating only its answer means no second parser can read
|
|
||||||
the host differently from the one that connects.
|
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
infos = _getaddrinfo(host, port, type=socket.SOCK_STREAM)
|
ip_obj = ipaddress.ip_address(host)
|
||||||
except (OSError, UnicodeError) as e:
|
if ip_obj.version == 6:
|
||||||
raise HostResolutionError(host=host, detail=str(e)) from e
|
return f"[{host}]"
|
||||||
return _require_public(host, port, _collect_addresses(host, infos))
|
return host
|
||||||
|
except ValueError:
|
||||||
|
return host
|
||||||
async def aresolve_public_addresses(
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
) -> tuple[IPAddress, ...]:
|
|
||||||
"""Async variant of resolve_public_addresses."""
|
|
||||||
try:
|
|
||||||
infos = await _agetaddrinfo(host, port, type=socket.SOCK_STREAM)
|
|
||||||
except (OSError, UnicodeError) as e:
|
|
||||||
raise HostResolutionError(host=host, detail=str(e)) from e
|
|
||||||
return _require_public(host, port, _collect_addresses(host, infos))
|
|
||||||
|
|
||||||
|
|
||||||
MAX_ADDRESSES_TRIED: Final = 8
|
|
||||||
MIN_ATTEMPT_TIMEOUT: Final = 2.0
|
|
||||||
MAX_ATTEMPT_TIMEOUT: Final = 10.0
|
|
||||||
|
|
||||||
|
|
||||||
def _require_positive_timeout(host: str, timeout: float | None) -> None:
|
|
||||||
# A zero timeout makes the socket non-blocking and a negative one is
|
|
||||||
# rejected by settimeout; neither can produce a useful connection attempt.
|
|
||||||
if timeout is not None and timeout <= 0:
|
|
||||||
raise httpcore.ConnectTimeout(
|
|
||||||
f"Connect timeout for {host} must be positive, got {timeout}",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _deadline(timeout: float | None) -> float:
|
|
||||||
return math.inf if timeout is None else _monotonic() + timeout
|
|
||||||
|
|
||||||
|
|
||||||
def _attempt_order(addresses: tuple[IPAddress, ...]) -> list[IPAddress]:
|
|
||||||
# Alternate address families, starting with the resolver's first family
|
|
||||||
# (RFC 8305 section 4), so one unreachable family cannot delay the other.
|
|
||||||
first_version = addresses[0].version
|
|
||||||
primary = [a for a in addresses if a.version == first_version]
|
|
||||||
secondary = [a for a in addresses if a.version != first_version]
|
|
||||||
ordered: list[IPAddress] = []
|
|
||||||
for index in range(max(len(primary), len(secondary))):
|
|
||||||
ordered.extend(primary[index : index + 1])
|
|
||||||
ordered.extend(secondary[index : index + 1])
|
|
||||||
return ordered[:MAX_ADDRESSES_TRIED]
|
|
||||||
|
|
||||||
|
|
||||||
def _attempt_timeout(remaining: float, attempts_left: int) -> float:
|
|
||||||
"""
|
|
||||||
Budget for the next attempt. Once the budget is too small to split, or on
|
|
||||||
the last address, the attempt gets everything left. Otherwise it gets an
|
|
||||||
equal share clamped to [MIN, MAX], always leaving MIN for a later attempt.
|
|
||||||
The floor survives one lost SYN; the ceiling bounds how long a black-holed
|
|
||||||
address delays the next one.
|
|
||||||
"""
|
|
||||||
if attempts_left == 1 or remaining < 2 * MIN_ATTEMPT_TIMEOUT:
|
|
||||||
return remaining
|
|
||||||
share = remaining / attempts_left
|
|
||||||
return min(
|
|
||||||
MAX_ATTEMPT_TIMEOUT,
|
|
||||||
max(MIN_ATTEMPT_TIMEOUT, share),
|
|
||||||
remaining - MIN_ATTEMPT_TIMEOUT,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _as_httpcore_timeout(seconds: float) -> float | None:
|
|
||||||
return None if math.isinf(seconds) else seconds
|
|
||||||
|
|
||||||
|
|
||||||
def _log_block(error: OutboundRequestBlockedError) -> None:
|
|
||||||
logger.warning("Blocked outbound connection: %s", error)
|
|
||||||
|
|
||||||
|
|
||||||
def _budget_exhausted(host: str, tried: int, total: int) -> httpcore.ConnectTimeout:
|
|
||||||
return httpcore.ConnectTimeout(
|
|
||||||
f"Timed out connecting to {host} after trying {tried} of {total} addresses",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _next_attempt_budget(
|
|
||||||
host: str,
|
|
||||||
deadline: float,
|
|
||||||
candidates: list[IPAddress],
|
|
||||||
index: int,
|
|
||||||
) -> float:
|
|
||||||
"""Budget for the attempt at index, or a timeout if none is left."""
|
|
||||||
remaining = deadline - _monotonic()
|
|
||||||
if remaining <= 0:
|
|
||||||
raise _budget_exhausted(host, index, len(candidates))
|
|
||||||
return _attempt_timeout(remaining, len(candidates) - index)
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_for_connect(host: str, port: int) -> tuple[IPAddress, ...]:
|
|
||||||
try:
|
|
||||||
return resolve_public_addresses(host, port)
|
|
||||||
except OutboundRequestBlockedError as e:
|
|
||||||
_log_block(e)
|
|
||||||
raise
|
|
||||||
except HostResolutionError as e:
|
|
||||||
raise httpcore.ConnectError(str(e)) from e
|
|
||||||
|
|
||||||
|
|
||||||
async def _aresolve_for_connect(
|
|
||||||
host: str,
|
|
||||||
port: int,
|
|
||||||
timeout: float | None,
|
|
||||||
) -> tuple[IPAddress, ...]:
|
|
||||||
# The scope closes before dialling; attempts are not nested inside it.
|
|
||||||
try:
|
|
||||||
with anyio.fail_after(timeout):
|
|
||||||
return await aresolve_public_addresses(host, port)
|
|
||||||
except TimeoutError as e:
|
|
||||||
raise httpcore.ConnectTimeout(f"Timed out resolving {host}") from e
|
|
||||||
except OutboundRequestBlockedError as e:
|
|
||||||
_log_block(e)
|
|
||||||
raise
|
|
||||||
except HostResolutionError as e:
|
|
||||||
raise httpcore.ConnectError(str(e)) from e
|
|
||||||
|
|
||||||
|
|
||||||
class _GuardedSyncBackend(httpcore.NetworkBackend):
|
|
||||||
"""
|
|
||||||
Wraps httpcore's sync backend. With internal addresses disallowed, it
|
|
||||||
resolves the origin host itself, rejects the name if any address is
|
|
||||||
non-public, and dials the validated literals so the checked address is
|
|
||||||
the connected one. TLS still verifies against the origin hostname.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, inner: httpcore.NetworkBackend, *, allow_internal: bool) -> None:
|
|
||||||
self._inner = inner
|
|
||||||
self._allow_internal = allow_internal
|
|
||||||
|
|
||||||
def connect_tcp(
|
|
||||||
self,
|
|
||||||
host: str,
|
|
||||||
port: int,
|
|
||||||
timeout: float | None = None,
|
|
||||||
local_address: str | None = None,
|
|
||||||
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
||||||
) -> httpcore.NetworkStream:
|
|
||||||
if self._allow_internal:
|
|
||||||
return self._inner.connect_tcp(
|
|
||||||
host,
|
|
||||||
port,
|
|
||||||
timeout=timeout,
|
|
||||||
local_address=local_address,
|
|
||||||
socket_options=socket_options,
|
|
||||||
)
|
|
||||||
_require_positive_timeout(host, timeout)
|
|
||||||
# Resolution is not charged to the budget, matching the stock backend.
|
|
||||||
candidates = _attempt_order(_resolve_for_connect(host, port))
|
|
||||||
deadline = _deadline(timeout)
|
|
||||||
last_error: httpcore.ConnectError | httpcore.ConnectTimeout | None = None
|
|
||||||
for index, address in enumerate(candidates):
|
|
||||||
budget = _next_attempt_budget(host, deadline, candidates, index)
|
|
||||||
try:
|
|
||||||
return self._inner.connect_tcp(
|
|
||||||
str(address),
|
|
||||||
port,
|
|
||||||
timeout=_as_httpcore_timeout(budget),
|
|
||||||
local_address=local_address,
|
|
||||||
socket_options=socket_options,
|
|
||||||
)
|
|
||||||
except (httpcore.ConnectError, httpcore.ConnectTimeout) as e:
|
|
||||||
logger.debug("Connecting to %s via %s failed: %s", host, address, e)
|
|
||||||
last_error = e
|
|
||||||
# candidates is never empty, so every address was tried and failed
|
|
||||||
raise last_error or _budget_exhausted(host, len(candidates), len(candidates))
|
|
||||||
|
|
||||||
def connect_unix_socket(
|
|
||||||
self,
|
|
||||||
path: str,
|
|
||||||
timeout: float | None = None,
|
|
||||||
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
||||||
) -> httpcore.NetworkStream:
|
|
||||||
error = OutboundRequestBlockedError(
|
|
||||||
host=path,
|
|
||||||
port=None,
|
|
||||||
reason=BlockReason.UNIX_SOCKET,
|
|
||||||
)
|
|
||||||
_log_block(error)
|
|
||||||
raise error
|
|
||||||
|
|
||||||
def sleep(self, seconds: float) -> None:
|
|
||||||
self._inner.sleep(seconds)
|
|
||||||
|
|
||||||
|
|
||||||
class _GuardedAsyncBackend(httpcore.AsyncNetworkBackend):
|
|
||||||
"""Async twin of _GuardedSyncBackend."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
inner: httpcore.AsyncNetworkBackend,
|
|
||||||
*,
|
|
||||||
allow_internal: bool,
|
|
||||||
) -> None:
|
|
||||||
self._inner = inner
|
|
||||||
self._allow_internal = allow_internal
|
|
||||||
|
|
||||||
async def connect_tcp(
|
|
||||||
self,
|
|
||||||
host: str,
|
|
||||||
port: int,
|
|
||||||
timeout: float | None = None,
|
|
||||||
local_address: str | None = None,
|
|
||||||
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
||||||
) -> httpcore.AsyncNetworkStream:
|
|
||||||
if self._allow_internal:
|
|
||||||
return await self._inner.connect_tcp(
|
|
||||||
host,
|
|
||||||
port,
|
|
||||||
timeout=timeout,
|
|
||||||
local_address=local_address,
|
|
||||||
socket_options=socket_options,
|
|
||||||
)
|
|
||||||
_require_positive_timeout(host, timeout)
|
|
||||||
# Resolution counts against the budget, matching the stock backend.
|
|
||||||
deadline = _deadline(timeout)
|
|
||||||
candidates = _attempt_order(await _aresolve_for_connect(host, port, timeout))
|
|
||||||
last_error: httpcore.ConnectError | httpcore.ConnectTimeout | None = None
|
|
||||||
for index, address in enumerate(candidates):
|
|
||||||
budget = _next_attempt_budget(host, deadline, candidates, index)
|
|
||||||
try:
|
|
||||||
return await self._inner.connect_tcp(
|
|
||||||
str(address),
|
|
||||||
port,
|
|
||||||
timeout=_as_httpcore_timeout(budget),
|
|
||||||
local_address=local_address,
|
|
||||||
socket_options=socket_options,
|
|
||||||
)
|
|
||||||
except (httpcore.ConnectError, httpcore.ConnectTimeout) as e:
|
|
||||||
logger.debug("Connecting to %s via %s failed: %s", host, address, e)
|
|
||||||
last_error = e
|
|
||||||
raise last_error or _budget_exhausted(host, len(candidates), len(candidates))
|
|
||||||
|
|
||||||
async def connect_unix_socket(
|
|
||||||
self,
|
|
||||||
path: str,
|
|
||||||
timeout: float | None = None,
|
|
||||||
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
||||||
) -> httpcore.AsyncNetworkStream:
|
|
||||||
error = OutboundRequestBlockedError(
|
|
||||||
host=path,
|
|
||||||
port=None,
|
|
||||||
reason=BlockReason.UNIX_SOCKET,
|
|
||||||
)
|
|
||||||
_log_block(error)
|
|
||||||
raise error
|
|
||||||
|
|
||||||
async def sleep(self, seconds: float) -> None:
|
|
||||||
await self._inner.sleep(seconds)
|
|
||||||
|
|
||||||
|
|
||||||
_LAYOUT_ERROR = (
|
|
||||||
"Unexpected httpx transport layout; refusing to create a transport "
|
|
||||||
"without the outbound connection guard"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class GuardedHTTPTransport(httpx.HTTPTransport):
|
|
||||||
"""
|
|
||||||
httpx transport whose connections pass through the outbound guard.
|
|
||||||
|
|
||||||
Deliberately accepts no proxy, uds or retries options: a proxy would be
|
|
||||||
dialled instead of the destination, and a unix socket bypasses TCP
|
|
||||||
entirely. Adding an option here is a reviewed change, not a pass-through.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, *, allow_internal: bool) -> None:
|
|
||||||
super().__init__()
|
|
||||||
# httpx has no public hook for the network backend. Check the exact
|
|
||||||
# layout before swapping so an httpx or httpcore change fails loudly.
|
|
||||||
pool = self._pool
|
|
||||||
if (
|
|
||||||
type(pool) is not httpcore.ConnectionPool
|
|
||||||
or type(pool._network_backend) is not httpcore.SyncBackend
|
|
||||||
):
|
|
||||||
raise RuntimeError(_LAYOUT_ERROR)
|
|
||||||
pool._network_backend = _GuardedSyncBackend(
|
|
||||||
pool._network_backend,
|
|
||||||
allow_internal=allow_internal,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class GuardedAsyncHTTPTransport(httpx.AsyncHTTPTransport):
|
|
||||||
"""Async twin of GuardedHTTPTransport."""
|
|
||||||
|
|
||||||
def __init__(self, *, allow_internal: bool) -> None:
|
|
||||||
super().__init__()
|
|
||||||
pool = self._pool
|
|
||||||
if (
|
|
||||||
type(pool) is not httpcore.AsyncConnectionPool
|
|
||||||
or type(pool._network_backend) is not AutoBackend
|
|
||||||
):
|
|
||||||
raise RuntimeError(_LAYOUT_ERROR)
|
|
||||||
pool._network_backend = _GuardedAsyncBackend(
|
|
||||||
pool._network_backend,
|
|
||||||
allow_internal=allow_internal,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def create_guarded_httpx_client(
|
|
||||||
url: str,
|
|
||||||
*,
|
|
||||||
allow_internal: bool,
|
|
||||||
timeout: float,
|
|
||||||
) -> httpx.Client:
|
|
||||||
"""
|
|
||||||
Validate ``url`` up front, then build a client that re-checks at connect
|
|
||||||
time. The up-front check turns static misconfiguration into a ValueError
|
|
||||||
before any retry layer sees it.
|
|
||||||
"""
|
|
||||||
validate_outbound_http_url(url, allow_internal=allow_internal)
|
|
||||||
return httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=allow_internal),
|
|
||||||
timeout=timeout,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def create_guarded_async_httpx_client(
|
|
||||||
url: str,
|
|
||||||
*,
|
|
||||||
allow_internal: bool,
|
|
||||||
timeout: float,
|
|
||||||
) -> httpx.AsyncClient:
|
|
||||||
"""Async twin of create_guarded_httpx_client."""
|
|
||||||
validate_outbound_http_url(url, allow_internal=allow_internal)
|
|
||||||
return httpx.AsyncClient(
|
|
||||||
transport=GuardedAsyncHTTPTransport(allow_internal=allow_internal),
|
|
||||||
timeout=timeout,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# urllib3 treats a backslash as ending the authority while urlparse and httpx do
|
|
||||||
# not, so the host checked here could differ from the one that is dialled.
|
|
||||||
# Control and whitespace characters are refused for the same reason.
|
|
||||||
_UNSAFE_URL_CHARS = re.compile(r"[\\\x00-\x1f\x7f\s]")
|
|
||||||
|
|
||||||
|
|
||||||
def _dns_name(url: str) -> str:
|
|
||||||
"""
|
|
||||||
The ASCII hostname that httpx and urllib3 look up for ``url``.
|
|
||||||
|
|
||||||
urlparse keeps a non-ASCII hostname as typed, and getaddrinfo would then
|
|
||||||
encode it with the stdlib IDNA 2003 codec. That maps some characters
|
|
||||||
differently from the IDNA 2008 encoding the HTTP clients use ("faß"
|
|
||||||
becomes "fass" instead of "xn--fa-hia"), so the check would resolve a
|
|
||||||
different name from the one that is connected to.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
return httpx.URL(url).raw_host.decode("ascii")
|
|
||||||
except (httpx.InvalidURL, UnicodeError) as e:
|
|
||||||
raise ValueError("Invalid URL scheme or hostname.") from e
|
|
||||||
|
|
||||||
|
|
||||||
def validate_outbound_http_url(
|
def validate_outbound_http_url(
|
||||||
@@ -553,17 +81,128 @@ def validate_outbound_http_url(
|
|||||||
raise ValueError("Destination port not permitted.")
|
raise ValueError("Destination port not permitted.")
|
||||||
|
|
||||||
if not allow_internal:
|
if not allow_internal:
|
||||||
if _UNSAFE_URL_CHARS.search(url):
|
for ip_str in resolve_hostname_ips(parsed.hostname):
|
||||||
raise ValueError("Invalid URL scheme or hostname.")
|
if not is_public_ip(ip_str):
|
||||||
host = _dns_name(url)
|
raise ValueError(
|
||||||
# HTTP clients may percent-decode the host before resolving it, so the
|
f"Connection blocked: {parsed.hostname} resolves to a non-public address",
|
||||||
# checked name could differ from the dialled one. An IPv6 zone id is the
|
)
|
||||||
# only legitimate use, and link-local addresses are non-public anyway.
|
|
||||||
if "%" in host:
|
|
||||||
raise ValueError("Invalid URL scheme or hostname.")
|
|
||||||
try:
|
|
||||||
resolve_public_addresses(host, port)
|
|
||||||
except (OutboundRequestBlockedError, HostResolutionError) as e:
|
|
||||||
raise ValueError(blocked_message(e)) from e
|
|
||||||
|
|
||||||
return parsed
|
return parsed
|
||||||
|
|
||||||
|
|
||||||
|
def _rewrite_request_to_pinned_ip(
|
||||||
|
request: httpx.Request,
|
||||||
|
*,
|
||||||
|
allow_internal: bool,
|
||||||
|
) -> httpx.Request:
|
||||||
|
hostname = request.url.host
|
||||||
|
|
||||||
|
if not hostname:
|
||||||
|
raise httpx.ConnectError("No hostname in request URL")
|
||||||
|
|
||||||
|
try:
|
||||||
|
ips = resolve_hostname_ips(hostname)
|
||||||
|
except ValueError as e:
|
||||||
|
raise httpx.ConnectError(str(e)) from e
|
||||||
|
|
||||||
|
if not allow_internal:
|
||||||
|
for ip_str in ips:
|
||||||
|
if not is_public_ip(ip_str):
|
||||||
|
raise httpx.ConnectError(
|
||||||
|
f"Connection blocked: {hostname} resolves to a non-public address",
|
||||||
|
)
|
||||||
|
|
||||||
|
ip_str = ips[0]
|
||||||
|
formatted_ip = format_host_for_url(ip_str)
|
||||||
|
|
||||||
|
new_headers = httpx.Headers(request.headers)
|
||||||
|
if "host" in new_headers:
|
||||||
|
del new_headers["host"]
|
||||||
|
host_header = format_host_for_url(hostname)
|
||||||
|
default_port = 443 if request.url.scheme == "https" else 80
|
||||||
|
if request.url.port and request.url.port != default_port:
|
||||||
|
host_header = f"{host_header}:{request.url.port}"
|
||||||
|
new_headers["Host"] = host_header
|
||||||
|
new_url = request.url.copy_with(host=formatted_ip)
|
||||||
|
|
||||||
|
rewritten_request = httpx.Request(
|
||||||
|
method=request.method,
|
||||||
|
url=new_url,
|
||||||
|
headers=new_headers,
|
||||||
|
stream=request.stream,
|
||||||
|
extensions=request.extensions,
|
||||||
|
)
|
||||||
|
rewritten_request.extensions["sni_hostname"] = hostname
|
||||||
|
|
||||||
|
return rewritten_request
|
||||||
|
|
||||||
|
|
||||||
|
class PinnedHostHTTPTransport(httpx.HTTPTransport):
|
||||||
|
"""
|
||||||
|
HTTP transport that resolves/validates hostnames per request and connects to
|
||||||
|
a vetted IP while preserving the original Host header and TLS SNI hostname.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*args,
|
||||||
|
allow_internal: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.allow_internal = allow_internal
|
||||||
|
|
||||||
|
def handle_request(self, request: httpx.Request) -> httpx.Response:
|
||||||
|
request = _rewrite_request_to_pinned_ip(
|
||||||
|
request,
|
||||||
|
allow_internal=self.allow_internal,
|
||||||
|
)
|
||||||
|
return super().handle_request(request)
|
||||||
|
|
||||||
|
|
||||||
|
class PinnedHostAsyncHTTPTransport(httpx.AsyncHTTPTransport):
|
||||||
|
"""
|
||||||
|
Async variant of PinnedHostHTTPTransport.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*args,
|
||||||
|
allow_internal: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.allow_internal = allow_internal
|
||||||
|
|
||||||
|
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||||
|
request = _rewrite_request_to_pinned_ip(
|
||||||
|
request,
|
||||||
|
allow_internal=self.allow_internal,
|
||||||
|
)
|
||||||
|
return await super().handle_async_request(request)
|
||||||
|
|
||||||
|
|
||||||
|
def create_pinned_httpx_client(
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
allow_internal: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
) -> httpx.Client:
|
||||||
|
validate_outbound_http_url(url, allow_internal=allow_internal)
|
||||||
|
return httpx.Client(
|
||||||
|
transport=PinnedHostHTTPTransport(allow_internal=allow_internal),
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def create_pinned_async_httpx_client(
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
allow_internal: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
) -> httpx.AsyncClient:
|
||||||
|
validate_outbound_http_url(url, allow_internal=allow_internal)
|
||||||
|
return httpx.AsyncClient(
|
||||||
|
transport=PinnedHostAsyncHTTPTransport(allow_internal=allow_internal),
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|||||||
@@ -22,8 +22,6 @@ from pathlib import Path
|
|||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
from typing import Self
|
from typing import Self
|
||||||
|
|
||||||
from bleach import clean
|
|
||||||
from bleach import linkify
|
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from django.utils.timezone import is_naive
|
from django.utils.timezone import is_naive
|
||||||
@@ -38,6 +36,9 @@ from humanize import naturalsize
|
|||||||
from imap_tools import MailAttachment
|
from imap_tools import MailAttachment
|
||||||
from imap_tools import MailMessage
|
from imap_tools import MailMessage
|
||||||
from tika_client import TikaClient
|
from tika_client import TikaClient
|
||||||
|
from turbohtml.clean import Linkify
|
||||||
|
from turbohtml.clean import linkify
|
||||||
|
from turbohtml.migration.bleach import clean
|
||||||
|
|
||||||
from documents.parsers import ParseError
|
from documents.parsers import ParseError
|
||||||
from documents.parsers import make_thumbnail_from_pdf
|
from documents.parsers import make_thumbnail_from_pdf
|
||||||
@@ -58,12 +59,6 @@ _SUPPORTED_MIME_TYPES: dict[str, str] = {
|
|||||||
"message/rfc822": ".eml",
|
"message/rfc822": ".eml",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Bleach's email-address linkifier uses a superlinear regular expression. Keep
|
|
||||||
# email linkification for ordinary headers and short messages, but never run it
|
|
||||||
# over an unbounded attacker-controlled field. URL linkification remains enabled
|
|
||||||
# for longer text.
|
|
||||||
_MAX_EMAIL_LINKIFY_LENGTH = 2048
|
|
||||||
|
|
||||||
|
|
||||||
class MailDocumentParser:
|
class MailDocumentParser:
|
||||||
"""Parse .eml email files for Paperless-ngx.
|
"""Parse .eml email files for Paperless-ngx.
|
||||||
@@ -633,10 +628,7 @@ class MailDocumentParser:
|
|||||||
text = str(text)
|
text = str(text)
|
||||||
text = escape(text)
|
text = escape(text)
|
||||||
text = clean(text)
|
text = clean(text)
|
||||||
text = linkify(
|
text = linkify(text, Linkify(parse_email=True))
|
||||||
text,
|
|
||||||
parse_email="@" in text and len(text) <= _MAX_EMAIL_LINKIFY_LENGTH,
|
|
||||||
)
|
|
||||||
text = text.replace("\n", "<br>")
|
text = text.replace("\n", "<br>")
|
||||||
return text
|
return text
|
||||||
|
|
||||||
|
|||||||
@@ -216,6 +216,11 @@ class ApplicationConfigurationSerializer(
|
|||||||
externally_configured_variables = serializers.SerializerMethodField()
|
externally_configured_variables = serializers.SerializerMethodField()
|
||||||
user_args = serializers.JSONField(binary=True, allow_null=True)
|
user_args = serializers.JSONField(binary=True, allow_null=True)
|
||||||
barcode_tag_mapping = serializers.JSONField(binary=True, allow_null=True)
|
barcode_tag_mapping = serializers.JSONField(binary=True, allow_null=True)
|
||||||
|
llm_embedding_api_key = ObfuscatedPasswordField(
|
||||||
|
required=False,
|
||||||
|
allow_null=True,
|
||||||
|
max_length=1024,
|
||||||
|
)
|
||||||
llm_api_key = ObfuscatedPasswordField(
|
llm_api_key = ObfuscatedPasswordField(
|
||||||
required=False,
|
required=False,
|
||||||
allow_null=True,
|
allow_null=True,
|
||||||
@@ -227,7 +232,11 @@ class ApplicationConfigurationSerializer(
|
|||||||
max_length=1024,
|
max_length=1024,
|
||||||
)
|
)
|
||||||
|
|
||||||
OBFUSCATED_FIELDS = ("llm_api_key", "remote_ocr_api_key")
|
OBFUSCATED_FIELDS = (
|
||||||
|
"llm_embedding_api_key",
|
||||||
|
"llm_api_key",
|
||||||
|
"remote_ocr_api_key",
|
||||||
|
)
|
||||||
|
|
||||||
def get_externally_configured_variables(
|
def get_externally_configured_variables(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import multiprocessing
|
|||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
from typing import Final
|
from typing import Final
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
@@ -1081,6 +1082,25 @@ CLASSIFIER_LANGUAGES: Final[dict[str, str]] = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_llm_extra_params() -> dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Parse PAPERLESS_AI_LLM_EXTRA_PARAMS, a JSON object passed straight through
|
||||||
|
to the LLM backend's request body.
|
||||||
|
"""
|
||||||
|
raw = os.getenv("PAPERLESS_AI_LLM_EXTRA_PARAMS", "{}")
|
||||||
|
try:
|
||||||
|
parsed = json.loads(raw)
|
||||||
|
except json.JSONDecodeError as e:
|
||||||
|
raise ImproperlyConfigured(
|
||||||
|
"PAPERLESS_AI_LLM_EXTRA_PARAMS must be valid JSON",
|
||||||
|
) from e
|
||||||
|
if not isinstance(parsed, dict):
|
||||||
|
raise ImproperlyConfigured(
|
||||||
|
"PAPERLESS_AI_LLM_EXTRA_PARAMS must be a JSON object",
|
||||||
|
)
|
||||||
|
return parsed
|
||||||
|
|
||||||
|
|
||||||
def _get_classifier_language_setting(ocr_lang: str) -> str | None:
|
def _get_classifier_language_setting(ocr_lang: str) -> str | None:
|
||||||
"""
|
"""
|
||||||
Maps the primary Tesseract language to the classifier's stemming
|
Maps the primary Tesseract language to the classifier's stemming
|
||||||
@@ -1216,6 +1236,7 @@ LLM_EMBEDDING_BACKEND = get_choice_from_env(
|
|||||||
{"huggingface", "openai-like", "ollama"},
|
{"huggingface", "openai-like", "ollama"},
|
||||||
)
|
)
|
||||||
LLM_EMBEDDING_MODEL = os.getenv("PAPERLESS_AI_LLM_EMBEDDING_MODEL")
|
LLM_EMBEDDING_MODEL = os.getenv("PAPERLESS_AI_LLM_EMBEDDING_MODEL")
|
||||||
|
LLM_EMBEDDING_API_KEY = os.getenv("PAPERLESS_AI_LLM_EMBEDDING_API_KEY")
|
||||||
LLM_EMBEDDING_ENDPOINT = os.getenv("PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT")
|
LLM_EMBEDDING_ENDPOINT = os.getenv("PAPERLESS_AI_LLM_EMBEDDING_ENDPOINT")
|
||||||
LLM_EMBEDDING_CHUNK_SIZE = get_int_from_env(
|
LLM_EMBEDDING_CHUNK_SIZE = get_int_from_env(
|
||||||
"PAPERLESS_AI_LLM_EMBEDDING_CHUNK_SIZE",
|
"PAPERLESS_AI_LLM_EMBEDDING_CHUNK_SIZE",
|
||||||
@@ -1241,3 +1262,4 @@ LLM_ALLOW_INTERNAL_ENDPOINTS = get_bool_from_env(
|
|||||||
"PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS",
|
"PAPERLESS_AI_LLM_ALLOW_INTERNAL_ENDPOINTS",
|
||||||
"true",
|
"true",
|
||||||
)
|
)
|
||||||
|
LLM_EXTRA_PARAMS = _get_llm_extra_params()
|
||||||
|
|||||||
@@ -735,33 +735,24 @@ class TestParser:
|
|||||||
|
|
||||||
assert expected_html == actual_html
|
assert expected_html == actual_html
|
||||||
|
|
||||||
def test_mail_to_html_bounds_email_linkification(
|
def test_mail_to_html_linkifies_email_in_long_text(
|
||||||
self,
|
self,
|
||||||
mail_parser: MailDocumentParser,
|
mail_parser: MailDocumentParser,
|
||||||
) -> None:
|
) -> None:
|
||||||
mail = mock.Mock(
|
mail = mock.Mock(
|
||||||
subject="sender@example.com",
|
subject="",
|
||||||
from_values=None,
|
from_values=None,
|
||||||
to_values=[],
|
to_values=[],
|
||||||
cc_values=[],
|
cc_values=[],
|
||||||
bcc_values=[],
|
bcc_values=[],
|
||||||
attachments=[],
|
attachments=[],
|
||||||
date=timezone.now(),
|
date=timezone.now(),
|
||||||
text=("a." * 1500) + "@example.com",
|
text=("a." * 1500) + " sender@example.com",
|
||||||
)
|
)
|
||||||
|
|
||||||
with mock.patch(
|
html_file = mail_parser.mail_to_html(mail)
|
||||||
"paperless.parsers.mail.linkify",
|
|
||||||
side_effect=lambda text, **kwargs: text,
|
|
||||||
) as mock_linkify:
|
|
||||||
mail_parser.mail_to_html(mail)
|
|
||||||
|
|
||||||
parse_email_by_text = {
|
assert 'href="mailto:sender@example.com"' in html_file.read_text()
|
||||||
call.args[0]: call.kwargs["parse_email"]
|
|
||||||
for call in mock_linkify.call_args_list
|
|
||||||
}
|
|
||||||
assert parse_email_by_text["sender@example.com"] is True
|
|
||||||
assert parse_email_by_text[mail.text] is False
|
|
||||||
|
|
||||||
def test_generate_pdf_from_mail(
|
def test_generate_pdf_from_mail(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -17,6 +17,24 @@ class TestRemoteUser(DirectoriesMixin, APITestCase):
|
|||||||
|
|
||||||
self.user = UserFactory(username="temp_admin", superuser=True)
|
self.user = UserFactory(username="temp_admin", superuser=True)
|
||||||
|
|
||||||
|
# _parse_remote_user_settings() mutates these shared lists in place,
|
||||||
|
# so undo that after the test instead of leaking remote-user auth
|
||||||
|
# into every test that runs afterward.
|
||||||
|
original_middleware = list(settings.MIDDLEWARE)
|
||||||
|
original_auth_backends = list(settings.AUTHENTICATION_BACKENDS)
|
||||||
|
original_auth_classes = list(
|
||||||
|
settings.REST_FRAMEWORK["DEFAULT_AUTHENTICATION_CLASSES"],
|
||||||
|
)
|
||||||
|
|
||||||
|
def _restore_remote_user_settings() -> None:
|
||||||
|
settings.MIDDLEWARE[:] = original_middleware
|
||||||
|
settings.AUTHENTICATION_BACKENDS[:] = original_auth_backends
|
||||||
|
settings.REST_FRAMEWORK["DEFAULT_AUTHENTICATION_CLASSES"][:] = (
|
||||||
|
original_auth_classes
|
||||||
|
)
|
||||||
|
|
||||||
|
self.addCleanup(_restore_remote_user_settings)
|
||||||
|
|
||||||
def test_remote_user(self) -> None:
|
def test_remote_user(self) -> None:
|
||||||
"""
|
"""
|
||||||
GIVEN:
|
GIVEN:
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from django.core.exceptions import ImproperlyConfigured
|
|||||||
|
|
||||||
from paperless.settings import _get_allauth_trusted_proxy_count
|
from paperless.settings import _get_allauth_trusted_proxy_count
|
||||||
from paperless.settings import _get_classifier_language_setting
|
from paperless.settings import _get_classifier_language_setting
|
||||||
|
from paperless.settings import _get_llm_extra_params
|
||||||
from paperless.settings import _get_search_language_setting
|
from paperless.settings import _get_search_language_setting
|
||||||
from paperless.settings import _parse_paperless_url
|
from paperless.settings import _parse_paperless_url
|
||||||
from paperless.settings import default_threads_per_worker
|
from paperless.settings import default_threads_per_worker
|
||||||
@@ -166,3 +167,45 @@ class TestPaperlessURLSettings(TestCase):
|
|||||||
|
|
||||||
self.assertIn(url, settings.CSRF_TRUSTED_ORIGINS)
|
self.assertIn(url, settings.CSRF_TRUSTED_ORIGINS)
|
||||||
self.assertIn(url, settings.CORS_ALLOWED_ORIGINS)
|
self.assertIn(url, settings.CORS_ALLOWED_ORIGINS)
|
||||||
|
|
||||||
|
|
||||||
|
class TestLlmExtraParams:
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("env_value", "expected"),
|
||||||
|
[
|
||||||
|
pytest.param(None, {}, id="unset"),
|
||||||
|
pytest.param(
|
||||||
|
'{"reasoning_effort": "none"}',
|
||||||
|
{"reasoning_effort": "none"},
|
||||||
|
id="json-object",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_parses(
|
||||||
|
self,
|
||||||
|
monkeypatch,
|
||||||
|
env_value,
|
||||||
|
expected,
|
||||||
|
):
|
||||||
|
if env_value is None:
|
||||||
|
monkeypatch.delenv("PAPERLESS_AI_LLM_EXTRA_PARAMS", raising=False)
|
||||||
|
else:
|
||||||
|
monkeypatch.setenv("PAPERLESS_AI_LLM_EXTRA_PARAMS", env_value)
|
||||||
|
assert _get_llm_extra_params() == expected
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("env_value", "match"),
|
||||||
|
[
|
||||||
|
pytest.param("reasoning_effort=none", "valid JSON", id="invalid-json"),
|
||||||
|
pytest.param('["none"]', "JSON object", id="not-an-object"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_invalid_raises(
|
||||||
|
self,
|
||||||
|
monkeypatch,
|
||||||
|
env_value,
|
||||||
|
match,
|
||||||
|
):
|
||||||
|
monkeypatch.setenv("PAPERLESS_AI_LLM_EXTRA_PARAMS", env_value)
|
||||||
|
with pytest.raises(ImproperlyConfigured, match=match):
|
||||||
|
_get_llm_extra_params()
|
||||||
|
|||||||
@@ -0,0 +1,68 @@
|
|||||||
|
import time
|
||||||
|
|
||||||
|
from allauth.mfa import app_settings as mfa_settings
|
||||||
|
from allauth.mfa.totp.internal import auth as totp_auth
|
||||||
|
from django.test import TestCase
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
from paperless_testing.factories import UserFactory
|
||||||
|
|
||||||
|
|
||||||
|
class TestAdminAuth(TestCase):
|
||||||
|
def test_admin_login_redirects_to_allauth(self):
|
||||||
|
user = UserFactory(staff=True, password="testpassword")
|
||||||
|
admin_url = reverse("admin:index")
|
||||||
|
login_url = reverse("admin:login")
|
||||||
|
expected_url = f"{reverse('account_login')}?next={admin_url}"
|
||||||
|
|
||||||
|
response = self.client.get(login_url, {"next": admin_url})
|
||||||
|
self.assertRedirects(response, expected_url)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
login_url,
|
||||||
|
{"username": user.username, "password": "testpassword", "next": admin_url},
|
||||||
|
)
|
||||||
|
self.assertRedirects(response, expected_url)
|
||||||
|
self.assertNotIn("_auth_user_id", self.client.session)
|
||||||
|
|
||||||
|
def test_admin_access_requires_totp_for_enrolled_staff(self):
|
||||||
|
user = UserFactory(staff=True, password="testpassword")
|
||||||
|
secret = totp_auth.generate_totp_secret()
|
||||||
|
totp_auth.TOTP.activate(user, secret)
|
||||||
|
admin_url = reverse("admin:index")
|
||||||
|
mfa_url = reverse("mfa_authenticate")
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("account_login"),
|
||||||
|
{"login": user.username, "password": "testpassword", "next": admin_url},
|
||||||
|
)
|
||||||
|
self.assertRedirects(response, mfa_url)
|
||||||
|
self.assertNotIn("_auth_user_id", self.client.session)
|
||||||
|
self.assertRedirects(
|
||||||
|
self.client.get(admin_url),
|
||||||
|
f"{reverse('admin:login')}?next={admin_url}",
|
||||||
|
fetch_redirect_response=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.post(mfa_url, {"code": "invalid"})
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertNotIn("_auth_user_id", self.client.session)
|
||||||
|
|
||||||
|
code = totp_auth.format_hotp_value(
|
||||||
|
totp_auth.hotp_value(secret, int(time.time()) // mfa_settings.TOTP_PERIOD),
|
||||||
|
)
|
||||||
|
response = self.client.post(mfa_url, {"code": code})
|
||||||
|
self.assertRedirects(response, admin_url)
|
||||||
|
self.assertEqual(self.client.session["_auth_user_id"], str(user.pk))
|
||||||
|
|
||||||
|
def test_staff_without_totp_can_still_log_in(self):
|
||||||
|
user = UserFactory(staff=True, password="testpassword")
|
||||||
|
admin_url = reverse("admin:index")
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("account_login"),
|
||||||
|
{"login": user.username, "password": "testpassword", "next": admin_url},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertRedirects(response, admin_url)
|
||||||
|
self.assertEqual(self.client.session["_auth_user_id"], str(user.pk))
|
||||||
@@ -30,3 +30,27 @@ class TestBooleanConfigPrecedence(TestCase):
|
|||||||
config.save()
|
config.save()
|
||||||
|
|
||||||
self.assertTrue(AIConfig().ai_enabled)
|
self.assertTrue(AIConfig().ai_enabled)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAIConfigPrecedence(TestCase):
|
||||||
|
@override_settings(LLM_EMBEDDING_API_KEY="environment-embedding-key")
|
||||||
|
def test_database_embedding_api_key_overrides_environment_setting(self) -> None:
|
||||||
|
config, _ = ApplicationConfiguration.objects.get_or_create()
|
||||||
|
config.llm_embedding_api_key = "database-embedding-key"
|
||||||
|
config.save()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
AIConfig().llm_embedding_api_key,
|
||||||
|
"database-embedding-key",
|
||||||
|
)
|
||||||
|
|
||||||
|
@override_settings(LLM_EMBEDDING_API_KEY="environment-embedding-key")
|
||||||
|
def test_null_embedding_api_key_uses_environment_setting(self) -> None:
|
||||||
|
config, _ = ApplicationConfiguration.objects.get_or_create()
|
||||||
|
config.llm_embedding_api_key = None
|
||||||
|
config.save()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
AIConfig().llm_embedding_api_key,
|
||||||
|
"environment-embedding-key",
|
||||||
|
)
|
||||||
|
|||||||
+79
-1456
File diff suppressed because it is too large
Load Diff
@@ -1,375 +0,0 @@
|
|||||||
import ipaddress
|
|
||||||
import os
|
|
||||||
|
|
||||||
import httpcore
|
|
||||||
import httpx
|
|
||||||
import pytest
|
|
||||||
from pytest_mock import MockerFixture
|
|
||||||
|
|
||||||
from paperless.network import GuardedAsyncHTTPTransport
|
|
||||||
from paperless.network import GuardedHTTPTransport
|
|
||||||
from paperless.network import OutboundRequestBlockedError
|
|
||||||
from paperless.network import create_guarded_httpx_client
|
|
||||||
from paperless_testing.outbound import DialRecorder
|
|
||||||
from paperless_testing.outbound import FakeDNS
|
|
||||||
from paperless_testing.outbound import LocalHTTPServer
|
|
||||||
from paperless_testing.outbound import running_http_server
|
|
||||||
|
|
||||||
|
|
||||||
class TestGuardedTransportSync:
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_pinned_connection_falls_back_to_next_address(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A hostname resolving to ::1 then 127.0.0.1
|
|
||||||
- A server listening on 127.0.0.1 only
|
|
||||||
- Internal addresses disallowed, with loopback treated as public
|
|
||||||
WHEN:
|
|
||||||
- A request is made
|
|
||||||
THEN:
|
|
||||||
- ::1 fails, 127.0.0.1 is dialled next and the request succeeds
|
|
||||||
"""
|
|
||||||
fake_dns.add("dual-stack.test", "::1", "127.0.0.1")
|
|
||||||
|
|
||||||
with httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
response = client.get(f"http://dual-stack.test:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert dial_recorder.hosts() == ["::1", "127.0.0.1"]
|
|
||||||
|
|
||||||
def test_allow_internal_uses_stock_resolution(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Internal addresses allowed
|
|
||||||
WHEN:
|
|
||||||
- A request is made to localhost
|
|
||||||
THEN:
|
|
||||||
- It succeeds without the guard resolving anything
|
|
||||||
"""
|
|
||||||
with httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=True),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
response = client.get(f"http://localhost:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert fake_dns.lookups == []
|
|
||||||
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_host_header_is_the_hostname(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A pinned connection to a named host
|
|
||||||
WHEN:
|
|
||||||
- A request is made
|
|
||||||
THEN:
|
|
||||||
- The server receives the hostname in Host, not the dialled IP
|
|
||||||
"""
|
|
||||||
fake_dns.add("pinned.test", "127.0.0.1")
|
|
||||||
|
|
||||||
with httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
client.get(f"http://pinned.test:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert local_http_server.requests[0].headers["host"] == (
|
|
||||||
f"pinned.test:{local_http_server.port}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_redirect_to_internal_host_is_blocked(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- An allowed origin that redirects to a host resolving to a blocked
|
|
||||||
address, and a client that follows redirects
|
|
||||||
WHEN:
|
|
||||||
- The origin is requested
|
|
||||||
THEN:
|
|
||||||
- The redirect hop is blocked without dialling the blocked address
|
|
||||||
"""
|
|
||||||
allowed = ipaddress.ip_address("127.0.0.1")
|
|
||||||
mocker.patch(
|
|
||||||
"paperless.network.is_public_ip",
|
|
||||||
side_effect=lambda address: address == allowed,
|
|
||||||
)
|
|
||||||
fake_dns.add("origin.test", "127.0.0.1")
|
|
||||||
fake_dns.add("internal.test", "127.0.0.2")
|
|
||||||
local_http_server.redirect_to = (
|
|
||||||
f"http://internal.test:{local_http_server.port}/"
|
|
||||||
)
|
|
||||||
|
|
||||||
with (
|
|
||||||
httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
follow_redirects=True,
|
|
||||||
) as client,
|
|
||||||
pytest.raises(OutboundRequestBlockedError) as exc_info,
|
|
||||||
):
|
|
||||||
client.get(f"http://origin.test:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert exc_info.value.address == ipaddress.ip_address("127.0.0.2")
|
|
||||||
assert dial_recorder.hosts() == ["127.0.0.1"]
|
|
||||||
assert len(local_http_server.requests) == 1
|
|
||||||
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_connections_are_not_shared_between_hosts_on_one_address(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Two hostnames resolving to the same address
|
|
||||||
- Internal addresses disallowed, with loopback treated as public
|
|
||||||
WHEN:
|
|
||||||
- One client requests the first host twice, then the second host
|
|
||||||
THEN:
|
|
||||||
- The first host's connection is reused for its second request
|
|
||||||
- The second host gets its own connection, so its certificate would
|
|
||||||
be checked rather than inheriting the first host's session
|
|
||||||
"""
|
|
||||||
fake_dns.add("first.test", "127.0.0.1")
|
|
||||||
fake_dns.add("second.test", "127.0.0.1")
|
|
||||||
|
|
||||||
with httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
client.get(f"http://first.test:{local_http_server.port}/")
|
|
||||||
client.get(f"http://first.test:{local_http_server.port}/")
|
|
||||||
client.get(f"http://second.test:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert dial_recorder.hosts() == ["127.0.0.1", "127.0.0.1"]
|
|
||||||
assert local_http_server.connections == 2
|
|
||||||
assert [request.headers["host"] for request in local_http_server.requests] == [
|
|
||||||
f"first.test:{local_http_server.port}",
|
|
||||||
f"first.test:{local_http_server.port}",
|
|
||||||
f"second.test:{local_http_server.port}",
|
|
||||||
]
|
|
||||||
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_tls_uses_the_hostname_not_the_dialled_address(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A pinned HTTPS connection to a named host
|
|
||||||
- A plain HTTP server, so the handshake itself fails
|
|
||||||
WHEN:
|
|
||||||
- A request is made
|
|
||||||
THEN:
|
|
||||||
- The validated address is dialled
|
|
||||||
- TLS is started with the hostname for SNI and certificate checks
|
|
||||||
"""
|
|
||||||
fake_dns.add("pinned.test", "127.0.0.1")
|
|
||||||
start_tls = mocker.spy(httpcore._backends.sync.SyncStream, "start_tls")
|
|
||||||
|
|
||||||
with (
|
|
||||||
httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client,
|
|
||||||
pytest.raises(httpx.ConnectError),
|
|
||||||
):
|
|
||||||
client.get(f"https://pinned.test:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert dial_recorder.hosts() == ["127.0.0.1"]
|
|
||||||
start_tls.assert_called_once()
|
|
||||||
assert start_tls.call_args.kwargs["server_hostname"] == "pinned.test"
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"host",
|
|
||||||
[
|
|
||||||
pytest.param("localhost", id="name"),
|
|
||||||
pytest.param("2130706433", id="decimal"),
|
|
||||||
pytest.param("0x7f.1", id="hex-short"),
|
|
||||||
pytest.param("127.1", id="short-dotted"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_blocks_internal_host_without_connecting(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
host: str,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Internal addresses disallowed
|
|
||||||
- A URL whose host reaches loopback, by name or by a
|
|
||||||
non-canonical spelling of 127.0.0.1
|
|
||||||
WHEN:
|
|
||||||
- A request is made through the transport
|
|
||||||
THEN:
|
|
||||||
- The resolved address is checked, the request is blocked and the
|
|
||||||
server never sees a connection
|
|
||||||
"""
|
|
||||||
with (
|
|
||||||
httpx.Client(
|
|
||||||
transport=GuardedHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client,
|
|
||||||
pytest.raises(OutboundRequestBlockedError),
|
|
||||||
):
|
|
||||||
client.get(f"http://{host}:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert local_http_server.connections == 0
|
|
||||||
assert dial_recorder.hosts() == []
|
|
||||||
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
def test_environment_proxy_is_not_used(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Proxy variables in the environment pointing at a second local server
|
|
||||||
- Internal addresses disallowed
|
|
||||||
WHEN:
|
|
||||||
- A request is made through the production client factory to an
|
|
||||||
allowed origin
|
|
||||||
THEN:
|
|
||||||
- The origin server receives the request directly and the proxy
|
|
||||||
server never sees a connection
|
|
||||||
"""
|
|
||||||
with running_http_server() as proxy_server:
|
|
||||||
mocker.patch.dict(
|
|
||||||
os.environ,
|
|
||||||
{
|
|
||||||
"HTTP_PROXY": f"http://127.0.0.1:{proxy_server.port}",
|
|
||||||
"HTTPS_PROXY": f"http://127.0.0.1:{proxy_server.port}",
|
|
||||||
"ALL_PROXY": f"http://127.0.0.1:{proxy_server.port}",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
fake_dns.add("origin.test", "127.0.0.1")
|
|
||||||
|
|
||||||
url = f"http://origin.test:{local_http_server.port}/"
|
|
||||||
with create_guarded_httpx_client(
|
|
||||||
url,
|
|
||||||
allow_internal=False,
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
response = client.get(url)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert len(local_http_server.requests) == 1
|
|
||||||
assert local_http_server.requests[0].headers["host"] == (
|
|
||||||
f"origin.test:{local_http_server.port}"
|
|
||||||
)
|
|
||||||
assert proxy_server.connections == 0
|
|
||||||
assert proxy_server.requests == []
|
|
||||||
assert dial_recorder.hosts() == ["127.0.0.1"]
|
|
||||||
|
|
||||||
|
|
||||||
class TestGuardedTransportAsync:
|
|
||||||
@pytest.fixture(autouse=True)
|
|
||||||
def anyio_backend(self) -> str:
|
|
||||||
return "asyncio"
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
@pytest.mark.usefixtures("every_address_is_public")
|
|
||||||
async def test_pinned_connection_falls_back_to_next_address(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A hostname resolving to ::1 then 127.0.0.1
|
|
||||||
- A server listening on 127.0.0.1 only
|
|
||||||
- Internal addresses disallowed, with loopback treated as public
|
|
||||||
WHEN:
|
|
||||||
- An async request is made
|
|
||||||
THEN:
|
|
||||||
- ::1 fails, 127.0.0.1 is dialled next and the request succeeds
|
|
||||||
"""
|
|
||||||
fake_dns.add("dual-stack.test", "::1", "127.0.0.1")
|
|
||||||
|
|
||||||
async with httpx.AsyncClient(
|
|
||||||
transport=GuardedAsyncHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
response = await client.get(
|
|
||||||
f"http://dual-stack.test:{local_http_server.port}/",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert dial_recorder.hosts() == ["::1", "127.0.0.1"]
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_allow_internal_uses_stock_resolution(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
fake_dns: FakeDNS,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Internal addresses allowed
|
|
||||||
WHEN:
|
|
||||||
- An async request is made to localhost
|
|
||||||
THEN:
|
|
||||||
- It succeeds without the guard resolving anything
|
|
||||||
"""
|
|
||||||
async with httpx.AsyncClient(
|
|
||||||
transport=GuardedAsyncHTTPTransport(allow_internal=True),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
response = await client.get(f"http://localhost:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert fake_dns.lookups == []
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_blocks_internal_host_without_connecting(
|
|
||||||
self,
|
|
||||||
local_http_server: LocalHTTPServer,
|
|
||||||
dial_recorder: DialRecorder,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- Internal addresses disallowed
|
|
||||||
WHEN:
|
|
||||||
- An async request is made to localhost through the transport
|
|
||||||
THEN:
|
|
||||||
- It is blocked and the server never sees a connection
|
|
||||||
"""
|
|
||||||
async with httpx.AsyncClient(
|
|
||||||
transport=GuardedAsyncHTTPTransport(allow_internal=False),
|
|
||||||
timeout=5.0,
|
|
||||||
) as client:
|
|
||||||
with pytest.raises(OutboundRequestBlockedError):
|
|
||||||
await client.get(f"http://localhost:{local_http_server.port}/")
|
|
||||||
|
|
||||||
assert local_http_server.connections == 0
|
|
||||||
assert dial_recorder.hosts() == []
|
|
||||||
@@ -1,4 +1,5 @@
|
|||||||
from allauth.account import views as allauth_account_views
|
from allauth.account import views as allauth_account_views
|
||||||
|
from allauth.account.decorators import secure_admin_login
|
||||||
from allauth.mfa.base import views as allauth_mfa_views
|
from allauth.mfa.base import views as allauth_mfa_views
|
||||||
from allauth.socialaccount import views as allauth_social_account_views
|
from allauth.socialaccount import views as allauth_social_account_views
|
||||||
from allauth.urls import build_provider_urlpatterns
|
from allauth.urls import build_provider_urlpatterns
|
||||||
@@ -68,6 +69,8 @@ from paperless_mail.views import MailRuleViewSet
|
|||||||
from paperless_mail.views import OauthCallbackView
|
from paperless_mail.views import OauthCallbackView
|
||||||
from paperless_mail.views import ProcessedMailViewSet
|
from paperless_mail.views import ProcessedMailViewSet
|
||||||
|
|
||||||
|
admin.site.login = secure_admin_login(admin.site.login)
|
||||||
|
|
||||||
api_router = DefaultRouter()
|
api_router = DefaultRouter()
|
||||||
api_router.register(r"correspondents", CorrespondentViewSet)
|
api_router.register(r"correspondents", CorrespondentViewSet)
|
||||||
api_router.register(r"document_types", DocumentTypeViewSet)
|
api_router.register(r"document_types", DocumentTypeViewSet)
|
||||||
@@ -297,7 +300,7 @@ urlpatterns = [
|
|||||||
),
|
),
|
||||||
re_path(r"^share/(?P<slug>\w+)/?$", SharedLinkView.as_view()),
|
re_path(r"^share/(?P<slug>\w+)/?$", SharedLinkView.as_view()),
|
||||||
re_path(r"^favicon.ico$", FaviconView.as_view(), name="favicon"),
|
re_path(r"^favicon.ico$", FaviconView.as_view(), name="favicon"),
|
||||||
re_path(r"admin/", admin.site.urls),
|
re_path(r"^admin/", admin.site.urls),
|
||||||
re_path(
|
re_path(
|
||||||
r"^fetch/",
|
r"^fetch/",
|
||||||
include(
|
include(
|
||||||
|
|||||||
+10
-29
@@ -14,16 +14,14 @@ if TYPE_CHECKING:
|
|||||||
from llama_index.llms.openai_like import OpenAILike
|
from llama_index.llms.openai_like import OpenAILike
|
||||||
|
|
||||||
from paperless.config import AIConfig
|
from paperless.config import AIConfig
|
||||||
from paperless.network import GuardedAsyncHTTPTransport
|
from paperless.network import PinnedHostAsyncHTTPTransport
|
||||||
from paperless.network import GuardedHTTPTransport
|
from paperless.network import PinnedHostHTTPTransport
|
||||||
from paperless.network import OutboundRequestBlockedError
|
from paperless.network import create_pinned_async_httpx_client
|
||||||
from paperless.network import create_guarded_async_httpx_client
|
from paperless.network import create_pinned_httpx_client
|
||||||
from paperless.network import create_guarded_httpx_client
|
|
||||||
from paperless.network import validate_outbound_http_url
|
from paperless.network import validate_outbound_http_url
|
||||||
from paperless_ai.base_model import ClassificationSuggestions
|
from paperless_ai.base_model import ClassificationSuggestions
|
||||||
from paperless_ai.base_model import DocumentClassifierSchema
|
from paperless_ai.base_model import DocumentClassifierSchema
|
||||||
from paperless_ai.base_model import model_to_classification_suggestions
|
from paperless_ai.base_model import model_to_classification_suggestions
|
||||||
from paperless_ai.exceptions import LLMBlockedError
|
|
||||||
from paperless_ai.exceptions import LLMProviderError
|
from paperless_ai.exceptions import LLMProviderError
|
||||||
from paperless_ai.exceptions import LLMTimeoutError
|
from paperless_ai.exceptions import LLMTimeoutError
|
||||||
|
|
||||||
@@ -45,19 +43,6 @@ LLM_SYSTEM_PROMPT = (
|
|||||||
PLACEHOLDER_API_KEY: Final = "fake"
|
PLACEHOLDER_API_KEY: Final = "fake"
|
||||||
|
|
||||||
|
|
||||||
def _find_blocked_cause(exc: BaseException) -> OutboundRequestBlockedError | None:
|
|
||||||
# The openai SDK wraps transport errors in APIConnectionError, so the
|
|
||||||
# block can sit anywhere in the __cause__ chain.
|
|
||||||
current: BaseException | None = exc
|
|
||||||
seen: set[int] = set()
|
|
||||||
while current is not None and id(current) not in seen:
|
|
||||||
if isinstance(current, OutboundRequestBlockedError):
|
|
||||||
return current
|
|
||||||
seen.add(id(current))
|
|
||||||
current = current.__cause__
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class AIClient:
|
class AIClient:
|
||||||
"""
|
"""
|
||||||
A client for interacting with an LLM backend.
|
A client for interacting with an LLM backend.
|
||||||
@@ -78,10 +63,10 @@ class AIClient:
|
|||||||
endpoint,
|
endpoint,
|
||||||
allow_internal=self.settings.llm_allow_internal_endpoints,
|
allow_internal=self.settings.llm_allow_internal_endpoints,
|
||||||
)
|
)
|
||||||
transport = GuardedHTTPTransport(
|
transport = PinnedHostHTTPTransport(
|
||||||
allow_internal=self.settings.llm_allow_internal_endpoints,
|
allow_internal=self.settings.llm_allow_internal_endpoints,
|
||||||
)
|
)
|
||||||
async_transport = GuardedAsyncHTTPTransport(
|
async_transport = PinnedHostAsyncHTTPTransport(
|
||||||
allow_internal=self.settings.llm_allow_internal_endpoints,
|
allow_internal=self.settings.llm_allow_internal_endpoints,
|
||||||
)
|
)
|
||||||
return Ollama(
|
return Ollama(
|
||||||
@@ -90,6 +75,7 @@ class AIClient:
|
|||||||
context_window=self.settings.llm_context_size,
|
context_window=self.settings.llm_context_size,
|
||||||
request_timeout=self.settings.llm_request_timeout,
|
request_timeout=self.settings.llm_request_timeout,
|
||||||
system_prompt=LLM_SYSTEM_PROMPT,
|
system_prompt=LLM_SYSTEM_PROMPT,
|
||||||
|
additional_kwargs=self.settings.llm_extra_params,
|
||||||
client=Client(
|
client=Client(
|
||||||
host=endpoint,
|
host=endpoint,
|
||||||
timeout=self.settings.llm_request_timeout,
|
timeout=self.settings.llm_request_timeout,
|
||||||
@@ -108,12 +94,12 @@ class AIClient:
|
|||||||
http_client = None
|
http_client = None
|
||||||
async_http_client = None
|
async_http_client = None
|
||||||
if endpoint:
|
if endpoint:
|
||||||
http_client = create_guarded_httpx_client(
|
http_client = create_pinned_httpx_client(
|
||||||
endpoint,
|
endpoint,
|
||||||
allow_internal=self.settings.llm_allow_internal_endpoints,
|
allow_internal=self.settings.llm_allow_internal_endpoints,
|
||||||
timeout=self.settings.llm_request_timeout,
|
timeout=self.settings.llm_request_timeout,
|
||||||
)
|
)
|
||||||
async_http_client = create_guarded_async_httpx_client(
|
async_http_client = create_pinned_async_httpx_client(
|
||||||
endpoint,
|
endpoint,
|
||||||
allow_internal=self.settings.llm_allow_internal_endpoints,
|
allow_internal=self.settings.llm_allow_internal_endpoints,
|
||||||
timeout=self.settings.llm_request_timeout,
|
timeout=self.settings.llm_request_timeout,
|
||||||
@@ -126,6 +112,7 @@ class AIClient:
|
|||||||
is_chat_model=True,
|
is_chat_model=True,
|
||||||
is_function_calling_model=True,
|
is_function_calling_model=True,
|
||||||
system_prompt=LLM_SYSTEM_PROMPT,
|
system_prompt=LLM_SYSTEM_PROMPT,
|
||||||
|
additional_kwargs=self.settings.llm_extra_params,
|
||||||
http_client=http_client,
|
http_client=http_client,
|
||||||
async_http_client=async_http_client,
|
async_http_client=async_http_client,
|
||||||
)
|
)
|
||||||
@@ -194,12 +181,6 @@ class AIClient:
|
|||||||
except httpx.TimeoutException as exc:
|
except httpx.TimeoutException as exc:
|
||||||
raise LLMTimeoutError from exc
|
raise LLMTimeoutError from exc
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
blocked = _find_blocked_cause(exc)
|
|
||||||
if blocked is not None:
|
|
||||||
raise LLMBlockedError(
|
|
||||||
"AI backend request was blocked by the outbound request "
|
|
||||||
f"policy: {blocked}",
|
|
||||||
) from exc
|
|
||||||
if self._is_openai_timeout(exc):
|
if self._is_openai_timeout(exc):
|
||||||
raise LLMTimeoutError from exc
|
raise LLMTimeoutError from exc
|
||||||
if self._is_provider_error(exc):
|
if self._is_provider_error(exc):
|
||||||
|
|||||||
@@ -9,10 +9,10 @@ if TYPE_CHECKING:
|
|||||||
from documents.models import Document
|
from documents.models import Document
|
||||||
from paperless.config import AIConfig
|
from paperless.config import AIConfig
|
||||||
from paperless.models import LLMEmbeddingBackend
|
from paperless.models import LLMEmbeddingBackend
|
||||||
from paperless.network import GuardedAsyncHTTPTransport
|
from paperless.network import PinnedHostAsyncHTTPTransport
|
||||||
from paperless.network import GuardedHTTPTransport
|
from paperless.network import PinnedHostHTTPTransport
|
||||||
from paperless.network import create_guarded_async_httpx_client
|
from paperless.network import create_pinned_async_httpx_client
|
||||||
from paperless.network import create_guarded_httpx_client
|
from paperless.network import create_pinned_httpx_client
|
||||||
from paperless.network import validate_outbound_http_url
|
from paperless.network import validate_outbound_http_url
|
||||||
from paperless_ai.client import PLACEHOLDER_API_KEY
|
from paperless_ai.client import PLACEHOLDER_API_KEY
|
||||||
|
|
||||||
@@ -29,19 +29,21 @@ def get_embedding_model(config: AIConfig) -> "BaseEmbedding":
|
|||||||
http_client = None
|
http_client = None
|
||||||
async_http_client = None
|
async_http_client = None
|
||||||
if endpoint:
|
if endpoint:
|
||||||
http_client = create_guarded_httpx_client(
|
http_client = create_pinned_httpx_client(
|
||||||
endpoint,
|
endpoint,
|
||||||
allow_internal=config.llm_allow_internal_endpoints,
|
allow_internal=config.llm_allow_internal_endpoints,
|
||||||
timeout=config.llm_request_timeout,
|
timeout=config.llm_request_timeout,
|
||||||
)
|
)
|
||||||
async_http_client = create_guarded_async_httpx_client(
|
async_http_client = create_pinned_async_httpx_client(
|
||||||
endpoint,
|
endpoint,
|
||||||
allow_internal=config.llm_allow_internal_endpoints,
|
allow_internal=config.llm_allow_internal_endpoints,
|
||||||
timeout=config.llm_request_timeout,
|
timeout=config.llm_request_timeout,
|
||||||
)
|
)
|
||||||
return OpenAILikeEmbedding(
|
return OpenAILikeEmbedding(
|
||||||
model_name=config.llm_embedding_model or "text-embedding-3-small",
|
model_name=config.llm_embedding_model or "text-embedding-3-small",
|
||||||
api_key=config.llm_api_key or PLACEHOLDER_API_KEY,
|
api_key=config.llm_embedding_api_key
|
||||||
|
or config.llm_api_key
|
||||||
|
or PLACEHOLDER_API_KEY,
|
||||||
api_base=endpoint,
|
api_base=endpoint,
|
||||||
timeout=config.llm_request_timeout,
|
timeout=config.llm_request_timeout,
|
||||||
http_client=http_client,
|
http_client=http_client,
|
||||||
@@ -77,14 +79,14 @@ def get_embedding_model(config: AIConfig) -> "BaseEmbedding":
|
|||||||
embedding._client = Client(
|
embedding._client = Client(
|
||||||
host=endpoint,
|
host=endpoint,
|
||||||
timeout=config.llm_request_timeout,
|
timeout=config.llm_request_timeout,
|
||||||
transport=GuardedHTTPTransport(
|
transport=PinnedHostHTTPTransport(
|
||||||
allow_internal=config.llm_allow_internal_endpoints,
|
allow_internal=config.llm_allow_internal_endpoints,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
embedding._async_client = AsyncClient(
|
embedding._async_client = AsyncClient(
|
||||||
host=endpoint,
|
host=endpoint,
|
||||||
timeout=config.llm_request_timeout,
|
timeout=config.llm_request_timeout,
|
||||||
transport=GuardedAsyncHTTPTransport(
|
transport=PinnedHostAsyncHTTPTransport(
|
||||||
allow_internal=config.llm_allow_internal_endpoints,
|
allow_internal=config.llm_allow_internal_endpoints,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4,7 +4,3 @@ class LLMTimeoutError(Exception):
|
|||||||
|
|
||||||
class LLMProviderError(Exception):
|
class LLMProviderError(Exception):
|
||||||
"""The LLM backend rejected the request."""
|
"""The LLM backend rejected the request."""
|
||||||
|
|
||||||
|
|
||||||
class LLMBlockedError(Exception):
|
|
||||||
"""The outbound request policy refused the connection to the LLM backend."""
|
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
import ipaddress
|
|
||||||
import json
|
import json
|
||||||
from unittest.mock import ANY
|
from unittest.mock import ANY
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
@@ -10,15 +9,11 @@ import openai
|
|||||||
import pytest
|
import pytest
|
||||||
from llama_index.core.llms.llm import ToolSelection
|
from llama_index.core.llms.llm import ToolSelection
|
||||||
|
|
||||||
from paperless.network import BlockReason
|
|
||||||
from paperless.network import OutboundRequestBlockedError
|
|
||||||
from paperless_ai.client import LLM_SYSTEM_PROMPT
|
from paperless_ai.client import LLM_SYSTEM_PROMPT
|
||||||
from paperless_ai.client import PLACEHOLDER_API_KEY
|
from paperless_ai.client import PLACEHOLDER_API_KEY
|
||||||
from paperless_ai.client import AIClient
|
from paperless_ai.client import AIClient
|
||||||
from paperless_ai.exceptions import LLMBlockedError
|
|
||||||
from paperless_ai.exceptions import LLMProviderError
|
from paperless_ai.exceptions import LLMProviderError
|
||||||
from paperless_ai.exceptions import LLMTimeoutError
|
from paperless_ai.exceptions import LLMTimeoutError
|
||||||
from paperless_testing.outbound import guard_of
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -28,6 +23,7 @@ def mock_ai_config():
|
|||||||
mock_config.llm_allow_internal_endpoints = True
|
mock_config.llm_allow_internal_endpoints = True
|
||||||
mock_config.llm_context_size = 8192
|
mock_config.llm_context_size = 8192
|
||||||
mock_config.llm_request_timeout = 120
|
mock_config.llm_request_timeout = 120
|
||||||
|
mock_config.llm_extra_params = {}
|
||||||
MockAIConfig.return_value = mock_config
|
MockAIConfig.return_value = mock_config
|
||||||
yield mock_config
|
yield mock_config
|
||||||
|
|
||||||
@@ -57,6 +53,7 @@ def test_get_llm_ollama(mock_ai_config, mock_ollama_llm):
|
|||||||
context_window=8192,
|
context_window=8192,
|
||||||
request_timeout=120,
|
request_timeout=120,
|
||||||
system_prompt=LLM_SYSTEM_PROMPT,
|
system_prompt=LLM_SYSTEM_PROMPT,
|
||||||
|
additional_kwargs={},
|
||||||
client=ANY,
|
client=ANY,
|
||||||
async_client=ANY,
|
async_client=ANY,
|
||||||
)
|
)
|
||||||
@@ -79,6 +76,7 @@ def test_get_llm_openai(mock_ai_config, mock_openai_llm):
|
|||||||
is_chat_model=True,
|
is_chat_model=True,
|
||||||
is_function_calling_model=True,
|
is_function_calling_model=True,
|
||||||
system_prompt=LLM_SYSTEM_PROMPT,
|
system_prompt=LLM_SYSTEM_PROMPT,
|
||||||
|
additional_kwargs={},
|
||||||
http_client=ANY,
|
http_client=ANY,
|
||||||
async_http_client=ANY,
|
async_http_client=ANY,
|
||||||
)
|
)
|
||||||
@@ -201,6 +199,36 @@ def test_run_llm_query_openai_uses_tools(mock_ai_config, mock_openai_llm):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("backend", "llm_fixture"),
|
||||||
|
[
|
||||||
|
pytest.param("openai-like", "mock_openai_llm", id="openai-like"),
|
||||||
|
pytest.param("ollama", "mock_ollama_llm", id="ollama"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_get_llm_passes_extra_params(request, mock_ai_config, backend, llm_fixture):
|
||||||
|
"""
|
||||||
|
GIVEN:
|
||||||
|
- Extra LLM params configured, e.g. for a provider that needs a
|
||||||
|
parameter we do not set ourselves
|
||||||
|
WHEN:
|
||||||
|
- The client builds the LLM
|
||||||
|
THEN:
|
||||||
|
- They are handed to the backend as additional_kwargs
|
||||||
|
"""
|
||||||
|
llm_mock = request.getfixturevalue(llm_fixture)
|
||||||
|
mock_ai_config.llm_backend = backend
|
||||||
|
mock_ai_config.llm_model = "gpt-5.6-luna"
|
||||||
|
mock_ai_config.llm_endpoint = "http://test-url"
|
||||||
|
mock_ai_config.llm_extra_params = {"reasoning_effort": "none"}
|
||||||
|
|
||||||
|
AIClient()
|
||||||
|
|
||||||
|
assert llm_mock.call_args.kwargs["additional_kwargs"] == {
|
||||||
|
"reasoning_effort": "none",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def test_run_llm_query_openai_timeout_raises_local_error(
|
def test_run_llm_query_openai_timeout_raises_local_error(
|
||||||
mock_ai_config,
|
mock_ai_config,
|
||||||
mock_openai_llm,
|
mock_openai_llm,
|
||||||
@@ -282,142 +310,3 @@ def test_run_llm_query_httpx_timeout_raises_local_error(
|
|||||||
|
|
||||||
with pytest.raises(LLMTimeoutError):
|
with pytest.raises(LLMTimeoutError):
|
||||||
client.run_llm_query("test_prompt")
|
client.run_llm_query("test_prompt")
|
||||||
|
|
||||||
|
|
||||||
class TestGuardedLLMClients:
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
("endpoint", "allow_internal"),
|
|
||||||
[
|
|
||||||
pytest.param("http://test-url", True, id="internal-allowed"),
|
|
||||||
pytest.param("http://93.184.216.34:11434", False, id="internal-blocked"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_ollama_clients_are_guarded(
|
|
||||||
self,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
mock_ollama_llm: MagicMock,
|
|
||||||
endpoint: str,
|
|
||||||
*,
|
|
||||||
allow_internal: bool,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The Ollama backend
|
|
||||||
WHEN:
|
|
||||||
- The LLM is built
|
|
||||||
THEN:
|
|
||||||
- Its sync and async clients use guarded transports with the setting
|
|
||||||
"""
|
|
||||||
mock_ai_config.llm_backend = "ollama"
|
|
||||||
mock_ai_config.llm_model = "test_model"
|
|
||||||
mock_ai_config.llm_endpoint = endpoint
|
|
||||||
mock_ai_config.llm_allow_internal_endpoints = allow_internal
|
|
||||||
|
|
||||||
AIClient()
|
|
||||||
|
|
||||||
kwargs = mock_ollama_llm.call_args.kwargs
|
|
||||||
assert guard_of(kwargs["client"]._client)._allow_internal is allow_internal
|
|
||||||
assert (
|
|
||||||
guard_of(kwargs["async_client"]._client)._allow_internal is allow_internal
|
|
||||||
)
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
("endpoint", "allow_internal"),
|
|
||||||
[
|
|
||||||
pytest.param("http://test-url", True, id="internal-allowed"),
|
|
||||||
pytest.param("http://93.184.216.34:8080", False, id="internal-blocked"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_openai_like_clients_are_guarded(
|
|
||||||
self,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
mock_openai_llm: MagicMock,
|
|
||||||
endpoint: str,
|
|
||||||
*,
|
|
||||||
allow_internal: bool,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The OpenAI-like backend with an endpoint
|
|
||||||
WHEN:
|
|
||||||
- The LLM is built
|
|
||||||
THEN:
|
|
||||||
- Its sync and async http clients use guarded transports
|
|
||||||
"""
|
|
||||||
mock_ai_config.llm_backend = "openai-like"
|
|
||||||
mock_ai_config.llm_model = "test_model"
|
|
||||||
mock_ai_config.llm_api_key = "key"
|
|
||||||
mock_ai_config.llm_endpoint = endpoint
|
|
||||||
mock_ai_config.llm_allow_internal_endpoints = allow_internal
|
|
||||||
|
|
||||||
AIClient()
|
|
||||||
|
|
||||||
kwargs = mock_openai_llm.call_args.kwargs
|
|
||||||
assert guard_of(kwargs["http_client"])._allow_internal is allow_internal
|
|
||||||
assert guard_of(kwargs["async_http_client"])._allow_internal is allow_internal
|
|
||||||
|
|
||||||
|
|
||||||
def _block() -> OutboundRequestBlockedError:
|
|
||||||
return OutboundRequestBlockedError(
|
|
||||||
host="llm.example",
|
|
||||||
port=443,
|
|
||||||
reason=BlockReason.NON_PUBLIC_ADDRESS,
|
|
||||||
address=ipaddress.ip_address("10.0.0.1"),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class TestBlockedLLMRequests:
|
|
||||||
def test_ollama_block_becomes_llm_blocked_error(
|
|
||||||
self,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
mock_ollama_llm: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The Ollama backend and a connection blocked by policy
|
|
||||||
WHEN:
|
|
||||||
- An LLM query runs
|
|
||||||
THEN:
|
|
||||||
- LLMBlockedError is raised with a message, chained to the block
|
|
||||||
- The message, which tracked tasks store, names the destination but
|
|
||||||
not the resolved internal address
|
|
||||||
"""
|
|
||||||
mock_ai_config.llm_backend = "ollama"
|
|
||||||
mock_ai_config.llm_model = "test_model"
|
|
||||||
mock_ai_config.llm_endpoint = "http://test-url"
|
|
||||||
block = _block()
|
|
||||||
mock_ollama_llm.return_value.chat.side_effect = block
|
|
||||||
|
|
||||||
with pytest.raises(LLMBlockedError) as exc_info:
|
|
||||||
AIClient().run_llm_query("test_prompt")
|
|
||||||
|
|
||||||
assert exc_info.value.__cause__ is block
|
|
||||||
assert "llm.example:443" in str(exc_info.value)
|
|
||||||
assert "10.0.0.1" not in str(exc_info.value)
|
|
||||||
|
|
||||||
def test_openai_wrapped_block_becomes_llm_blocked_error(
|
|
||||||
self,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
mock_openai_llm: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The OpenAI-like backend, whose SDK wraps the block in
|
|
||||||
APIConnectionError
|
|
||||||
WHEN:
|
|
||||||
- An LLM query runs
|
|
||||||
THEN:
|
|
||||||
- LLMBlockedError is raised
|
|
||||||
"""
|
|
||||||
mock_ai_config.llm_backend = "openai-like"
|
|
||||||
mock_ai_config.llm_model = "test_model"
|
|
||||||
mock_ai_config.llm_api_key = "key"
|
|
||||||
mock_ai_config.llm_endpoint = "http://test-url"
|
|
||||||
wrapped = openai.APIConnectionError(
|
|
||||||
request=httpx.Request("POST", "http://test-url/v1/chat/completions"),
|
|
||||||
)
|
|
||||||
wrapped.__cause__ = _block()
|
|
||||||
mock_openai_llm.return_value.chat_with_tools.side_effect = wrapped
|
|
||||||
|
|
||||||
with pytest.raises(LLMBlockedError):
|
|
||||||
AIClient().run_llm_query("test_prompt")
|
|
||||||
|
|||||||
@@ -1,12 +1,9 @@
|
|||||||
from typing import TYPE_CHECKING
|
|
||||||
from typing import cast
|
|
||||||
from unittest.mock import ANY
|
from unittest.mock import ANY
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
from pytest_mock import MockerFixture
|
|
||||||
|
|
||||||
from documents.models import Document
|
from documents.models import Document
|
||||||
from paperless.models import LLMEmbeddingBackend
|
from paperless.models import LLMEmbeddingBackend
|
||||||
@@ -15,15 +12,12 @@ from paperless_ai.embedding import _normalize_llm_index_text
|
|||||||
from paperless_ai.embedding import build_llm_index_text
|
from paperless_ai.embedding import build_llm_index_text
|
||||||
from paperless_ai.embedding import get_configured_model_name
|
from paperless_ai.embedding import get_configured_model_name
|
||||||
from paperless_ai.embedding import get_embedding_model
|
from paperless_ai.embedding import get_embedding_model
|
||||||
from paperless_testing.outbound import guard_of
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from llama_index.embeddings.ollama import OllamaEmbedding
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_ai_config():
|
def mock_ai_config():
|
||||||
with patch("paperless_ai.embedding.AIConfig") as MockAIConfig:
|
with patch("paperless_ai.embedding.AIConfig") as MockAIConfig:
|
||||||
|
MockAIConfig.return_value.llm_embedding_api_key = None
|
||||||
MockAIConfig.return_value.llm_embedding_endpoint = None
|
MockAIConfig.return_value.llm_embedding_endpoint = None
|
||||||
MockAIConfig.return_value.llm_allow_internal_endpoints = True
|
MockAIConfig.return_value.llm_allow_internal_endpoints = True
|
||||||
MockAIConfig.return_value.llm_context_size = 8192
|
MockAIConfig.return_value.llm_context_size = 8192
|
||||||
@@ -70,6 +64,7 @@ def mock_document():
|
|||||||
def test_get_embedding_model_openai(mock_ai_config):
|
def test_get_embedding_model_openai(mock_ai_config):
|
||||||
mock_ai_config.return_value.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
mock_ai_config.return_value.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
||||||
mock_ai_config.return_value.llm_embedding_model = "text-embedding-3-small"
|
mock_ai_config.return_value.llm_embedding_model = "text-embedding-3-small"
|
||||||
|
mock_ai_config.return_value.llm_embedding_api_key = "test_embedding_api_key"
|
||||||
mock_ai_config.return_value.llm_api_key = "test_api_key"
|
mock_ai_config.return_value.llm_api_key = "test_api_key"
|
||||||
mock_ai_config.return_value.llm_endpoint = "http://test-url"
|
mock_ai_config.return_value.llm_endpoint = "http://test-url"
|
||||||
|
|
||||||
@@ -79,7 +74,7 @@ def test_get_embedding_model_openai(mock_ai_config):
|
|||||||
model = get_embedding_model(mock_ai_config.return_value)
|
model = get_embedding_model(mock_ai_config.return_value)
|
||||||
MockOpenAIEmbedding.assert_called_once_with(
|
MockOpenAIEmbedding.assert_called_once_with(
|
||||||
model_name="text-embedding-3-small",
|
model_name="text-embedding-3-small",
|
||||||
api_key="test_api_key",
|
api_key="test_embedding_api_key",
|
||||||
api_base="http://test-url",
|
api_base="http://test-url",
|
||||||
timeout=120,
|
timeout=120,
|
||||||
http_client=ANY,
|
http_client=ANY,
|
||||||
@@ -88,6 +83,20 @@ def test_get_embedding_model_openai(mock_ai_config):
|
|||||||
assert model == MockOpenAIEmbedding.return_value
|
assert model == MockOpenAIEmbedding.return_value
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_embedding_model_openai_falls_back_to_llm_api_key(mock_ai_config):
|
||||||
|
mock_ai_config.return_value.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
||||||
|
mock_ai_config.return_value.llm_embedding_model = "text-embedding-3-small"
|
||||||
|
mock_ai_config.return_value.llm_api_key = "test_api_key"
|
||||||
|
mock_ai_config.return_value.llm_endpoint = "http://test-url"
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"llama_index.embeddings.openai_like.OpenAILikeEmbedding",
|
||||||
|
) as MockOpenAIEmbedding:
|
||||||
|
get_embedding_model(mock_ai_config.return_value)
|
||||||
|
|
||||||
|
assert MockOpenAIEmbedding.call_args.kwargs["api_key"] == "test_api_key"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("configured_key", [None, ""])
|
@pytest.mark.parametrize("configured_key", [None, ""])
|
||||||
def test_get_embedding_model_openai_without_api_key_sends_placeholder(
|
def test_get_embedding_model_openai_without_api_key_sends_placeholder(
|
||||||
mock_ai_config,
|
mock_ai_config,
|
||||||
@@ -96,6 +105,7 @@ def test_get_embedding_model_openai_without_api_key_sends_placeholder(
|
|||||||
"""Same required key handling as the LLM client, see #13831."""
|
"""Same required key handling as the LLM client, see #13831."""
|
||||||
mock_ai_config.return_value.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
mock_ai_config.return_value.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
||||||
mock_ai_config.return_value.llm_embedding_model = "text-embedding-3-small"
|
mock_ai_config.return_value.llm_embedding_model = "text-embedding-3-small"
|
||||||
|
mock_ai_config.return_value.llm_embedding_api_key = configured_key
|
||||||
mock_ai_config.return_value.llm_api_key = configured_key
|
mock_ai_config.return_value.llm_api_key = configured_key
|
||||||
mock_ai_config.return_value.llm_endpoint = "http://test-url"
|
mock_ai_config.return_value.llm_endpoint = "http://test-url"
|
||||||
|
|
||||||
@@ -290,61 +300,3 @@ def test_normalize_llm_index_text_collapses_ocr_leaders_without_joining_lines():
|
|||||||
|
|
||||||
def test_normalize_llm_index_text_collapses_non_breaking_spaces():
|
def test_normalize_llm_index_text_collapses_non_breaking_spaces():
|
||||||
assert _normalize_llm_index_text("A\u00a0........\u00a0B") == "A B"
|
assert _normalize_llm_index_text("A\u00a0........\u00a0B") == "A B"
|
||||||
|
|
||||||
|
|
||||||
class TestGuardedEmbeddingClients:
|
|
||||||
def test_ollama_embedding_clients_are_guarded(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The Ollama embedding backend
|
|
||||||
WHEN:
|
|
||||||
- The embedding model is built
|
|
||||||
THEN:
|
|
||||||
- The clients swapped onto it use guarded transports
|
|
||||||
"""
|
|
||||||
config = mock_ai_config.return_value
|
|
||||||
config.llm_embedding_backend = LLMEmbeddingBackend.OLLAMA
|
|
||||||
config.llm_embedding_model = "embeddinggemma"
|
|
||||||
config.llm_endpoint = "http://93.184.216.34:11434"
|
|
||||||
config.llm_allow_internal_endpoints = False
|
|
||||||
|
|
||||||
mocker.patch("llama_index.embeddings.ollama.OllamaEmbedding")
|
|
||||||
|
|
||||||
model = cast("OllamaEmbedding", get_embedding_model(config))
|
|
||||||
|
|
||||||
assert guard_of(model._client._client)._allow_internal is False
|
|
||||||
assert guard_of(model._async_client._client)._allow_internal is False
|
|
||||||
|
|
||||||
def test_openai_like_embedding_clients_are_guarded(
|
|
||||||
self,
|
|
||||||
mocker: MockerFixture,
|
|
||||||
mock_ai_config: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- The OpenAI-like embedding backend with an endpoint
|
|
||||||
WHEN:
|
|
||||||
- The embedding model is built
|
|
||||||
THEN:
|
|
||||||
- Its http clients use guarded transports
|
|
||||||
"""
|
|
||||||
config = mock_ai_config.return_value
|
|
||||||
config.llm_embedding_backend = LLMEmbeddingBackend.OPENAI_LIKE
|
|
||||||
config.llm_embedding_model = "text-embedding-3-small"
|
|
||||||
config.llm_api_key = "key"
|
|
||||||
config.llm_endpoint = "http://93.184.216.34:8080"
|
|
||||||
config.llm_allow_internal_endpoints = False
|
|
||||||
|
|
||||||
embedding_class = mocker.patch(
|
|
||||||
"llama_index.embeddings.openai_like.OpenAILikeEmbedding",
|
|
||||||
)
|
|
||||||
|
|
||||||
get_embedding_model(config)
|
|
||||||
|
|
||||||
kwargs = embedding_class.call_args.kwargs
|
|
||||||
assert guard_of(kwargs["http_client"])._allow_internal is False
|
|
||||||
assert guard_of(kwargs["async_http_client"])._allow_internal is False
|
|
||||||
|
|||||||
+21
-43
@@ -45,11 +45,8 @@ from documents.models import Correspondent
|
|||||||
from documents.models import PaperlessTask
|
from documents.models import PaperlessTask
|
||||||
from documents.parsers import is_mime_type_supported
|
from documents.parsers import is_mime_type_supported
|
||||||
from documents.tasks import consume_file
|
from documents.tasks import consume_file
|
||||||
from paperless.network import HostResolutionError
|
from paperless.network import is_public_ip
|
||||||
from paperless.network import IPAddress
|
from paperless.network import resolve_hostname_ips
|
||||||
from paperless.network import OutboundRequestBlockedError
|
|
||||||
from paperless.network import blocked_message
|
|
||||||
from paperless.network import resolve_public_addresses
|
|
||||||
from paperless_mail.models import MailAccount
|
from paperless_mail.models import MailAccount
|
||||||
from paperless_mail.models import MailRule
|
from paperless_mail.models import MailRule
|
||||||
from paperless_mail.models import ProcessedMail
|
from paperless_mail.models import ProcessedMail
|
||||||
@@ -448,34 +445,18 @@ class PinnedIMAP4(imaplib.IMAP4):
|
|||||||
|
|
||||||
Without pinned addresses, and with the ssl_context of the matching imaplib
|
Without pinned addresses, and with the ssl_context of the matching imaplib
|
||||||
class, this behaves exactly like imaplib.IMAP4 / imaplib.IMAP4_SSL.
|
class, this behaves exactly like imaplib.IMAP4 / imaplib.IMAP4_SSL.
|
||||||
|
|
||||||
``pinned_ips`` of ``None`` means no pinning was requested and the stock
|
|
||||||
imaplib connection path is used. An empty tuple means pinning was requested
|
|
||||||
and yielded nothing, and the connection fails without opening a socket
|
|
||||||
rather than falling back to a hostname lookup.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(self, host, port, pinned_ips, ssl_context=None, timeout=None) -> None:
|
||||||
self,
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
pinned_ips: tuple[IPAddress, ...] | None,
|
|
||||||
ssl_context: ssl.SSLContext | None = None,
|
|
||||||
timeout: float | None = None,
|
|
||||||
) -> None:
|
|
||||||
self._pinned_ips = pinned_ips
|
self._pinned_ips = pinned_ips
|
||||||
self.ssl_context = ssl_context
|
self.ssl_context = ssl_context
|
||||||
super().__init__(host, port, timeout=timeout)
|
super().__init__(host, port, timeout=timeout)
|
||||||
|
|
||||||
def _connect_pinned(
|
def _connect_pinned(self, timeout):
|
||||||
self,
|
|
||||||
pinned_ips: tuple[IPAddress, ...],
|
|
||||||
timeout: float | None,
|
|
||||||
) -> socket.socket:
|
|
||||||
last_error: OSError | None = None
|
last_error: OSError | None = None
|
||||||
for ip in pinned_ips:
|
for ip_str in self._pinned_ips:
|
||||||
try:
|
try:
|
||||||
address = (str(ip), self.port)
|
address = (ip_str, self.port)
|
||||||
if timeout is not None:
|
if timeout is not None:
|
||||||
return socket.create_connection(address, timeout)
|
return socket.create_connection(address, timeout)
|
||||||
return socket.create_connection(address)
|
return socket.create_connection(address)
|
||||||
@@ -483,9 +464,9 @@ class PinnedIMAP4(imaplib.IMAP4):
|
|||||||
last_error = e
|
last_error = e
|
||||||
raise last_error or OSError(f"Could not connect to {self.host}")
|
raise last_error or OSError(f"Could not connect to {self.host}")
|
||||||
|
|
||||||
def _create_socket(self, timeout: float | None) -> socket.socket:
|
def _create_socket(self, timeout):
|
||||||
if self._pinned_ips is not None:
|
if self._pinned_ips:
|
||||||
sock = self._connect_pinned(self._pinned_ips, timeout)
|
sock = self._connect_pinned(timeout)
|
||||||
else:
|
else:
|
||||||
sock = super()._create_socket(timeout)
|
sock = super()._create_socket(timeout)
|
||||||
if self.ssl_context is None:
|
if self.ssl_context is None:
|
||||||
@@ -496,12 +477,7 @@ class PinnedIMAP4(imaplib.IMAP4):
|
|||||||
class PinnedClientMixin:
|
class PinnedClientMixin:
|
||||||
"""Builds the imaplib client against the pre-resolved addresses, if any."""
|
"""Builds the imaplib client against the pre-resolved addresses, if any."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(self, *args, pinned_ips: list[str] | None, **kwargs) -> None:
|
||||||
self,
|
|
||||||
*args,
|
|
||||||
pinned_ips: tuple[IPAddress, ...] | None,
|
|
||||||
**kwargs,
|
|
||||||
) -> None:
|
|
||||||
self._pinned_ips = pinned_ips
|
self._pinned_ips = pinned_ips
|
||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
@@ -539,20 +515,22 @@ class PinnedMailBoxStartTls(PinnedClientMixin, MailBoxStartTls):
|
|||||||
return client
|
return client
|
||||||
|
|
||||||
|
|
||||||
def get_mailbox(
|
def get_mailbox(server, port, security) -> MailBox:
|
||||||
server: str,
|
|
||||||
port: int | None,
|
|
||||||
security: int,
|
|
||||||
) -> MailBox:
|
|
||||||
"""
|
"""
|
||||||
Returns the correct MailBox instance for the given configuration.
|
Returns the correct MailBox instance for the given configuration.
|
||||||
"""
|
"""
|
||||||
pinned_ips: tuple[IPAddress, ...] | None = None
|
pinned_ips: list[str] | None = None
|
||||||
if not settings.EMAIL_ALLOW_INTERNAL_HOSTS:
|
if not settings.EMAIL_ALLOW_INTERNAL_HOSTS:
|
||||||
try:
|
try:
|
||||||
pinned_ips = resolve_public_addresses(server, port)
|
pinned_ips = resolve_hostname_ips(server)
|
||||||
except (OutboundRequestBlockedError, HostResolutionError) as e:
|
except ValueError as e:
|
||||||
raise MailError(blocked_message(e)) from e
|
raise MailError(str(e)) from e
|
||||||
|
|
||||||
|
for ip_str in pinned_ips:
|
||||||
|
if not is_public_ip(ip_str):
|
||||||
|
raise MailError(
|
||||||
|
f"Connection blocked: {server} resolves to a non-public address",
|
||||||
|
)
|
||||||
|
|
||||||
ssl_context = ssl.create_default_context()
|
ssl_context = ssl.create_default_context()
|
||||||
if settings.EMAIL_CERTIFICATE_FILE is not None: # pragma: no cover
|
if settings.EMAIL_CERTIFICATE_FILE is not None: # pragma: no cover
|
||||||
|
|||||||
@@ -2,14 +2,14 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import factory
|
import factory
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from factory.django import DjangoModelFactory
|
|
||||||
|
|
||||||
from paperless_mail.models import MailAccount
|
from paperless_mail.models import MailAccount
|
||||||
from paperless_mail.models import MailRule
|
from paperless_mail.models import MailRule
|
||||||
from paperless_mail.models import ProcessedMail
|
from paperless_mail.models import ProcessedMail
|
||||||
|
from paperless_testing.typed_factory import TypedModelFactory
|
||||||
|
|
||||||
|
|
||||||
class MailAccountFactory(DjangoModelFactory[MailAccount]):
|
class MailAccountFactory(TypedModelFactory[MailAccount]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = MailAccount
|
model = MailAccount
|
||||||
|
|
||||||
@@ -24,7 +24,7 @@ class MailAccountFactory(DjangoModelFactory[MailAccount]):
|
|||||||
is_token = False
|
is_token = False
|
||||||
|
|
||||||
|
|
||||||
class MailRuleFactory(DjangoModelFactory[MailRule]):
|
class MailRuleFactory(TypedModelFactory[MailRule]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = MailRule
|
model = MailRule
|
||||||
|
|
||||||
@@ -44,7 +44,7 @@ class MailRuleFactory(DjangoModelFactory[MailRule]):
|
|||||||
stop_processing = False
|
stop_processing = False
|
||||||
|
|
||||||
|
|
||||||
class ProcessedMailFactory(DjangoModelFactory[ProcessedMail]):
|
class ProcessedMailFactory(TypedModelFactory[ProcessedMail]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = ProcessedMail
|
model = ProcessedMail
|
||||||
|
|
||||||
|
|||||||
@@ -1,12 +1,9 @@
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import ipaddress
|
|
||||||
import socket
|
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections import namedtuple
|
from collections import namedtuple
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from django.contrib.auth.models import Permission
|
from django.contrib.auth.models import Permission
|
||||||
@@ -28,7 +25,6 @@ from documents.models import MatchingModel
|
|||||||
from paperless_mail import tasks
|
from paperless_mail import tasks
|
||||||
from paperless_mail.mail import MailAccountHandler
|
from paperless_mail.mail import MailAccountHandler
|
||||||
from paperless_mail.mail import MailError
|
from paperless_mail.mail import MailError
|
||||||
from paperless_mail.mail import PinnedIMAP4
|
|
||||||
from paperless_mail.mail import TagMailAction
|
from paperless_mail.mail import TagMailAction
|
||||||
from paperless_mail.mail import apply_mail_action
|
from paperless_mail.mail import apply_mail_action
|
||||||
from paperless_mail.mail import error_callback
|
from paperless_mail.mail import error_callback
|
||||||
@@ -1569,7 +1565,12 @@ class TestMail(
|
|||||||
("electronic", None, "invoices@mycompany.com", None, 1),
|
("electronic", None, "invoices@mycompany.com", None, 1),
|
||||||
(None, "amazon", "me@myselfandi.com", None, 1),
|
(None, "amazon", "me@myselfandi.com", None, 1),
|
||||||
]:
|
]:
|
||||||
with self.subTest(f_body=f_body, f_from=f_from, f_subject=f_subject):
|
with self.subTest(
|
||||||
|
f_body=f_body,
|
||||||
|
f_from=f_from,
|
||||||
|
f_to=f_to,
|
||||||
|
f_subject=f_subject,
|
||||||
|
):
|
||||||
MailRule.objects.all().delete()
|
MailRule.objects.all().delete()
|
||||||
_ = MailRule.objects.create(
|
_ = MailRule.objects.create(
|
||||||
name="testrule3",
|
name="testrule3",
|
||||||
@@ -1810,7 +1811,7 @@ class TestPostConsumeAction(TestCase):
|
|||||||
|
|
||||||
with (
|
with (
|
||||||
self.assertRaises(errors.ImapToolsError),
|
self.assertRaises(errors.ImapToolsError),
|
||||||
self.assertLogs("paperless.mail", level="ERROR") as cm,
|
self.assertLogs("paperless_mail", level="ERROR") as cm,
|
||||||
):
|
):
|
||||||
apply_mail_action(
|
apply_mail_action(
|
||||||
result=[],
|
result=[],
|
||||||
@@ -1819,9 +1820,10 @@ class TestPostConsumeAction(TestCase):
|
|||||||
message_subject=self.message_subject,
|
message_subject=self.message_subject,
|
||||||
message_date=self.message_date,
|
message_date=self.message_date,
|
||||||
)
|
)
|
||||||
error_str = cm.output[0]
|
|
||||||
expected_str = "Error while processing mail action during post_consume"
|
error_str = cm.output[0]
|
||||||
self.assertIn(expected_str, error_str)
|
expected_str = "Error while processing mail action during post_consume"
|
||||||
|
self.assertIn(expected_str, error_str)
|
||||||
|
|
||||||
processed_mail = ProcessedMail.objects.get(uid=self.message_uid)
|
processed_mail = ProcessedMail.objects.get(uid=self.message_uid)
|
||||||
self.assertEqual(processed_mail.status, "FAILED")
|
self.assertEqual(processed_mail.status, "FAILED")
|
||||||
@@ -2049,13 +2051,10 @@ class TestMailAccountTestView(APITestCase):
|
|||||||
self.assertEqual(response.content.decode(), "Unable to connect to server")
|
self.assertEqual(response.content.decode(), "Unable to connect to server")
|
||||||
|
|
||||||
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
||||||
@mock.patch(
|
@mock.patch("paperless_mail.mail.resolve_hostname_ips", return_value=["127.0.0.1"])
|
||||||
"paperless.network._getaddrinfo",
|
|
||||||
return_value=[(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 993))],
|
|
||||||
)
|
|
||||||
def test_mail_account_test_view_blocks_internal_host_when_disabled(
|
def test_mail_account_test_view_blocks_internal_host_when_disabled(
|
||||||
self,
|
self,
|
||||||
_mock_getaddrinfo: MagicMock,
|
_mock_resolve_hostname_ips,
|
||||||
) -> None:
|
) -> None:
|
||||||
data = {
|
data = {
|
||||||
"imap_server": "internal.example",
|
"imap_server": "internal.example",
|
||||||
@@ -2212,10 +2211,10 @@ class TestGetMailboxHostPinning(TestCase):
|
|||||||
|
|
||||||
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
||||||
@mock.patch(
|
@mock.patch(
|
||||||
"paperless_mail.mail.resolve_public_addresses",
|
"paperless_mail.mail.resolve_hostname_ips",
|
||||||
return_value=(ipaddress.ip_address("93.184.216.34"),),
|
return_value=["93.184.216.34"],
|
||||||
)
|
)
|
||||||
def test_connects_to_validated_ip(self, _mock_resolve: MagicMock) -> None:
|
def test_connects_to_validated_ip(self, _mock_resolve) -> None:
|
||||||
with mock.patch(
|
with mock.patch(
|
||||||
"paperless_mail.mail.socket.create_connection",
|
"paperless_mail.mail.socket.create_connection",
|
||||||
side_effect=OSError("no connection in tests"),
|
side_effect=OSError("no connection in tests"),
|
||||||
@@ -2232,13 +2231,10 @@ class TestGetMailboxHostPinning(TestCase):
|
|||||||
|
|
||||||
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
||||||
@mock.patch(
|
@mock.patch(
|
||||||
"paperless_mail.mail.resolve_public_addresses",
|
"paperless_mail.mail.resolve_hostname_ips",
|
||||||
return_value=(ipaddress.ip_address("93.184.216.34"),),
|
return_value=["93.184.216.34"],
|
||||||
)
|
)
|
||||||
def test_ssl_pins_ip_but_keeps_hostname_for_sni(
|
def test_ssl_pins_ip_but_keeps_hostname_for_sni(self, _mock_resolve) -> None:
|
||||||
self,
|
|
||||||
_mock_resolve: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
ssl_context = mock.MagicMock()
|
ssl_context = mock.MagicMock()
|
||||||
ssl_context.wrap_socket.return_value.makefile.side_effect = OSError(
|
ssl_context.wrap_socket.return_value.makefile.side_effect = OSError(
|
||||||
"no connection in tests",
|
"no connection in tests",
|
||||||
@@ -2269,51 +2265,13 @@ class TestGetMailboxHostPinning(TestCase):
|
|||||||
|
|
||||||
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
|
||||||
@mock.patch(
|
@mock.patch(
|
||||||
"paperless.network._getaddrinfo",
|
"paperless_mail.mail.resolve_hostname_ips",
|
||||||
return_value=[
|
return_value=["93.184.216.34", "127.0.0.1"],
|
||||||
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 993)),
|
|
||||||
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 993)),
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
def test_blocks_when_any_resolved_address_is_internal(
|
def test_blocks_when_any_resolved_address_is_internal(self, _mock_resolve) -> None:
|
||||||
self,
|
with self.assertRaises(MailError):
|
||||||
_mock_resolve: MagicMock,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A mail host resolving to one public and one loopback address
|
|
||||||
- EMAIL_ALLOW_INTERNAL_HOSTS is False
|
|
||||||
WHEN:
|
|
||||||
- A mailbox is requested
|
|
||||||
THEN:
|
|
||||||
- The whole host is blocked with the existing message
|
|
||||||
"""
|
|
||||||
with self.assertRaisesMessage(
|
|
||||||
MailError,
|
|
||||||
"Connection blocked: mail.example.com resolves to a non-public address",
|
|
||||||
):
|
|
||||||
get_mailbox("mail.example.com", 993, MailAccount.ImapSecurity.SSL)
|
get_mailbox("mail.example.com", 993, MailAccount.ImapSecurity.SSL)
|
||||||
|
|
||||||
def test_empty_pin_list_never_falls_back_to_hostname_lookup(self) -> None:
|
|
||||||
"""
|
|
||||||
GIVEN:
|
|
||||||
- A pinned IMAP client given an empty tuple of addresses
|
|
||||||
WHEN:
|
|
||||||
- It connects
|
|
||||||
THEN:
|
|
||||||
- It fails without opening any socket, rather than resolving the
|
|
||||||
hostname itself
|
|
||||||
"""
|
|
||||||
with (
|
|
||||||
mock.patch("paperless_mail.mail.socket.create_connection") as pinned,
|
|
||||||
mock.patch("imaplib.IMAP4._create_socket") as unpinned,
|
|
||||||
self.assertRaises(OSError),
|
|
||||||
):
|
|
||||||
PinnedIMAP4("mail.example.com", 143, ())
|
|
||||||
|
|
||||||
pinned.assert_not_called()
|
|
||||||
unpinned.assert_not_called()
|
|
||||||
|
|
||||||
|
|
||||||
class TestMailAccountProcess(APITestCase):
|
class TestMailAccountProcess(APITestCase):
|
||||||
def setUp(self) -> None:
|
def setUp(self) -> None:
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ class PaperlessDirs:
|
|||||||
logging_dir: Path
|
logging_dir: Path
|
||||||
model_file: Path
|
model_file: Path
|
||||||
media_lock: Path
|
media_lock: Path
|
||||||
|
share_link_bundle_dir: Path
|
||||||
|
|
||||||
|
|
||||||
class DirSettings(TypedDict):
|
class DirSettings(TypedDict):
|
||||||
@@ -54,6 +55,7 @@ class DirSettings(TypedDict):
|
|||||||
STATIC_ROOT: Path
|
STATIC_ROOT: Path
|
||||||
MODEL_FILE: Path
|
MODEL_FILE: Path
|
||||||
MEDIA_LOCK: Path
|
MEDIA_LOCK: Path
|
||||||
|
SHARE_LINK_BUNDLE_DIR: Path
|
||||||
|
|
||||||
|
|
||||||
def build_paperless_dirs(root: Path) -> PaperlessDirs:
|
def build_paperless_dirs(root: Path) -> PaperlessDirs:
|
||||||
@@ -75,6 +77,7 @@ def build_paperless_dirs(root: Path) -> PaperlessDirs:
|
|||||||
logging_dir=data_dir / "log",
|
logging_dir=data_dir / "log",
|
||||||
model_file=data_dir / "classification_model.pickle",
|
model_file=data_dir / "classification_model.pickle",
|
||||||
media_lock=media_dir / "media.lock",
|
media_lock=media_dir / "media.lock",
|
||||||
|
share_link_bundle_dir=documents_dir / "share_link_bundles",
|
||||||
)
|
)
|
||||||
|
|
||||||
for directory in (
|
for directory in (
|
||||||
@@ -109,6 +112,7 @@ def dirs_settings(dirs: PaperlessDirs) -> DirSettings:
|
|||||||
STATIC_ROOT=dirs.static_dir,
|
STATIC_ROOT=dirs.static_dir,
|
||||||
MODEL_FILE=dirs.model_file,
|
MODEL_FILE=dirs.model_file,
|
||||||
MEDIA_LOCK=dirs.media_lock,
|
MEDIA_LOCK=dirs.media_lock,
|
||||||
|
SHARE_LINK_BUNDLE_DIR=dirs.share_link_bundle_dir,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,8 +5,7 @@ Factory-boy factories for documents app models.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import factory
|
import factory
|
||||||
from django.contrib.auth import get_user_model
|
from django.contrib.auth.models import User
|
||||||
from factory.django import DjangoModelFactory
|
|
||||||
|
|
||||||
from documents.models import Correspondent
|
from documents.models import Correspondent
|
||||||
from documents.models import Document
|
from documents.models import Document
|
||||||
@@ -15,11 +14,10 @@ from documents.models import MatchingModel
|
|||||||
from documents.models import PaperlessTask
|
from documents.models import PaperlessTask
|
||||||
from documents.models import StoragePath
|
from documents.models import StoragePath
|
||||||
from documents.models import Tag
|
from documents.models import Tag
|
||||||
|
from paperless_testing.typed_factory import TypedModelFactory
|
||||||
UserModelT = get_user_model()
|
|
||||||
|
|
||||||
|
|
||||||
class CorrespondentFactory(DjangoModelFactory[Correspondent]):
|
class CorrespondentFactory(TypedModelFactory[Correspondent]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = Correspondent
|
model = Correspondent
|
||||||
|
|
||||||
@@ -28,7 +26,7 @@ class CorrespondentFactory(DjangoModelFactory[Correspondent]):
|
|||||||
matching_algorithm = MatchingModel.MATCH_NONE
|
matching_algorithm = MatchingModel.MATCH_NONE
|
||||||
|
|
||||||
|
|
||||||
class DocumentTypeFactory(DjangoModelFactory[DocumentType]):
|
class DocumentTypeFactory(TypedModelFactory[DocumentType]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = DocumentType
|
model = DocumentType
|
||||||
|
|
||||||
@@ -37,7 +35,7 @@ class DocumentTypeFactory(DjangoModelFactory[DocumentType]):
|
|||||||
matching_algorithm = MatchingModel.MATCH_NONE
|
matching_algorithm = MatchingModel.MATCH_NONE
|
||||||
|
|
||||||
|
|
||||||
class TagFactory(DjangoModelFactory[Tag]):
|
class TagFactory(TypedModelFactory[Tag]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = Tag
|
model = Tag
|
||||||
|
|
||||||
@@ -47,7 +45,7 @@ class TagFactory(DjangoModelFactory[Tag]):
|
|||||||
is_inbox_tag = False
|
is_inbox_tag = False
|
||||||
|
|
||||||
|
|
||||||
class StoragePathFactory(DjangoModelFactory[StoragePath]):
|
class StoragePathFactory(TypedModelFactory[StoragePath]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = StoragePath
|
model = StoragePath
|
||||||
|
|
||||||
@@ -59,7 +57,7 @@ class StoragePathFactory(DjangoModelFactory[StoragePath]):
|
|||||||
matching_algorithm = MatchingModel.MATCH_NONE
|
matching_algorithm = MatchingModel.MATCH_NONE
|
||||||
|
|
||||||
|
|
||||||
class DocumentFactory(DjangoModelFactory[Document]):
|
class DocumentFactory(TypedModelFactory[Document]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = Document
|
model = Document
|
||||||
|
|
||||||
@@ -71,9 +69,9 @@ class DocumentFactory(DjangoModelFactory[Document]):
|
|||||||
storage_path = None
|
storage_path = None
|
||||||
|
|
||||||
|
|
||||||
class UserFactory(DjangoModelFactory[UserModelT]):
|
class UserFactory(TypedModelFactory[User]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = UserModelT
|
model = User
|
||||||
|
|
||||||
username = factory.Sequence(lambda n: f"user{n}")
|
username = factory.Sequence(lambda n: f"user{n}")
|
||||||
is_staff = False
|
is_staff = False
|
||||||
@@ -88,7 +86,7 @@ class UserFactory(DjangoModelFactory[UserModelT]):
|
|||||||
staff = factory.Trait(is_staff=True)
|
staff = factory.Trait(is_staff=True)
|
||||||
|
|
||||||
|
|
||||||
class PaperlessTaskFactory(DjangoModelFactory[PaperlessTask]):
|
class PaperlessTaskFactory(TypedModelFactory[PaperlessTask]):
|
||||||
class Meta:
|
class Meta:
|
||||||
model = PaperlessTask
|
model = PaperlessTask
|
||||||
|
|
||||||
|
|||||||
@@ -1,218 +0,0 @@
|
|||||||
"""
|
|
||||||
Real-socket helpers for tests of the outbound connection guard in
|
|
||||||
paperless.network: a local HTTP server, a per-hostname resolver fake and
|
|
||||||
spies recording which addresses were actually dialled.
|
|
||||||
|
|
||||||
The fixtures wrapping these live in the root conftest.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import http.server
|
|
||||||
import socket
|
|
||||||
import threading
|
|
||||||
from contextlib import contextmanager
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from dataclasses import field
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
from typing import Any
|
|
||||||
from typing import cast
|
|
||||||
|
|
||||||
import anyio
|
|
||||||
import httpcore
|
|
||||||
|
|
||||||
from paperless.network import GuardedAsyncHTTPTransport
|
|
||||||
from paperless.network import GuardedHTTPTransport
|
|
||||||
from paperless.network import _GuardedAsyncBackend
|
|
||||||
from paperless.network import _GuardedSyncBackend
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from collections.abc import Iterator
|
|
||||||
from unittest.mock import MagicMock
|
|
||||||
from unittest.mock import _Call
|
|
||||||
|
|
||||||
import httpx
|
|
||||||
from pytest_mock import MockerFixture
|
|
||||||
|
|
||||||
_REAL_GETADDRINFO = socket.getaddrinfo
|
|
||||||
_REAL_AGETADDRINFO = anyio.getaddrinfo
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class ReceivedRequest:
|
|
||||||
method: str
|
|
||||||
path: str
|
|
||||||
headers: dict[str, str]
|
|
||||||
body: bytes
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class LocalHTTPServer:
|
|
||||||
"""State of a threaded HTTP server bound to 127.0.0.1 on an ephemeral port."""
|
|
||||||
|
|
||||||
port: int
|
|
||||||
requests: list[ReceivedRequest] = field(default_factory=list)
|
|
||||||
connections: int = 0
|
|
||||||
redirect_to: str | None = None
|
|
||||||
|
|
||||||
|
|
||||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
|
||||||
# HTTP/1.1 keeps connections open, so tests can observe connection reuse.
|
|
||||||
# Every response sets Content-Length, which keep-alive requires.
|
|
||||||
protocol_version = "HTTP/1.1"
|
|
||||||
|
|
||||||
def _handle(self) -> None:
|
|
||||||
length = int(self.headers.get("Content-Length") or 0)
|
|
||||||
body = self.rfile.read(length) if length else b""
|
|
||||||
# BaseHTTPRequestHandler types server as the base socketserver.BaseServer;
|
|
||||||
# narrowing the attribute's declared type is a variance error, so the
|
|
||||||
# subclass is recovered here instead of on the class body.
|
|
||||||
server = cast("_RecordingHTTPServer", self.server)
|
|
||||||
state = server.state
|
|
||||||
state.requests.append(
|
|
||||||
ReceivedRequest(
|
|
||||||
method=self.command,
|
|
||||||
path=self.path,
|
|
||||||
headers={key.lower(): value for key, value in self.headers.items()},
|
|
||||||
body=body,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
if state.redirect_to is not None:
|
|
||||||
self.send_response(302)
|
|
||||||
self.send_header("Location", state.redirect_to)
|
|
||||||
self.send_header("Content-Length", "0")
|
|
||||||
self.end_headers()
|
|
||||||
return
|
|
||||||
self.send_response(200)
|
|
||||||
self.send_header("Content-Length", "2")
|
|
||||||
self.end_headers()
|
|
||||||
self.wfile.write(b"ok")
|
|
||||||
|
|
||||||
do_GET = _handle
|
|
||||||
do_POST = _handle
|
|
||||||
|
|
||||||
def log_message(self, format: str, *args: Any) -> None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class _RecordingHTTPServer(http.server.ThreadingHTTPServer):
|
|
||||||
daemon_threads = True
|
|
||||||
|
|
||||||
def __init__(self) -> None:
|
|
||||||
super().__init__(("127.0.0.1", 0), _Handler)
|
|
||||||
self.state = LocalHTTPServer(port=self.socket.getsockname()[1])
|
|
||||||
|
|
||||||
def verify_request(self, request: Any, client_address: Any) -> bool:
|
|
||||||
self.state.connections += 1
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def running_http_server() -> Iterator[LocalHTTPServer]:
|
|
||||||
"""Serve on 127.0.0.1 in a background thread until the block exits."""
|
|
||||||
server = _RecordingHTTPServer()
|
|
||||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
|
||||||
thread.start()
|
|
||||||
try:
|
|
||||||
yield server.state
|
|
||||||
finally:
|
|
||||||
server.shutdown()
|
|
||||||
server.server_close()
|
|
||||||
thread.join(timeout=5)
|
|
||||||
|
|
||||||
|
|
||||||
def _addrinfo(address: str, port: int | None) -> tuple[Any, ...]:
|
|
||||||
if ":" in address:
|
|
||||||
return (socket.AF_INET6, socket.SOCK_STREAM, 6, "", (address, port or 0, 0, 0))
|
|
||||||
return (socket.AF_INET, socket.SOCK_STREAM, 6, "", (address, port or 0))
|
|
||||||
|
|
||||||
|
|
||||||
class FakeDNS:
|
|
||||||
"""
|
|
||||||
Answers the guard's resolver hooks for registered names and delegates
|
|
||||||
every other name to the real resolver. The stock httpcore backends keep
|
|
||||||
using the unpatched socket.getaddrinfo.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self) -> None:
|
|
||||||
self._answers: dict[str, list[str]] = {}
|
|
||||||
self.lookups: list[str] = []
|
|
||||||
|
|
||||||
def add(self, hostname: str, *addresses: str) -> None:
|
|
||||||
self._answers[hostname] = list(addresses)
|
|
||||||
|
|
||||||
def getaddrinfo(
|
|
||||||
self,
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
*args: Any,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> list[tuple[Any, ...]]:
|
|
||||||
self.lookups.append(host)
|
|
||||||
if host in self._answers:
|
|
||||||
return [_addrinfo(address, port) for address in self._answers[host]]
|
|
||||||
return list(_REAL_GETADDRINFO(host, port, *args, **kwargs))
|
|
||||||
|
|
||||||
async def agetaddrinfo(
|
|
||||||
self,
|
|
||||||
host: str,
|
|
||||||
port: int | None,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> list[tuple[Any, ...]]:
|
|
||||||
self.lookups.append(host)
|
|
||||||
if host in self._answers:
|
|
||||||
return [_addrinfo(address, port) for address in self._answers[host]]
|
|
||||||
return list(await _REAL_AGETADDRINFO(host, port, **kwargs))
|
|
||||||
|
|
||||||
|
|
||||||
def install_fake_dns(mocker: MockerFixture) -> FakeDNS:
|
|
||||||
"""Patch the guard's resolver hooks with a FakeDNS for the current test."""
|
|
||||||
dns = FakeDNS()
|
|
||||||
mocker.patch("paperless.network._getaddrinfo", new=dns.getaddrinfo)
|
|
||||||
mocker.patch("paperless.network._agetaddrinfo", new=dns.agetaddrinfo)
|
|
||||||
return dns
|
|
||||||
|
|
||||||
|
|
||||||
def _dialled_host(call: _Call) -> str:
|
|
||||||
# The spy sits on the class, so args[0] is the backend instance.
|
|
||||||
if "host" in call.kwargs:
|
|
||||||
return str(call.kwargs["host"])
|
|
||||||
return str(call.args[1])
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class DialRecorder:
|
|
||||||
sync_spy: MagicMock
|
|
||||||
async_spy: MagicMock
|
|
||||||
|
|
||||||
def hosts(self) -> list[str]:
|
|
||||||
calls = [*self.sync_spy.call_args_list, *self.async_spy.call_args_list]
|
|
||||||
return [_dialled_host(call) for call in calls]
|
|
||||||
|
|
||||||
|
|
||||||
def install_dial_recorder(mocker: MockerFixture) -> DialRecorder:
|
|
||||||
"""Spy on the stock backends' connect_tcp for the current test."""
|
|
||||||
return DialRecorder(
|
|
||||||
sync_spy=mocker.spy(httpcore.SyncBackend, "connect_tcp"),
|
|
||||||
async_spy=mocker.spy(httpcore.AnyIOBackend, "connect_tcp"),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def allow_all_addresses(mocker: MockerFixture) -> None:
|
|
||||||
"""Patch the guard's public-address check to accept every address.
|
|
||||||
|
|
||||||
Loopback and other private addresses pass just like a public one, for
|
|
||||||
tests that exercise something other than the address policy itself.
|
|
||||||
"""
|
|
||||||
mocker.patch("paperless.network.is_public_ip", return_value=True)
|
|
||||||
|
|
||||||
|
|
||||||
def guard_of(
|
|
||||||
client: httpx.Client | httpx.AsyncClient,
|
|
||||||
) -> _GuardedSyncBackend | _GuardedAsyncBackend:
|
|
||||||
"""Return the guard installed on a client's transport."""
|
|
||||||
transport = client._transport
|
|
||||||
assert isinstance(transport, GuardedHTTPTransport | GuardedAsyncHTTPTransport)
|
|
||||||
backend = transport._pool._network_backend
|
|
||||||
assert isinstance(backend, _GuardedSyncBackend | _GuardedAsyncBackend)
|
|
||||||
return backend
|
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
"""
|
||||||
|
A DjangoModelFactory base whose calls are typed as the model they build.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
from typing import Any
|
||||||
|
from typing import TypeVar
|
||||||
|
|
||||||
|
from factory.django import DjangoModelFactory
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class TypedModelFactory(DjangoModelFactory[T]):
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
# factory-boy leaves Factory() unannotated, so mypy takes it to build a
|
||||||
|
# factory instance. At runtime it builds the model.
|
||||||
|
def __new__(cls, *args: Any, **kwargs: Any) -> T: ... # type: ignore[misc]
|
||||||
Reference in New Issue
Block a user