Compare commits

...
Author SHA1 Message Date
stumpylogandClaude Opus 5 296ddff37e Collapse near-duplicate outbound guard tests into parametrized cases
Several tests in test_network.py and test_network_integration.py had
grown into near-copies of each other, differing only in the URL or
host they fed the guard while asserting the same thing. That made the
suite noisier to read and to extend than the behaviour it covers
warrants.

Merge the invalid-host rejection tests into one parametrized test over
all seven unusual host forms, and likewise for their allow_internal
counterparts that assert nothing is rejected. Merge the
resolve_public_addresses tests that only check the host reaches the
resolver unchanged, and the ones that check a private answer blocks
regardless of how the host was spelled, keeping the stricter resolver
call assertion on both merged cases. Add localhost as a fourth case to
the numeric-host-forms test in the transport integration tests, so the
name and the non-canonical numeric spellings of loopback are covered
by one test.

No behaviour or assertion strength changes; this only reduces
duplicated test bodies while keeping every original case addressable
by its own parametrize id.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 13:49:35 -07:00
stumpylogandClaude Opus 5 de8edc31c5 Fold the outbound guard's public-address patch into a named fixture
Six tests in test_network_integration.py and one in test_workflows.py each
patched paperless.network.is_public_ip inline to make every resolved
address count as public, so the guard's own address policy would not get
in the way of whatever the test was actually checking. The line was
repeated verbatim at every call site and understated what it did: read
literally it makes every address public, while each docstring around it
already explained the intent as loopback being treated as public.

Add an every_address_is_public fixture to conftest.py, following the
existing pattern for the other outbound test fixtures: a thin wrapper
that imports a new allow_all_addresses helper from
paperless_testing.outbound and calls it with mocker. Apply it at each
call site with pytest.mark.usefixtures, dropping the mocker parameter
from tests that no longer need it directly. The test that patches
is_public_ip with a selective side effect to allow only one address is
left untouched, since that selectivity is load-bearing for what it
verifies.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:56:55 -07:00
stumpylogandClaude Opus 5 46c150eb67 Deduplicate the outbound connect guard's sync and async paths
The two connect_tcp implementations each repeated the per-attempt budget
arithmetic verbatim, and only the sync one delegated resolution and error
mapping to a helper, so a drift between the copies could silently give one
stack a different timeout policy from the other. In the IMAP client,
_connect_pinned re-asserted a fact its only caller had already established,
which reads as a runtime invariant check on a security-relevant path when it
is only a narrowing aid.

Move the budget arithmetic into one helper called from both loops, add an
async twin of the resolve helper so the two loop bodies differ only by await,
and pass the narrowed address tuple into _connect_pinned instead of asserting
it. The PinnedIMAP4 docstring now spells out that no pinning and pinning that
yielded nothing are different things, and the monotonic clock seam says why it
exists.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:06 -07:00
stumpylogandClaude Opus 5 58777076c5 Cover webhooks to internal hosts when they are allowed
Every webhook test that reaches a real socket ran with internal requests
disallowed, and the rest intercept above the transport. Wiring the
transport to always disallow internal addresses would therefore have
passed the suite while breaking webhooks to internal hosts on every
default install.

A new test sends a webhook to localhost with internal requests allowed
and checks that the payload arrives and that the guard resolved nothing.
The Host header test's docstring is also corrected: the header now comes
from the URL, not from a resolved hostname.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:06 -07:00
stumpylogandClaude Opus 5 d414297c05 Fix: Validate IP literals from the resolver's answer only
Outbound host checks parsed IP literals themselves and skipped the
resolver for them. That second parser is what let a host such as
8.8.8.8%2eexample pass as the public address 8.8.8.8, and even with
zone ids limited to IPv6 it remains one more place where the checked
host can be read differently from the connected one.

The separate literal parsing is removed. Every host now goes to
getaddrinfo, which answers numeric literals itself without a lookup, and
only the addresses it returns are classified. Zone ids are still dropped
from the resolver's answers, where a scoped IPv6 literal comes back as
fe80::1%1.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:06 -07:00
stumpylogandClaude Opus 5 d53a930aed Fix: Stop treating a dotted quad followed by "%" as an IP literal
The literal check stripped everything after the first "%" from any host,
so a name such as 8.8.8.8%2e169-254-169-254.sslip.io was accepted as the
public address 8.8.8.8 without a DNS lookup. requests percent-decodes the
host before connecting, so Remote OCR would then resolve
8.8.8.8.169-254-169-254.sslip.io and reach an internal address even with
internal endpoints disallowed.

A zone id is now stripped only when the part before "%" is an IPv6
address; any other host containing "%" is looked up as a name. URL
validation with internal addresses disallowed also rejects a host that
contains "%" at all, since the name checked there could otherwise differ
from the one an HTTP client decodes and dials.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:05 -07:00
stumpylogandClaude Opus 5 78ca920771 Document outbound connection policy for internal-address settings
The internal-address settings now describe that a hostname is blocked if any resolved address is non-public. Guarded requests connect directly without proxy variables, and webhook requests never follow redirects.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:05 -07:00
stumpylogandClaude Opus 5 4c5255a8dd Remove the URL-rewriting pinned transport
Every consumer now uses the guarded transports, so the request-rewriting
transport and its helpers are no longer needed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:05 -07:00
stumpylogandClaude Opus 5 283e49b1ee Match get_mailbox annotations to what callers actually pass
get_mailbox declared port and security as always-present types, but
MailAccount.imap_port is nullable and imap_security is stored as a plain
integer, so type checking flagged every call site as passing the wrong
type.

Widen the annotations to port: int | None and security: int, matching the
model fields; IntegerChoices members still compare equal to plain ints, so
the existing branching in get_mailbox is unaffected.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:05 -07:00
stumpylogandClaude Opus 5 3130fc3a7c Use shared outbound resolution for IMAP host pinning
get_mailbox validates the IMAP host with resolve_public_addresses and keeps
its existing error messages; the pinned client dials typed addresses.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:49:05 -07:00
stumpylogandClaude Opus 5 8f56c6167b Use the guarded transport for workflow webhooks
A blocked webhook now raises OutboundRequestBlockedError, which the task
treats as an expected failure and does not retry. The webhook security
tests run against a real local server instead of a patched resolver.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:48:44 -07:00
stumpylogandClaude Opus 5 eb644f3fac Report outbound policy blocks from AI requests as 502
AIClient raises LLMBlockedError when a request was refused by the outbound
connection policy, including when the openai SDK wraps the block in
APIConnectionError. ai_suggestions answers 502 instead of a 500.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:48:44 -07:00
stumpylogandClaude Opus 5 6ad00ca55a Use guarded transports for AI LLM and embedding clients
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:48:21 -07:00
stumpylogandClaude Opus 5 526adbad1a Fix two weak assertions in the outbound guard integration tests
The environment proxy test could not fail for any regression in the guard:
httpx ignores environment proxy variables whenever a transport is passed
explicitly, and the address the test pointed the proxy at was itself
internal, so a blocked direct connection was indistinguishable from a
blocked proxied one. It now builds the client through the production
factory, which does not pass a transport, points the proxy variables at a
second recording server, and asserts the real origin server receives the
request while the proxy server sees no connection at all.

Several tests asserting a blocked request never reached the target server
checked only a connection counter that is incremented after accept() in
the server thread, which does not rule out a guard that connects and then
fails validation afterward. Each of those tests now also asserts the dial
recorder saw no address dialled.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:48:21 -07:00
stumpylogandClaude Opus 5 d53453fb71 Add real-socket tests for the outbound connection guard
A local HTTP server, a per-hostname resolver fake and dial spies exercise
the guarded transports end to end: address fallback, blocking before any
connection, the Host header, TLS server name, per-host connection pooling,
numeric host spellings, environment proxies and redirects to blocked hosts.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:48:20 -07:00
stumpylogandClaude Opus 5 4eaf8032ad Add guarded httpx transports and client factories
The transports install the outbound guard on httpcore's connection pool
after checking its exact layout, and accept no proxy, uds or retries
options. httpx, httpcore and anyio become declared dependencies, pinned
narrowly where private attributes are relied on.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:47:40 -07:00
stumpylogandClaude Opus 5 9e1f938d56 Guard outbound connections in the httpcore network backend
With internal addresses disallowed, the backend resolves the origin host,
rejects it if any address is non-public, and dials the validated literals
in family-interleaved order under the caller's connect timeout. Resolver
failures surface as connect errors; unix sockets are always refused.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:47:40 -07:00
stumpylogandClaude Opus 5 0da50ad348 Reject outbound URLs that HTTP clients may split differently
urllib3 treats a backslash as the end of the URL authority, while urlparse
and httpx do not. For a URL such as http://127.0.0.1\@evil.example/ the
check resolved evil.example while urllib3 would connect to 127.0.0.1, so a
redirect to such a URL could reach an internal host.

When internal addresses are disallowed, validate_outbound_http_url now
rejects any URL containing a backslash, an ASCII control character or
whitespace before resolving it. URLs validated with internal addresses
allowed are unaffected.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:47:40 -07:00
stumpylogandClaude Opus 5 5e971bc0ce Resolve outbound hosts to validated public addresses
resolve_public_addresses and its async twin return every resolved address
in resolver order, de-duplicated and zone-stripped, and reject the whole
name if any address is non-public. validate_outbound_http_url uses them and
keeps its existing messages.

validate_outbound_http_url now resolves the hostname as httpx and urllib3
encode it (IDNA 2008). It previously let getaddrinfo apply the stdlib IDNA
2003 codec, which encodes characters such as "ß" differently, so a URL
could pass the check under one DNS name and be connected to under another.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:47:40 -07:00
stumpylogandClaude Opus 5 17a91c385d Type outbound block errors and classify addresses with is_global
is_public_ip now takes an ipaddress object and relies on is_global, keeping
multicast and the NAT64 well-known prefix as explicit extra exclusions.
Adds OutboundRequestBlockedError and HostResolutionError, both picklable so
Celery keeps them intact on task failure.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-22 12:47:40 -07:00
Trenton H 15b73b890c Chore: Ban cross-app imports of documents.tests (#14224)
Nothing outside the documents app imports documents.tests any more.

Ruff now bans the module everywhere except src/documents/tests, where in-app imports remain fine.
2026-09-22 08:02:18 -07:00
Trenton H 02d355061f Chore: Move helpers that other test modules imported out of test files (#14223)
The mail message and mailbox builders, the fake libmagic and the classifier preprocessor stub lived inside test_mail.py and test_classifier.py, so other test modules imported them by importing a test module.

They now live in helpers modules beside the tests that use them.
2026-09-22 08:02:18 -07:00
Trenton H d8b5b4d447 Chore: Stop the live-service retry helper sleeping after a success (#14222)
util_call_with_backoff made every call wait 20 seconds even when the first attempt succeeded

The helper now sleeps only after a failure that has another attempt left.
2026-09-22 08:02:18 -07:00
Trenton H 03ac4aed7e Chore: Move the cross-app test helpers into the shared layer and add a progress fixture (#14221)
Test modules in paperless, paperless_mail and documents imported filesystem assertions, the migration test base, the retry helper and the streaming-response reader out of documents/tests/utils.py, which kept each app's tests coupled to another app's test package.

They now live in paperless_testing, and the progress manager fake is renamed FakeProgressManager and now subclasses the real ProgressManager, overriding only the transport, so the payload it records is built by the production code. The twenty places that patched documents.tasks.ProgressManager by hand now use a fake_progress_manager fixture.
2026-09-22 08:02:17 -07:00
Trenton H 90d23bad9c Chore: Give the search index directory one owner in the tests (#14217)
Two fixtures created a temporary index directory and pointed INDEX_DIR at it, and paperless_dirs did the same, so a test that requested more than one got whichever assignment ran last. The search conftest no longer defines its own index_dir fixture, the tests that took it read paperless_dirs.index_dir instead, and _search_index is now a thin wrapper that requests paperless_dirs. The fixture that yields a Document is renamed from indexed_document to searchable_document so it no longer differs by one character from the index_document factory next to it.
2026-09-22 08:02:17 -07:00
Trenton H e4367b5648 Chore: Delete the in-memory Tantivy backend and run the search tests on disk (#14216)
TantivyBackend(path=None) built an in-memory index that is not used by production ever used. The backend now requires a path, the open, write-batch and rebuild branches are gone, and the shared backend fixture and the fulltext similar-documents fixture build a real index under the per-test directory layout.
2026-09-22 08:02:16 -07:00
Trenton H cb85441c2f Chore: Decrease backend test suite time (a little) (#14213)
* Chore: Speed up test setup by hashing passwords with MD5 and batching index writes

Don't use Django's default PBKDF2, about 600 ms per use and 110 uses across the suite. Switches to MD5 instead.

Also fixes a test that didn't batch update the search index

* Chore: Stop the invalid webhook params test from waiting on a Celery broker

test_workflow_webhook_action_url_invalid_params_headers left send_webhook.apply_async unpatched, so it tried to actually enqueu and waited for the timeout.
2026-09-21 20:17:40 +00:00
Trenton H cceaa559d4 Chore: Rename sample directory fixtures and drop unused parser ones (#14215)
Nineteen single-file fixtures in the parsers conftest had no consumers anywhere in the test tree.

Two fixtures were both called samples_dir and resolved one directory apart They are now document_samples_dir and parser_samples_dir
2026-09-21 12:59:32 -07:00
GitHub Actions a748d4c64f Auto translate strings 2026-09-21 19:01:49 +00:00
shamoon 452ed005bd Fix: handle legacy bulk edit page range with missing page_count (#14212) 2026-09-21 19:00:14 +00:00
Trenton H 40058ff7d5 Chore: Give the test suite a larger regex timeout (#14211)
Maybe the random ordering sometimes causes heavier tests to run alongside timed regex ones?
2026-09-21 11:38:54 -07:00
70 changed files with 4490 additions and 1809 deletions
+6 -1
View File
@@ -1576,6 +1576,9 @@ ports.
#### [`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).
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.
@@ -1584,7 +1587,7 @@ ports.
#### [`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
configured IMAP hostname resolves to a non-public address (for example,
configured IMAP hostname resolves to any non-public address (for example,
localhost, link-local, or RFC1918 private ranges).
Defaults to true, which allows internal hosts.
@@ -2214,6 +2217,8 @@ 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}
: 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.
+1 -1
View File
@@ -613,7 +613,7 @@ The following workflow action types are available:
- 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
[configuration settings](configuration.md#workflow-webhooks) to change this behavior. If you are allowing non-admins to create workflows,
[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,
you may want to adjust these settings to prevent abuse.
##### Move to Trash {#workflow-action-move-to-trash}
+9
View File
@@ -17,6 +17,7 @@ classifiers = [
# TODO: Move certain things to groups and then utilize that further
# This will allow testing to not install a webserver, mysql, etc
dependencies = [
"anyio>=4.12",
"azure-ai-documentintelligence>=1.0.2",
"babel>=2.17",
"bleach~=6.4.0",
@@ -47,6 +48,8 @@ dependencies = [
"filelock~=3.32.0",
"flower>=2.0.1,<2.2",
"gotenberg-client[httpx]~=1.0",
"httpcore~=1.0.9",
"httpx~=0.28.1",
"httpx-oauth~=0.17",
"ijson>=3.5.1",
"imap-tools>=1.14,<1.16",
@@ -247,6 +250,10 @@ per-file-ignores."docker/wait-for-redis.py" = [
per-file-ignores."src/documents/models.py" = [
"SIM115",
]
per-file-ignores."src/documents/tests/*.py" = [
"TID251",
]
flake8-tidy-imports.banned-api."documents.tests".msg = "Shared test infrastructure lives in src/paperless_testing/."
isort.force-single-line = true
[tool.codespell]
@@ -329,6 +336,8 @@ PAPERLESS_CACHE_BACKEND = "django.core.cache.backends.locmem.LocMemCache"
PAPERLESS_CHANNELS_BACKEND = "channels.layers.InMemoryChannelLayer"
# I don't think anything hits this, but just in case, basically infinite
PAPERLESS_TOKEN_THROTTLE_RATE = "1000/min"
# The 0.1s production default trips on a stalled CI runner, the date parsing tests then find no dates
PAPERLESS_MATCH_REGEX_TIMEOUT_SECONDS = "5"
[tool.coverage.run]
source = [
+67
View File
@@ -17,9 +17,14 @@ if TYPE_CHECKING:
from django.contrib.auth.models import User
from pytest_django.fixtures import Settings
from pytest_mock import MockerFixture
from rest_framework.test import APIClient
from paperless_testing.dirs import PaperlessDirs
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)
@@ -31,6 +36,18 @@ def faker_session_locale() -> str:
return "en_US"
@pytest.fixture(autouse=True)
def _fast_password_hasher(settings: Settings) -> None:
"""Hash test passwords with MD5 instead of Django's default PBKDF2.
PBKDF2 is deliberately slow, and every ``admin_user`` or
``create_superuser`` call pays for it: about 600 ms each. No test depends
on the hash format, only on ``check_password`` and on the stored value
changing when the password does.
"""
settings.PASSWORD_HASHERS = ["django.contrib.auth.hashers.MD5PasswordHasher"]
@pytest.fixture(autouse=True)
def _clear_content_type_caches() -> None:
"""Clear Django's ContentType cache and guardian's lru_cache before each test.
@@ -124,3 +141,53 @@ def user_client(rest_api_client: APIClient, regular_user: User) -> APIClient:
rest_api_client.force_authenticate(user=regular_user)
rest_api_client.credentials(HTTP_ACCEPT="application/json; version=10")
return rest_api_client
@pytest.fixture
def fake_progress_manager(
monkeypatch: pytest.MonkeyPatch,
) -> type[FakeProgressManager]:
"""Replace documents.tasks.ProgressManager with the fake, so consuming a file
in a test never tries to reach a broker."""
from paperless_testing.fakes.progress import FakeProgressManager
monkeypatch.setattr("documents.tasks.ProgressManager", 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)
+50 -63
View File
@@ -196,52 +196,49 @@ class WriteBatch:
return self._raw_writer
def __enter__(self) -> Self:
if self._backend._path is not None:
lock_path = self._backend._path / ".tantivy.lock"
self._lock = filelock.FileLock(str(lock_path))
for attempt in range(_LOCK_RETRY_ATTEMPTS):
try:
self._lock.acquire(timeout=self._lock_timeout)
break
except filelock.Timeout:
if attempt == _LOCK_RETRY_ATTEMPTS - 1:
raise SearchIndexLockError(
f"Could not acquire index lock after {_LOCK_RETRY_ATTEMPTS} "
f"attempts (timeout={self._lock_timeout}s each)",
)
sleep_s = random.uniform(
0,
min(_LOCK_BACKOFF_CAP, _LOCK_BACKOFF_BASE * (2**attempt)),
lock_path = self._backend._path / ".tantivy.lock"
self._lock = filelock.FileLock(str(lock_path))
for attempt in range(_LOCK_RETRY_ATTEMPTS):
try:
self._lock.acquire(timeout=self._lock_timeout)
break
except filelock.Timeout:
if attempt == _LOCK_RETRY_ATTEMPTS - 1:
raise SearchIndexLockError(
f"Could not acquire index lock after {_LOCK_RETRY_ATTEMPTS} "
f"attempts (timeout={self._lock_timeout}s each)",
)
logger.debug(
"Index lock contention; retrying in %.2fs (attempt %d/%d)",
sleep_s,
attempt + 1,
_LOCK_RETRY_ATTEMPTS,
)
time.sleep(sleep_s)
sleep_s = random.uniform(
0,
min(_LOCK_BACKOFF_CAP, _LOCK_BACKOFF_BASE * (2**attempt)),
)
logger.debug(
"Index lock contention; retrying in %.2fs (attempt %d/%d)",
sleep_s,
attempt + 1,
_LOCK_RETRY_ATTEMPTS,
)
time.sleep(sleep_s)
# Open a fresh Index (and thus a fresh Tantivy ManagedDirectory)
# for the write, rather than reusing the process-local cached
# index. ManagedDirectory loads its GC bookkeeping (.managed.json)
# once, at construction, and never re-reads it; paperless runs
# several long-lived processes (Granian workers, Celery workers)
# that take turns writing under the file lock above. A cached,
# long-lived writer index would carry a stale managed-files view
# and, on commit, overwrite .managed.json with that stale view -
# permanently losing track of segment files other processes
# registered in the meantime, so they can never be garbage
# collected. Reopening fresh here always picks up the current
# on-disk state. The long-lived self._backend._index is used for
# reads only and is reloaded (not reopened) after commit below.
write_index = tantivy.Index(
build_schema(),
path=str(self._backend._path),
)
register_tokenizers(write_index, settings.SEARCH_LANGUAGE)
self._raw_writer = write_index.writer()
else:
self._raw_writer = self._backend._index.writer()
# Open a fresh Index (and thus a fresh Tantivy ManagedDirectory)
# for the write, rather than reusing the process-local cached
# index. ManagedDirectory loads its GC bookkeeping (.managed.json)
# once, at construction, and never re-reads it; paperless runs
# several long-lived processes (Granian workers, Celery workers)
# that take turns writing under the file lock above. A cached,
# long-lived writer index would carry a stale managed-files view
# and, on commit, overwrite .managed.json with that stale view -
# permanently losing track of segment files other processes
# registered in the meantime, so they can never be garbage
# collected. Reopening fresh here always picks up the current
# on-disk state. The long-lived self._backend._index is used for
# reads only and is reloaded (not reopened) after commit below.
write_index = tantivy.Index(
build_schema(),
path=str(self._backend._path),
)
register_tokenizers(write_index, settings.SEARCH_LANGUAGE)
self._raw_writer = write_index.writer()
return self
def __exit__(self, exc_type, exc_val, exc_tb):
@@ -372,9 +369,8 @@ class TantivyBackend:
Tantivy search backend with explicit lifecycle management.
Provides full-text search capabilities using the Tantivy search engine.
Supports in-memory indexes (for testing) and persistent on-disk indexes
(for production use). Handles document indexing, search queries, autocompletion,
and "more like this" functionality.
Keeps a persistent on-disk index. Handles document indexing, search queries,
autocompletion, and "more like this" functionality.
The backend manages its own connection lifecycle and can be reset when
the underlying index directory changes (e.g., during test isolation).
@@ -408,9 +404,7 @@ class TantivyBackend:
},
)
def __init__(self, path: Path | None = None):
# path=None → in-memory index (for tests)
# path=some_dir → on-disk index (for production)
def __init__(self, path: Path):
self._path = path
self._raw_index: tantivy.Index | None = None
self._raw_schema: tantivy.Schema | None = None
@@ -429,16 +423,13 @@ class TantivyBackend:
"""
Open or rebuild the index as needed.
For disk-based indexes, checks if rebuilding is needed due to schema
version or language changes. Registers custom tokenizers after opening.
Checks if rebuilding is needed due to schema version or language
changes. Registers custom tokenizers after opening.
Safe to call multiple times - subsequent calls are no-ops.
"""
if self._raw_index is not None:
return # pragma: no cover
if self._path is not None:
self._raw_index = open_or_rebuild_index(self._path)
else:
self._raw_index = tantivy.Index(build_schema())
self._raw_index = open_or_rebuild_index(self._path)
register_tokenizers(self._raw_index, settings.SEARCH_LANGUAGE)
self._raw_schema = self._raw_index.schema
@@ -1102,13 +1093,9 @@ class TantivyBackend:
writer's threads). Larger values buffer more docs in RAM before
flushing a segment, deferring merge work; they do not avoid it.
"""
# Create new index (on-disk or in-memory)
if self._path is not None:
wipe_index(self._path)
new_index = tantivy.Index(build_schema(), path=str(self._path))
_write_sentinels(self._path)
else:
new_index = tantivy.Index(build_schema())
wipe_index(self._path)
new_index = tantivy.Index(build_schema(), path=str(self._path))
_write_sentinels(self._path)
register_tokenizers(new_index, settings.SEARCH_LANGUAGE)
# Point instance at the new index so _build_tantivy_doc uses it
+3 -1
View File
@@ -2098,6 +2098,8 @@ class BulkEditSerializer(
if not isinstance(parameters["pages"], str):
raise serializers.ValidationError("invalid pages specified")
page_count = Document.objects.get(id=document_id).page_count
if not page_count:
raise serializers.ValidationError("document page count is unknown")
pages = []
for group in parameters["pages"].split(","):
start, is_range, end = group.partition("-")
@@ -2107,7 +2109,7 @@ class BulkEditSerializer(
except ValueError as e:
raise serializers.ValidationError("invalid pages specified") from e
# Bound the range before building it, a huge one would exhaust memory
if not 1 <= first <= last or (page_count and last > page_count):
if not 1 <= first <= last <= page_count:
raise serializers.ValidationError("invalid pages specified")
pages.append(list(range(first, last + 1)))
parameters["pages"] = pages
+11 -24
View File
@@ -1,11 +1,9 @@
import shutil
from collections.abc import Generator
from pathlib import Path
from typing import TYPE_CHECKING
import filelock
import pytest
from pytest_django.fixtures import Settings
from paperless_testing.factories import DocumentFactory
@@ -15,7 +13,7 @@ if TYPE_CHECKING:
@pytest.fixture(scope="session")
def samples_dir() -> Path:
def document_samples_dir() -> Path:
"""Path to the shared test sample documents."""
return Path(__file__).parent / "samples" / "documents"
@@ -23,20 +21,20 @@ def samples_dir() -> Path:
@pytest.fixture()
def sample_doc(
paperless_dirs: "PaperlessDirs",
samples_dir: Path,
document_samples_dir: Path,
) -> "Document":
"""Create a document with valid files and matching checksums."""
with filelock.FileLock(paperless_dirs.media_lock):
shutil.copy(
samples_dir / "originals" / "0000001.pdf",
document_samples_dir / "originals" / "0000001.pdf",
paperless_dirs.originals_dir / "0000001.pdf",
)
shutil.copy(
samples_dir / "archive" / "0000001.pdf",
document_samples_dir / "archive" / "0000001.pdf",
paperless_dirs.archive_dir / "0000001.pdf",
)
shutil.copy(
samples_dir / "thumbnails" / "0000001.webp",
document_samples_dir / "thumbnails" / "0000001.webp",
paperless_dirs.thumbnail_dir / "0000001.webp",
)
@@ -52,28 +50,17 @@ def sample_doc(
)
@pytest.fixture()
def _search_index(
tmp_path: Path,
settings: Settings,
) -> Generator[None, None, None]:
"""Create a temp index directory and point INDEX_DIR at it.
@pytest.fixture
def _search_index(paperless_dirs: "PaperlessDirs") -> None:
"""Point the search backend at a fresh, empty index directory.
Resets the backend singleton before and after so each test gets a clean
index rather than reusing a stale singleton from another test.
paperless_dirs owns INDEX_DIR and resets the backend singleton on both
sides of the test, so requesting it is all that is needed.
"""
from documents.search import reset_backend
index_dir = tmp_path / "index"
index_dir.mkdir()
settings.INDEX_DIR = index_dir
reset_backend()
yield
reset_backend()
@pytest.fixture
def indexed_document(_search_index: None) -> "Document":
def searchable_document(_search_index: None) -> "Document":
"""One searchable document, for tests about what the search endpoint
returns rather than about what it finds.
"""
+10
View File
@@ -0,0 +1,10 @@
import re
def dummy_preprocess(content: str) -> str:
"""
Simpler, faster pre-processing for testing purposes
"""
content = content.lower().strip()
content = re.sub(r"\s+", " ", content)
return content
+3 -11
View File
@@ -14,24 +14,16 @@ from paperless_testing.factories import DocumentFactory
if TYPE_CHECKING:
from collections.abc import Callable
from collections.abc import Generator
from pathlib import Path
from pytest_django.fixtures import Settings
from documents.models import Document
from paperless_testing.dirs import PaperlessDirs
@pytest.fixture
def index_dir(tmp_path: Path, settings: Settings) -> Path:
path = tmp_path / "index"
path.mkdir()
settings.INDEX_DIR = path
return path
@pytest.fixture
def backend() -> Generator[TantivyBackend, None, None]:
b = TantivyBackend() # path=None → in-memory index
def backend(paperless_dirs: PaperlessDirs) -> Generator[TantivyBackend, None, None]:
b = TantivyBackend(path=paperless_dirs.index_dir)
b.open()
try:
yield b
+4 -2
View File
@@ -947,7 +947,8 @@ class TestSingleton:
yield
reset_backend()
def test_returns_same_instance_on_repeated_calls(self, index_dir) -> None:
@pytest.mark.usefixtures("paperless_dirs")
def test_returns_same_instance_on_repeated_calls(self) -> None:
"""Singleton pattern: repeated calls to get_backend() must return the same instance."""
assert get_backend() is get_backend()
@@ -964,7 +965,8 @@ class TestSingleton:
assert b1 is not b2
assert b2._path == tmp_path / "b"
def test_reset_forces_new_instance(self, index_dir) -> None:
@pytest.mark.usefixtures("paperless_dirs")
def test_reset_forces_new_instance(self) -> None:
"""reset_backend() must force creation of a new backend instance on next get_backend() call."""
b1 = get_backend()
reset_backend()
@@ -269,7 +269,7 @@ class TestDocumentedDateForms:
yield
@pytest.fixture
def dated(self, index_document: Callable[..., Document]) -> dict[str, int]:
def dated(self, backend: TantivyBackend) -> dict[str, int]:
stamps = {
"today": datetime(2026, 6, 15, 9, 0, tzinfo=UTC),
"yesterday": datetime(2026, 6, 14, 9, 0, tzinfo=UTC),
@@ -279,14 +279,14 @@ class TestDocumentedDateForms:
"january": datetime(2026, 1, 10, 10, 0, tzinfo=UTC),
"old": datetime(2005, 3, 4, 15, 30, tzinfo=UTC),
}
return {
label: index_document(
title=label,
content="dated body",
added=stamp,
).pk
docs = {
label: DocumentFactory(title=label, content="dated body", added=stamp)
for label, stamp in stamps.items()
}
with backend.batch_update() as batch:
for doc in docs.values():
batch.add_or_update(doc)
return {label: doc.pk for label, doc in docs.items()}
@pytest.mark.parametrize(
("query", "label"),
@@ -1,6 +1,6 @@
import pytest
from documents.tests.utils import TestMigrations
from paperless_testing.migrations import TestMigrations
pytestmark = pytest.mark.search
+22 -19
View File
@@ -13,11 +13,11 @@ from documents.search._schema import needs_rebuild
from documents.search._schema import schema_fingerprint
if TYPE_CHECKING:
from pathlib import Path
import tantivy
from pytest_django.fixtures import Settings
from paperless_testing.dirs import PaperlessDirs
pytestmark = pytest.mark.search
@@ -25,16 +25,19 @@ pytestmark = pytest.mark.search
class TestNeedsRebuild:
"""needs_rebuild covers all sentinel-file states that require a full reindex."""
def test_returns_true_when_settings_file_missing(self, index_dir: Path) -> None:
assert needs_rebuild(index_dir) is True
def test_returns_true_when_settings_file_missing(
self,
paperless_dirs: PaperlessDirs,
) -> None:
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_returns_false_when_version_and_language_match(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: Settings,
) -> None:
settings.SEARCH_LANGUAGE = "en"
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps(
{
"schema_version": SCHEMA_VERSION,
@@ -43,51 +46,51 @@ class TestNeedsRebuild:
},
),
)
assert needs_rebuild(index_dir) is False
assert needs_rebuild(paperless_dirs.index_dir) is False
def test_returns_true_on_schema_version_mismatch(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: Settings,
) -> None:
settings.SEARCH_LANGUAGE = None
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps({"schema_version": SCHEMA_VERSION - 1, "language": None}),
)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_returns_true_when_version_is_not_an_integer(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: Settings,
) -> None:
settings.SEARCH_LANGUAGE = None
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps({"schema_version": "not-a-number", "language": None}),
)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_returns_true_when_language_key_missing(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: Settings,
) -> None:
settings.SEARCH_LANGUAGE = "en"
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps({"schema_version": SCHEMA_VERSION}),
)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_returns_true_when_language_differs(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: Settings,
) -> None:
settings.SEARCH_LANGUAGE = "de"
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps({"schema_version": SCHEMA_VERSION, "language": "en"}),
)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def _schema_fields(schema: tantivy.Schema) -> dict[str, dict]:
@@ -35,6 +35,8 @@ if TYPE_CHECKING:
from pytest_django.fixtures import SettingsWrapper
from paperless_testing.dirs import PaperlessDirs
pytestmark = pytest.mark.search
# The on-disk field layout of a v2 index, pinned as data. Any edit here is an
@@ -469,7 +471,7 @@ def _fingerprint_of(descriptors: list[FieldDescriptor]) -> str:
class TestNeedsRebuildOnFingerprint:
def test_matching_fingerprint_does_not_rebuild(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: SettingsWrapper,
) -> None:
"""
@@ -482,13 +484,13 @@ class TestNeedsRebuildOnFingerprint:
- It returns False
"""
settings.SEARCH_LANGUAGE = None
_sentinels(index_dir)
_sentinels(paperless_dirs.index_dir)
assert needs_rebuild(index_dir) is False
assert needs_rebuild(paperless_dirs.index_dir) is False
def test_stale_fingerprint_rebuilds_despite_a_matching_version(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: SettingsWrapper,
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -505,7 +507,7 @@ class TestNeedsRebuildOnFingerprint:
every subsequent write would raise
"""
settings.SEARCH_LANGUAGE = None
_sentinels(index_dir)
_sentinels(paperless_dirs.index_dir)
extended = [
*field_descriptors(),
FieldDescriptor(
@@ -519,11 +521,11 @@ class TestNeedsRebuildOnFingerprint:
]
monkeypatch.setattr(_schema, "field_descriptors", lambda: extended)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_reordered_schema_rebuilds(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: SettingsWrapper,
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -538,16 +540,16 @@ class TestNeedsRebuildOnFingerprint:
- It returns True
"""
settings.SEARCH_LANGUAGE = None
_sentinels(index_dir)
_sentinels(paperless_dirs.index_dir)
reordered = field_descriptors()
reordered[1], reordered[2] = reordered[2], reordered[1]
monkeypatch.setattr(_schema, "field_descriptors", lambda: reordered)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_missing_fingerprint_rebuilds(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: SettingsWrapper,
) -> None:
"""
@@ -561,15 +563,15 @@ class TestNeedsRebuildOnFingerprint:
is rebuilt rather than trusted
"""
settings.SEARCH_LANGUAGE = None
(index_dir / ".index_settings.json").write_text(
(paperless_dirs.index_dir / ".index_settings.json").write_text(
json.dumps({"schema_version": SCHEMA_VERSION, "language": None}),
)
assert needs_rebuild(index_dir) is True
assert needs_rebuild(paperless_dirs.index_dir) is True
def test_written_sentinels_satisfy_the_check(
self,
index_dir: Path,
paperless_dirs: PaperlessDirs,
settings: SettingsWrapper,
) -> None:
"""
@@ -582,6 +584,6 @@ class TestNeedsRebuildOnFingerprint:
- It returns False
"""
settings.SEARCH_LANGUAGE = "en"
_write_sentinels(index_dir)
_write_sentinels(paperless_dirs.index_dir)
assert needs_rebuild(index_dir) is False
assert needs_rebuild(paperless_dirs.index_dir) is False
+1 -1
View File
@@ -10,11 +10,11 @@ from PIL.PngImagePlugin import PngInfo
from rest_framework import status
from rest_framework.test import APITestCase
from documents.tests.utils import read_streaming_response
from paperless.models import ApplicationConfiguration
from paperless.models import ColorConvertChoices
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import UserFactory
from paperless_testing.http import read_streaming_response
class TestApiAppConfig(DirectoriesMixin, APITestCase):
@@ -13,9 +13,9 @@ from documents.models import Correspondent
from documents.models import Document
from documents.models import DocumentType
from documents.tests.utils import SampleDirMixin
from documents.tests.utils import read_streaming_response
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import UserFactory
from paperless_testing.http import read_streaming_response
from paperless_testing.permissions import grant_global
+30
View File
@@ -1784,6 +1784,36 @@ class TestBulkEditAPI(DirectoriesMixin, APITestCase):
self.assertIn(b"invalid pages specified", response.content)
m.assert_not_called()
@mock.patch("documents.serialisers.bulk_edit.split")
def test_bulk_edit_split_rejects_unknown_page_count(self, m) -> None:
"""
GIVEN:
- A legacy split bulk edit of a document without a page count
WHEN:
- API to bulk edit is called
THEN:
- API returns HTTP 400
- split is not called
"""
self.setup_mock(m, "split")
for pages in ("1", "1-5000000"):
with self.subTest(pages=pages):
response = self.client.post(
"/api/documents/bulk_edit/",
json.dumps(
{
"documents": [self.doc1.id],
"method": "split",
"parameters": {"pages": pages},
},
),
content_type="application/json",
)
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn(b"document page count is unknown", response.content)
m.assert_not_called()
@mock.patch("documents.serialisers.bulk_edit.split")
def test_bulk_edit_split_parses_pages(self, m) -> None:
"""
@@ -16,11 +16,11 @@ from documents.data_models import DocumentSource
from documents.filters import EffectiveContentFilter
from documents.filters import TitleContentFilter
from documents.models import Document
from documents.tests.utils import read_streaming_response
from documents.versioning import annotate_effective_content
from documents.views import DocumentSelectionMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import UserFactory
from paperless_testing.http import read_streaming_response
from paperless_testing.permissions import grant_global
if TYPE_CHECKING:
+1 -1
View File
@@ -48,11 +48,11 @@ from documents.models import WorkflowAction
from documents.models import WorkflowTrigger
from documents.signals.handlers import run_workflows
from documents.tests.utils import ConsumeTaskMixin
from documents.tests.utils import read_streaming_response
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import DocumentFactory
from paperless_testing.factories import TagFactory
from paperless_testing.factories import UserFactory
from paperless_testing.http import read_streaming_response
from paperless_testing.permissions import grant_all_global
from paperless_testing.permissions import grant_global
from paperless_testing.permissions import grant_object
+10 -10
View File
@@ -33,7 +33,7 @@ class TestSearchQueryErrorStillBecomesA400:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -68,7 +68,7 @@ class TestLibraryDefectsPropagate:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -98,7 +98,7 @@ class TestLibraryDefectsPropagate:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -141,7 +141,7 @@ class TestSelectionPathsAgreeWithSearch:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -181,7 +181,7 @@ class TestSelectionPathsAgreeWithSearch:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -221,7 +221,7 @@ class TestSelectionPathsAgreeWithSearch:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -259,7 +259,7 @@ class TestSelectionPathsAgreeWithSearch:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -287,7 +287,7 @@ class TestSelectionPathsAgreeWithSearch:
{
"documents": [],
"all": True,
"filters": {"more_like_id": indexed_document.pk},
"filters": {"more_like_id": searchable_document.pk},
},
format="json",
)
@@ -298,7 +298,7 @@ class TestSelectionPathsAgreeWithSearch:
self,
admin_client: APIClient,
monkeypatch: pytest.MonkeyPatch,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -328,7 +328,7 @@ class TestSelectionPathsAgreeWithSearch:
{
"documents": [],
"all": True,
"filters": {"more_like_id": indexed_document.pk},
"filters": {"more_like_id": searchable_document.pk},
},
format="json",
)
@@ -36,7 +36,7 @@ class TestGetSearchEndpointEnforcesTheCap:
def test_query_one_over_the_cap_is_a_400(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -66,7 +66,7 @@ class TestGetSearchEndpointEnforcesTheCap:
def test_query_at_exactly_the_cap_is_accepted(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -86,7 +86,7 @@ class TestGetSearchEndpointEnforcesTheCap:
def test_an_ordinary_query_is_unaffected(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -112,7 +112,7 @@ class TestPostSelectionPathsEnforceTheCap:
def test_bulk_edit_query_one_over_the_cap_is_a_400(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -154,7 +154,7 @@ class TestPostSelectionPathsEnforceTheCap:
self,
bulk_update_task_mock: mock.MagicMock,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -187,7 +187,7 @@ class TestPostSelectionPathsEnforceTheCap:
def test_bulk_download_query_one_over_the_cap_is_a_400(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -236,7 +236,7 @@ class TestGlobalSearchEnforcesTheCapToo:
def test_query_one_over_the_cap_is_a_400(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -260,7 +260,7 @@ class TestGlobalSearchEnforcesTheCapToo:
def test_query_at_exactly_the_cap_is_accepted(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
@@ -38,7 +38,7 @@ class TestUnterminatedBracketReturnsA400:
def test_unterminated_bracket_is_a_400(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
query: str,
) -> None:
"""
@@ -59,7 +59,7 @@ class TestUnterminatedBracketReturnsA400:
def test_properly_closed_bracket_still_searches_cleanly(
self,
admin_client: APIClient,
indexed_document: Document,
searchable_document: Document,
) -> None:
"""
GIVEN:
+74 -74
View File
@@ -2,8 +2,8 @@ import shutil
from collections.abc import Generator
from contextlib import contextmanager
from pathlib import Path
from unittest import mock
import pytest
from django.conf import settings
from django.test import TestCase
from django.test import override_settings
@@ -18,11 +18,11 @@ from documents.models import Document
from documents.models import Tag
from documents.plugins.base import StopConsumeTaskError
from documents.tests.utils import ConsumeTaskMixin
from documents.tests.utils import DummyProgressManager
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.utils import SampleDirMixin
from paperless.models import ApplicationConfiguration
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.fakes.progress import FakeProgressManager
class GetReaderPluginMixin:
@@ -31,7 +31,7 @@ class GetReaderPluginMixin:
reader = BarcodePlugin(
ConsumableDocument(DocumentSource.ConsumeFolder, original_file=filepath),
DocumentMetadataOverrides(),
DummyProgressManager(filepath.name, None),
FakeProgressManager(filepath.name, None),
self.dirs.scratch_dir,
"task-id",
)
@@ -86,6 +86,7 @@ class TestBarcode(
self.assertDictEqual(separator_page_numbers, {1: False})
@override_settings(CONSUMER_ENABLE_ASN_BARCODE=True)
@pytest.mark.usefixtures("fake_progress_manager")
def test_asn_barcode_duplicate_in_trash_fails(self) -> None:
"""
GIVEN:
@@ -110,15 +111,14 @@ class TestBarcode(
dupe_asn = settings.SCRATCH_DIR / "barcode-39-asn-123-second.pdf"
shutil.copy(test_file, dupe_asn)
with mock.patch("documents.tasks.ProgressManager", DummyProgressManager):
with self.assertRaisesRegex(ConsumerError, r"ASN 123.*trash"):
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dupe_asn,
),
None,
)
with self.assertRaisesRegex(ConsumerError, r"ASN 123.*trash"):
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dupe_asn,
),
None,
)
@override_settings(
CONSUMER_BARCODE_TIFF_SUPPORT=True,
@@ -606,6 +606,7 @@ class TestBarcodeNewConsume(
TestCase,
):
@override_settings(CONSUMER_ENABLE_BARCODES=True)
@pytest.mark.usefixtures("fake_progress_manager")
def test_consume_barcode_file(self) -> None:
"""
GIVEN:
@@ -624,34 +625,33 @@ class TestBarcodeNewConsume(
overrides = DocumentMetadataOverrides(tag_ids=[1, 2, 9])
with mock.patch("documents.tasks.ProgressManager", DummyProgressManager):
self.assertEqual(
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=temp_copy,
),
overrides,
self.assertEqual(
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=temp_copy,
),
{"reason": "Barcode splitting complete!"},
)
# 2 new document consume tasks created
self.assertEqual(self.consume_file_mock.call_count, 2)
overrides,
),
{"reason": "Barcode splitting complete!"},
)
# 2 new document consume tasks created
self.assertEqual(self.consume_file_mock.call_count, 2)
self.assertIsNotFile(temp_copy)
self.assertIsNotFile(temp_copy)
# Check the split files exist
# Check the original_path is set
# Check the source is unchanged
# Check the overrides are unchanged
for (
new_input_doc,
new_doc_overrides,
) in self.get_all_consume_task_call_args():
self.assertIsFile(new_input_doc.original_file)
self.assertEqual(new_input_doc.original_path, temp_copy)
self.assertEqual(new_input_doc.source, DocumentSource.ConsumeFolder)
self.assertEqual(overrides, new_doc_overrides)
# Check the split files exist
# Check the original_path is set
# Check the source is unchanged
# Check the overrides are unchanged
for (
new_input_doc,
new_doc_overrides,
) in self.get_all_consume_task_call_args():
self.assertIsFile(new_input_doc.original_file)
self.assertEqual(new_input_doc.original_path, temp_copy)
self.assertEqual(new_input_doc.source, DocumentSource.ConsumeFolder)
self.assertEqual(overrides, new_doc_overrides)
class TestAsnBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, TestCase):
@@ -660,7 +660,7 @@ class TestAsnBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
reader = BarcodePlugin(
ConsumableDocument(DocumentSource.ConsumeFolder, original_file=filepath),
DocumentMetadataOverrides(),
DummyProgressManager(filepath.name, None),
FakeProgressManager(filepath.name, None),
self.dirs.scratch_dir,
"task-id",
)
@@ -745,6 +745,7 @@ class TestAsnBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
self.assertEqual(asn, None)
@override_settings(CONSUMER_ENABLE_ASN_BARCODE=True)
@pytest.mark.usefixtures("fake_progress_manager")
def test_consume_barcode_file_asn_assignment(self) -> None:
"""
GIVEN:
@@ -762,19 +763,18 @@ class TestAsnBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
dst = settings.SCRATCH_DIR / "barcode-39-asn-123.pdf"
shutil.copy(test_file, dst)
with mock.patch("documents.tasks.ProgressManager", DummyProgressManager):
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
document = Document.objects.first()
assert document is not None
document = Document.objects.first()
assert document is not None
self.assertEqual(document.archive_serial_number, 123)
self.assertEqual(document.archive_serial_number, 123)
def test_scan_file_for_qrcode_without_upscale(self) -> None:
"""
@@ -819,7 +819,7 @@ class TestTagBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
reader = BarcodePlugin(
ConsumableDocument(DocumentSource.ConsumeFolder, original_file=filepath),
DocumentMetadataOverrides(),
DummyProgressManager(filepath.name, None),
FakeProgressManager(filepath.name, None),
self.dirs.scratch_dir,
"task-id",
)
@@ -1024,6 +1024,7 @@ class TestTagBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
CELERY_TASK_ALWAYS_EAGER=True,
OCR_MODE="auto",
)
@pytest.mark.usefixtures("fake_progress_manager")
def test_consume_barcode_file_tag_split_and_assignment(self) -> None:
"""
GIVEN:
@@ -1042,34 +1043,33 @@ class TestTagBarcode(DirectoriesMixin, SampleDirMixin, GetReaderPluginMixin, Tes
dst = settings.SCRATCH_DIR / "split-by-tag-basic.pdf"
shutil.copy(test_file, dst)
with mock.patch("documents.tasks.ProgressManager", DummyProgressManager):
result = tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
result = tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
self.assertEqual(result, {"reason": "Barcode splitting complete!"})
self.assertEqual(result, {"reason": "Barcode splitting complete!"})
documents = Document.objects.all().order_by("id")
self.assertEqual(documents.count(), 3)
documents = Document.objects.all().order_by("id")
self.assertEqual(documents.count(), 3)
doc1 = documents[0]
self.assertEqual(doc1.tags.count(), 0)
doc1 = documents[0]
self.assertEqual(doc1.tags.count(), 0)
doc2 = documents[1]
self.assertEqual(doc2.tags.count(), 1)
_tag_1 = doc2.tags.first()
assert _tag_1 is not None
self.assertEqual(_tag_1.name, "invoice")
doc2 = documents[1]
self.assertEqual(doc2.tags.count(), 1)
_tag_1 = doc2.tags.first()
assert _tag_1 is not None
self.assertEqual(_tag_1.name, "invoice")
doc3 = documents[2]
self.assertEqual(doc3.tags.count(), 1)
_tag_2 = doc3.tags.first()
assert _tag_2 is not None
self.assertEqual(_tag_2.name, "receipt")
doc3 = documents[2]
self.assertEqual(doc3.tags.count(), 1)
_tag_2 = doc3.tags.first()
assert _tag_2 is not None
self.assertEqual(_tag_2.name, "receipt")
@override_settings(
CONSUMER_ENABLE_TAG_BARCODE=True,
+1 -10
View File
@@ -1,5 +1,4 @@
import pickle
import re
import warnings
from datetime import UTC
from datetime import datetime
@@ -28,6 +27,7 @@ from documents.models import DocumentType
from documents.models import MatchingModel
from documents.models import StoragePath
from documents.models import Tag
from documents.tests.helpers import dummy_preprocess
from paperless.settings import CLASSIFIER_LANGUAGES
from paperless.signed_pickle import HMAC_SIZE
from paperless.signed_pickle import signed_pickle_dumps
@@ -36,15 +36,6 @@ from paperless_testing.factories import DocumentFactory
from paperless_testing.factories import TagFactory
def dummy_preprocess(content: str) -> str:
"""
Simpler, faster pre-processing for testing purposes
"""
content = content.lower().strip()
content = re.sub(r"\s+", " ", content)
return content
class TestClassifier(DirectoriesMixin, TestCase):
def setUp(self) -> None:
super().setUp()
+5 -5
View File
@@ -30,12 +30,12 @@ from documents.models import Tag
from documents.parsers import ParseError
from documents.plugins.helpers import ProgressStatusOptions
from documents.tasks import sanity_check
from documents.tests.utils import DummyProgressManager
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.utils import GetConsumerMixin
from paperless_mail.models import MailRule
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import UserFactory
from paperless_testing.fakes.progress import FakeProgressManager
class _BaseNewStyleParser:
@@ -777,7 +777,7 @@ class TestConsumer(
)
version_file = self.get_test_file2()
status = DummyProgressManager(version_file.name, None)
status = FakeProgressManager(version_file.name, None)
overrides = DocumentMetadataOverrides(
version_label="v2",
actor_id=actor.pk,
@@ -840,7 +840,7 @@ class TestConsumer(
assert root_doc is not None
version_file = self.get_test_file2()
status = DummyProgressManager(version_file.name, None)
status = FakeProgressManager(version_file.name, None)
overrides = DocumentMetadataOverrides(
filename="valid_pdf_version-upload",
actor_id=999999,
@@ -897,7 +897,7 @@ class TestConsumer(
assert root_doc is not None
def consume_version(version_file: Path) -> Document:
status = DummyProgressManager(version_file.name, None)
status = FakeProgressManager(version_file.name, None)
overrides = DocumentMetadataOverrides()
doc = ConsumableDocument(
DocumentSource.ApiUpload,
+17 -17
View File
@@ -2,8 +2,8 @@ import datetime as dt
import os
import shutil
from pathlib import Path
from unittest import mock
import pytest
from django.test import TestCase
from django.test import override_settings
from pdfminer.high_level import extract_text
@@ -15,18 +15,22 @@ from documents.data_models import ConsumableDocument
from documents.data_models import DocumentSource
from documents.double_sided import STAGING_FILE_NAME
from documents.double_sided import TIMEOUT_MINUTES
from documents.tests.utils import DummyProgressManager
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.utils import SampleDirMixin
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
@pytest.mark.usefixtures("fake_progress_manager")
@override_settings(
CONSUMER_RECURSIVE=True,
CONSUMER_ENABLE_COLLATE_DOUBLE_SIDED=True,
)
class TestDoubleSided(DirectoriesMixin, FileSystemAssertsMixin, TestCase):
SAMPLE_DIR = Path(__file__).parent / "samples"
class TestDoubleSided(
DirectoriesMixin,
FileSystemAssertsMixin,
SampleDirMixin,
TestCase,
):
def setUp(self) -> None:
super().setUp()
self.double_sided_dir = self.dirs.consumption_dir / "double-sided"
@@ -42,17 +46,13 @@ class TestDoubleSided(DirectoriesMixin, FileSystemAssertsMixin, TestCase):
dst = self.double_sided_dir / dstname
dst.parent.mkdir(parents=True, exist_ok=True)
shutil.copy(src, dst)
with mock.patch(
"documents.tasks.ProgressManager",
DummyProgressManager,
):
msg = tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
msg = tasks.consume_file(
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=dst,
),
None,
)
self.assertIsNotFile(dst)
return msg
+1 -1
View File
@@ -29,7 +29,7 @@ from documents.models import DocumentType
from documents.models import StoragePath
from documents.serialisers import DocumentSerializer
from documents.tasks import empty_trash
from documents.tests.utils import FileSystemAssertsMixin
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import DocumentFactory
from paperless_testing.factories import UserFactory
+1 -1
View File
@@ -20,7 +20,7 @@ if TYPE_CHECKING:
from documents.file_handling import generate_filename
from documents.models import Document
from documents.tasks import update_document_content_maybe_archive_file
from documents.tests.utils import FileSystemAssertsMixin
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
sample_file: Path = Path(__file__).parent / "samples" / "simple.pdf"
@@ -45,9 +45,9 @@ from documents.models import WorkflowTrigger
from documents.sanity_checker import check_sanity
from documents.settings import EXPORTER_FILE_NAME
from documents.settings import EXPORTER_SHARE_LINK_BUNDLE_NAME
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.utils import SampleDirMixin
from paperless_mail.models import MailAccount
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.dirs import paperless_environment
from paperless_testing.permissions import grant_object
@@ -15,8 +15,8 @@ from documents.management.commands.document_importer import _deserialize_record
from documents.models import Document
from documents.settings import EXPORTER_ARCHIVE_NAME
from documents.settings import EXPORTER_FILE_NAME
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.utils import SampleDirMixin
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
@@ -9,7 +9,7 @@ from django.test import TestCase
from documents.management.commands.document_thumbnails import _process_document
from documents.models import Document
from documents.parsers import get_default_thumbnail
from documents.tests.utils import FileSystemAssertsMixin
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
@@ -1,4 +1,4 @@
from documents.tests.utils import TestMigrations
from paperless_testing.migrations import TestMigrations
SAVED_VIEWS_KEY = "saved_views"
DASHBOARD_VIEWS_VISIBLE_IDS_KEY = "dashboard_views_visible_ids"
@@ -7,7 +7,7 @@ from django.conf import settings
from django.db import connection
from django.test import override_settings
from documents.tests.utils import TestMigrations
from paperless_testing.migrations import TestMigrations
def _sha256(data: bytes) -> str:
@@ -1,4 +1,4 @@
from documents.tests.utils import TestMigrations
from paperless_testing.migrations import TestMigrations
class TestMigrateShareLinkBundlePermissions(TestMigrations):
+37 -2
View File
@@ -17,8 +17,9 @@ from documents.models import Tag
from documents.models import WorkflowAction
from documents.sanity_checker import SanityCheckFailedException
from documents.sanity_checker import SanityCheckMessages
from documents.tests.test_classifier import dummy_preprocess
from documents.tests.utils import FileSystemAssertsMixin
from documents.tests.helpers import dummy_preprocess
from paperless_ai.exceptions import LLMBlockedError
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
@@ -555,3 +556,37 @@ class TestApplyAISuggestionsTask(DirectoriesMixin, TestCase):
apply_suggestions.assert_not_called()
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()
+44 -1
View File
@@ -28,12 +28,13 @@ from documents.models import StoragePath
from documents.models import Tag
from documents.models import UiSettings
from documents.signals.handlers import update_llm_suggestions_cache
from documents.tests.utils import read_streaming_response
from paperless.models import ApplicationConfiguration
from paperless_ai.exceptions import LLMBlockedError
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import UserFactory
from paperless_testing.http import read_streaming_response
from paperless_testing.permissions import grant_global
from paperless_testing.permissions import grant_object
@@ -770,6 +771,48 @@ class TestAISuggestions(DirectoriesMixin, TestCase):
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")
@override_settings(
AI_ENABLED=True,
File diff suppressed because it is too large Load Diff
+2 -225
View File
@@ -1,126 +1,16 @@
import time
import warnings
from collections.abc import Callable
from collections.abc import Generator
from collections.abc import Iterator
from contextlib import contextmanager
from os import PathLike
from pathlib import Path
from typing import Any
from unittest import mock
import httpx
import pytest
from django.apps import apps
from django.db import connection
from django.db.migrations.executor import MigrationExecutor
from django.http import StreamingHttpResponse
from django.test import TransactionTestCase
from documents.consumer import AsnCheckPlugin
from documents.consumer import ConsumerPlugin
from documents.consumer import ConsumerPreflightPlugin
from documents.data_models import ConsumableDocument
from documents.data_models import DocumentMetadataOverrides
from documents.data_models import DocumentSource
from documents.parsers import ParseError
from documents.plugins.helpers import ProgressStatusOptions
def util_call_with_backoff(
method_or_callable: Callable,
args: list | tuple,
*,
skip_on_50x_err=True,
) -> tuple[bool, Any]:
"""
For whatever reason, the images started during the test pipeline like to
segfault sometimes, crash and otherwise fail randomly, when run with the
exact files that usually pass.
So, this function will retry the given method/function up to 3 times, with larger backoff
periods between each attempt, in hopes the issue resolves itself during
one attempt to parse.
This will wait the following:
- Attempt 1 - 20s following failure
- Attempt 2 - 40s following failure
- Attempt 3 - 80s following failure
"""
result = None
succeeded = False
retry_time = 20.0
retry_count = 0
status_codes = []
max_retry_count = 3
while retry_count < max_retry_count and not succeeded:
try:
result = method_or_callable(*args)
succeeded = True
except ParseError as e: # pragma: no cover
cause_exec = e.__cause__
if cause_exec is not None and isinstance(cause_exec, httpx.HTTPStatusError):
status_codes.append(cause_exec.response.status_code)
warnings.warn(
f"HTTP Exception for {cause_exec.request.url} - {cause_exec}",
)
else:
warnings.warn(f"Unexpected error: {e}")
except Exception as e: # pragma: no cover
warnings.warn(f"Unexpected error: {e}")
retry_count = retry_count + 1
time.sleep(retry_time)
retry_time = retry_time * 2.0
if (
not succeeded
and status_codes
and skip_on_50x_err
and all(httpx.codes.is_server_error(code) for code in status_codes)
):
pytest.skip("Repeated HTTP 50x for service") # pragma: no cover
return succeeded, result
def read_streaming_response(response: StreamingHttpResponse) -> bytes:
"""Consume a StreamingHttpResponse/FileResponse and close it."""
content = b"".join(response.streaming_content)
response.close()
return content
class FileSystemAssertsMixin:
"""
Utilities for checks various state information of the file system
"""
def assertIsFile(self, path: PathLike[str] | str) -> None:
self.assertTrue(Path(path).resolve().is_file(), f"File does not exist: {path}")
def assertIsNotFile(self, path: PathLike[str] | str) -> None:
self.assertFalse(Path(path).resolve().is_file(), f"File does exist: {path}")
def assertIsDir(self, path: PathLike[str] | str) -> None:
self.assertTrue(Path(path).resolve().is_dir(), f"Dir does not exist: {path}")
def assertIsNotDir(self, path: PathLike[str] | str) -> None:
self.assertFalse(Path(path).resolve().is_dir(), f"Dir does exist: {path}")
def assertFileCountInDir(self, path: PathLike[str] | str, count: int) -> None:
path = Path(path).resolve()
self.assertTrue(path.is_dir(), f"Path {path} is not a directory")
files = [x for x in path.iterdir() if x.is_file()]
self.assertEqual(
len(files),
count,
f"Path {path} contains {len(files)} files instead of {count} files",
)
from paperless_testing.fakes.progress import FakeProgressManager
class ConsumeTaskMixin:
@@ -158,59 +48,6 @@ class ConsumeTaskMixin:
yield (task_kwargs["input_doc"], task_kwargs["overrides"])
class TestMigrations(TransactionTestCase):
@property
def app(self):
return apps.get_containing_app_config(type(self).__module__).name
migrate_from = None
dependencies = None
migrate_to = None
def setUp(self) -> None:
super().setUp()
assert self.migrate_from and self.migrate_to, (
f"TestCase '{type(self).__name__}' must define migrate_from and migrate_to properties"
)
self.migrate_from = [(self.app, self.migrate_from)]
if self.dependencies is not None:
self.migrate_from.extend(self.dependencies)
self.migrate_to = [(self.app, self.migrate_to)]
executor = MigrationExecutor(connection)
old_apps = executor.loader.project_state(self.migrate_from).apps
# Reverse to the original migration
executor.migrate(self.migrate_from)
self.setUpBeforeMigration(old_apps)
self.apps = old_apps
# Run the migration to test
executor = MigrationExecutor(connection)
executor.loader.build_graph() # reload.
executor.migrate(self.migrate_to)
self.apps = executor.loader.project_state(self.migrate_to).apps
def setUpBeforeMigration(self, apps) -> None:
pass
def tearDown(self) -> None:
"""
Ensure the database schema is restored to the latest migration after
each migration test, so subsequent tests run against HEAD.
"""
try:
executor = MigrationExecutor(connection)
executor.loader.build_graph()
targets = executor.loader.graph.leaf_nodes()
executor.migrate(targets)
finally:
super().tearDown()
class SampleDirMixin:
SAMPLE_DIR = Path(__file__).parent / "samples"
@@ -227,7 +64,7 @@ class GetConsumerMixin:
mailrule_id: int | None = None,
) -> Generator[ConsumerPlugin, None, None]:
# Store this for verification
self.status = DummyProgressManager(filepath.name, None)
self.status = FakeProgressManager(filepath.name, None)
doc = ConsumableDocument(
source,
original_file=filepath,
@@ -263,63 +100,3 @@ class GetConsumerMixin:
yield reader
finally:
reader.cleanup()
class DummyProgressManager:
"""
A dummy handler for progress management that doesn't actually try to
connect to Redis. Payloads are stored for test assertions if needed.
Use it with
mock.patch("documents.tasks.ProgressManager", DummyProgressManager)
"""
def __init__(self, filename: str, task_id: str | None = None) -> None:
self.filename = filename
self.task_id = task_id
self.payloads = []
def __enter__(self):
self.open()
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
self.close()
def open(self) -> None:
pass
def close(self) -> None:
pass
def send_progress(
self,
status: ProgressStatusOptions,
message: str,
current_progress: int,
max_progress: int,
*,
document_id: int | None = None,
owner_id: int | None = None,
users_can_view: list[int] | None = None,
groups_can_view: list[int] | None = None,
) -> None:
# Ensure the layer is open
self.open()
payload = {
"type": "status_update",
"data": {
"filename": self.filename,
"task_id": self.task_id,
"current_progress": current_progress,
"max_progress": max_progress,
"status": status,
"message": message,
"document_id": document_id,
"owner_id": owner_id,
"users_can_view": users_can_view or [],
"groups_can_view": groups_can_view or [],
},
}
self.payloads.append(payload)
+18
View File
@@ -256,6 +256,7 @@ from paperless.views import StandardPagination
from paperless_ai.ai_classifier import get_ai_document_classification
from paperless_ai.ai_classifier import get_llm_output_language
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 LLMTimeoutError
from paperless_ai.matching import extract_unmatched_names
@@ -1697,6 +1698,23 @@ class DocumentViewSet(
},
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(
doc.pk,
llm_suggestions,
+6 -4
View File
@@ -4,7 +4,8 @@ import httpx
from celery import shared_task
from django.conf import settings
from paperless.network import PinnedHostHTTPTransport
from paperless.network import GuardedHTTPTransport
from paperless.network import OutboundRequestBlockedError
from paperless.network import validate_outbound_http_url
logger = logging.getLogger("paperless.workflows.webhooks")
@@ -14,7 +15,7 @@ logger = logging.getLogger("paperless.workflows.webhooks")
retry_backoff=True,
autoretry_for=(httpx.HTTPStatusError,),
max_retries=3,
throws=(httpx.HTTPError,),
throws=(httpx.HTTPError, OutboundRequestBlockedError),
)
def send_webhook(
url: str,
@@ -29,14 +30,15 @@ def send_webhook(
url,
allowed_schemes=settings.WEBHOOKS_ALLOWED_SCHEMES,
allowed_ports=settings.WEBHOOKS_ALLOWED_PORTS,
# Internal-address checks happen in transport to preserve ConnectError behavior.
# Scheme and port only; the transport enforces the internal-address
# policy at connect time, on the address actually dialled.
allow_internal=True,
)
except ValueError as e:
logger.warning("Webhook blocked: %s", e)
raise
transport = PinnedHostHTTPTransport(
transport = GuardedHTTPTransport(
allow_internal=settings.WEBHOOKS_ALLOW_INTERNAL_REQUESTS,
)
+10 -10
View File
@@ -2,7 +2,7 @@ msgid ""
msgstr ""
"Project-Id-Version: paperless-ngx\n"
"Report-Msgid-Bugs-To: \n"
"POT-Creation-Date: 2026-09-18 01:29+0000\n"
"POT-Creation-Date: 2026-09-21 19:00+0000\n"
"PO-Revision-Date: 2022-02-17 04:17\n"
"Last-Translator: \n"
"Language-Team: English\n"
@@ -1632,7 +1632,7 @@ msgid "workflow runs"
msgstr ""
#: documents/serialisers.py:514 documents/serialisers.py:871
#: documents/serialisers.py:2883 documents/views.py:343 documents/views.py:2726
#: documents/serialisers.py:2885 documents/views.py:343 documents/views.py:2726
#: paperless_mail/serialisers.py:156
msgid "Insufficient permissions."
msgstr ""
@@ -1641,39 +1641,39 @@ msgstr ""
msgid "Invalid color."
msgstr ""
#: documents/serialisers.py:2350
#: documents/serialisers.py:2352
#, python-format
msgid "File type %(type)s not supported"
msgstr ""
#: documents/serialisers.py:2394
#: documents/serialisers.py:2396
#, python-format
msgid "Custom field id must be an integer: %(id)s"
msgstr ""
#: documents/serialisers.py:2401
#: documents/serialisers.py:2403
#, python-format
msgid "Custom field with id %(id)s does not exist"
msgstr ""
#: documents/serialisers.py:2418 documents/serialisers.py:2428
#: documents/serialisers.py:2420 documents/serialisers.py:2430
msgid ""
"Custom fields must be a list of integers or an object mapping ids to values."
msgstr ""
#: documents/serialisers.py:2423
#: documents/serialisers.py:2425
msgid "Some custom fields don't exist or were specified twice."
msgstr ""
#: documents/serialisers.py:2570
#: documents/serialisers.py:2572
msgid "Invalid variable detected."
msgstr ""
#: documents/serialisers.py:2939
#: documents/serialisers.py:2941
msgid "Duplicate document identifiers are not allowed."
msgstr ""
#: documents/serialisers.py:2969 documents/views.py:4763
#: documents/serialisers.py:2971 documents/views.py:4763
#, python-format
msgid "Documents not found: %(ids)s"
msgstr ""
+519 -158
View File
@@ -1,61 +1,533 @@
import functools
import ipaddress
import logging
import math
import re
import socket
import time
from collections.abc import Callable
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 urlparse
import anyio
import httpcore
import httpx
# Ranges ipaddress does not report as private, but which routinely front
# internal infrastructure.
# Not exported by httpcore; the guard asserts it is still the async default.
from httpcore._backends.auto import AutoBackend
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 = (
# 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
# a NAT64 gateway exists.
# a NAT64 gateway exists, yet ipaddress classifies the prefix as global.
ipaddress.ip_network("64:ff9b::/96"),
)
def is_public_ip(ip: str | int) -> bool:
try:
obj = ipaddress.ip_address(ip)
return not (
obj.is_private
or obj.is_loopback
or obj.is_link_local
or obj.is_multicast
or obj.is_unspecified
or any(obj in network for network in _NON_PUBLIC_NETWORKS)
class BlockReason(StrEnum):
NON_PUBLIC_ADDRESS = "non_public_address"
UNIX_SOCKET = "unix_socket"
class OutboundRequestBlockedError(Exception):
"""
An outbound connection was refused by policy before any socket was opened.
For NON_PUBLIC_ADDRESS, ``host`` is the name or literal being connected to
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
def resolve_hostname_ips(hostname: str) -> list[str]:
try:
addr_info = socket.getaddrinfo(hostname, None)
except socket.gaierror as e:
raise ValueError(f"Could not resolve hostname: {hostname}") from e
class HostResolutionError(Exception):
"""The resolver returned no usable addresses for a host."""
ips = [info[4][0] for info in addr_info if info and info[4]]
if not ips:
raise ValueError(f"Could not resolve hostname: {hostname}")
return ips
def __init__(self, *, host: str, detail: str) -> None:
self.host = host
self.detail = detail
super().__init__(f"Could not resolve {host}: {detail}")
def __reduce__(self) -> tuple[Callable[..., Self], tuple[object, ...]]:
return (
functools.partial(type(self), host=self.host, detail=self.detail),
(),
)
def format_host_for_url(host: str) -> str:
def blocked_message(exc: OutboundRequestBlockedError | HostResolutionError) -> 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:
"""
Format IP address for URL use (wrap IPv6 in brackets).
True when ``ip`` is globally routable unicast and not in a range that
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:
ip_obj = ipaddress.ip_address(host)
if ip_obj.version == 6:
return f"[{host}]"
return host
except ValueError:
return host
infos = _getaddrinfo(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))
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(
@@ -81,128 +553,17 @@ def validate_outbound_http_url(
raise ValueError("Destination port not permitted.")
if not allow_internal:
for ip_str in resolve_hostname_ips(parsed.hostname):
if not is_public_ip(ip_str):
raise ValueError(
f"Connection blocked: {parsed.hostname} resolves to a non-public address",
)
if _UNSAFE_URL_CHARS.search(url):
raise ValueError("Invalid URL scheme or hostname.")
host = _dns_name(url)
# HTTP clients may percent-decode the host before resolving it, so the
# 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
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,
)
+4 -4
View File
@@ -22,11 +22,11 @@ if TYPE_CHECKING:
@pytest.fixture(scope="session")
def samples_dir() -> Path:
def parser_samples_dir() -> Path:
"""Absolute path to the shared parser sample files directory.
Sub-package conftest files derive format-specific paths from this root,
e.g. ``samples_dir / "text" / "test.txt"``.
e.g. ``parser_samples_dir / "text" / "test.txt"``.
Returns
-------
@@ -37,7 +37,7 @@ def samples_dir() -> Path:
@pytest.fixture(scope="session")
def tagged_no_text_pdf_file(samples_dir: Path) -> Path:
def tagged_no_text_pdf_file(parser_samples_dir: Path) -> Path:
"""Path to a tagged PDF whose only "text" is pdftotext layout padding.
Reproduces GH #13387: ``/MarkInfo /Marked true`` is set, but the only
@@ -50,7 +50,7 @@ def tagged_no_text_pdf_file(samples_dir: Path) -> Path:
Path
Absolute path to ``tesseract/tagged-but-no-text.pdf``.
"""
return samples_dir / "tesseract" / "tagged-but-no-text.pdf"
return parser_samples_dir / "tesseract" / "tagged-but-no-text.pdf"
@pytest.fixture(autouse=True)
+12 -240
View File
@@ -37,15 +37,15 @@ if TYPE_CHECKING:
@pytest.fixture(scope="session")
def text_samples_dir(samples_dir: Path) -> Path:
def text_samples_dir(parser_samples_dir: Path) -> Path:
"""Absolute path to the text parser sample files directory.
Returns
-------
Path
``<samples_dir>/text/``
``<parser_samples_dir>/text/``
"""
return samples_dir / "text"
return parser_samples_dir / "text"
@pytest.fixture(scope="session")
@@ -175,15 +175,15 @@ def no_engine_settings(
@pytest.fixture(scope="session")
def tika_samples_dir(samples_dir: Path) -> Path:
def tika_samples_dir(parser_samples_dir: Path) -> Path:
"""Absolute path to the Tika parser sample files directory.
Returns
-------
Path
``<samples_dir>/tika/``
``<parser_samples_dir>/tika/``
"""
return samples_dir / "tika"
return parser_samples_dir / "tika"
@pytest.fixture(scope="session")
@@ -258,15 +258,15 @@ def tika_parser() -> Generator[TikaDocumentParser, None, None]:
@pytest.fixture(scope="session")
def mail_samples_dir(samples_dir: Path) -> Path:
def mail_samples_dir(parser_samples_dir: Path) -> Path:
"""Absolute path to the mail parser sample files directory.
Returns
-------
Path
``<samples_dir>/mail/``
``<parser_samples_dir>/mail/``
"""
return samples_dir / "mail"
return parser_samples_dir / "mail"
@pytest.fixture(scope="session")
@@ -421,75 +421,15 @@ def nginx_base_url() -> Generator[str, None, None]:
@pytest.fixture(scope="session")
def tesseract_samples_dir(samples_dir: Path) -> Path:
def tesseract_samples_dir(parser_samples_dir: Path) -> Path:
"""Absolute path to the tesseract parser sample files directory.
Returns
-------
Path
``<samples_dir>/tesseract/``
``<parser_samples_dir>/tesseract/``
"""
return samples_dir / "tesseract"
@pytest.fixture(scope="session")
def document_webp_file(tesseract_samples_dir: Path) -> Path:
"""Path to a WebP document sample file.
Returns
-------
Path
Absolute path to ``tesseract/document.webp``.
"""
return tesseract_samples_dir / "document.webp"
@pytest.fixture(scope="session")
def encrypted_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to an encrypted PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/encrypted.pdf``.
"""
return tesseract_samples_dir / "encrypted.pdf"
@pytest.fixture(scope="session")
def multi_page_digital_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a multi-page digital PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/multi-page-digital.pdf``.
"""
return tesseract_samples_dir / "multi-page-digital.pdf"
@pytest.fixture(scope="session")
def multi_page_images_alpha_rgb_tiff_file(tesseract_samples_dir: Path) -> Path:
"""Path to a multi-page TIFF with alpha channel in RGB.
Returns
-------
Path
Absolute path to ``tesseract/multi-page-images-alpha-rgb.tiff``.
"""
return tesseract_samples_dir / "multi-page-images-alpha-rgb.tiff"
@pytest.fixture(scope="session")
def multi_page_images_alpha_tiff_file(tesseract_samples_dir: Path) -> Path:
"""Path to a multi-page TIFF with alpha channel.
Returns
-------
Path
Absolute path to ``tesseract/multi-page-images-alpha.tiff``.
"""
return tesseract_samples_dir / "multi-page-images-alpha.tiff"
return parser_samples_dir / "tesseract"
@pytest.fixture(scope="session")
@@ -504,90 +444,6 @@ def multi_page_images_pdf_file(tesseract_samples_dir: Path) -> Path:
return tesseract_samples_dir / "multi-page-images.pdf"
@pytest.fixture(scope="session")
def multi_page_images_tiff_file(tesseract_samples_dir: Path) -> Path:
"""Path to a multi-page TIFF sample file.
Returns
-------
Path
Absolute path to ``tesseract/multi-page-images.tiff``.
"""
return tesseract_samples_dir / "multi-page-images.tiff"
@pytest.fixture(scope="session")
def multi_page_mixed_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a multi-page mixed PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/multi-page-mixed.pdf``.
"""
return tesseract_samples_dir / "multi-page-mixed.pdf"
@pytest.fixture(scope="session")
def no_text_alpha_png_file(tesseract_samples_dir: Path) -> Path:
"""Path to a PNG with alpha channel and no text.
Returns
-------
Path
Absolute path to ``tesseract/no-text-alpha.png``.
"""
return tesseract_samples_dir / "no-text-alpha.png"
@pytest.fixture(scope="session")
def rotated_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a rotated PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/rotated.pdf``.
"""
return tesseract_samples_dir / "rotated.pdf"
@pytest.fixture(scope="session")
def rtl_test_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to an RTL test PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/rtl-test.pdf``.
"""
return tesseract_samples_dir / "rtl-test.pdf"
@pytest.fixture(scope="session")
def signed_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a signed PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/signed.pdf``.
"""
return tesseract_samples_dir / "signed.pdf"
@pytest.fixture(scope="session")
def simple_alpha_png_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple PNG with alpha channel.
Returns
-------
Path
Absolute path to ``tesseract/simple-alpha.png``.
"""
return tesseract_samples_dir / "simple-alpha.png"
@pytest.fixture(scope="session")
def simple_digital_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple digital PDF sample file.
@@ -612,54 +468,6 @@ def simple_no_dpi_png_file(tesseract_samples_dir: Path) -> Path:
return tesseract_samples_dir / "simple-no-dpi.png"
@pytest.fixture(scope="session")
def simple_bmp_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple BMP sample file.
Returns
-------
Path
Absolute path to ``tesseract/simple.bmp``.
"""
return tesseract_samples_dir / "simple.bmp"
@pytest.fixture(scope="session")
def simple_gif_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple GIF sample file.
Returns
-------
Path
Absolute path to ``tesseract/simple.gif``.
"""
return tesseract_samples_dir / "simple.gif"
@pytest.fixture(scope="session")
def simple_heic_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple HEIC sample file.
Returns
-------
Path
Absolute path to ``tesseract/simple.heic``.
"""
return tesseract_samples_dir / "simple.heic"
@pytest.fixture(scope="session")
def simple_jpg_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple JPG sample file.
Returns
-------
Path
Absolute path to ``tesseract/simple.jpg``.
"""
return tesseract_samples_dir / "simple.jpg"
@pytest.fixture(scope="session")
def simple_png_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple PNG sample file.
@@ -672,42 +480,6 @@ def simple_png_file(tesseract_samples_dir: Path) -> Path:
return tesseract_samples_dir / "simple.png"
@pytest.fixture(scope="session")
def simple_tif_file(tesseract_samples_dir: Path) -> Path:
"""Path to a simple TIF sample file.
Returns
-------
Path
Absolute path to ``tesseract/simple.tif``.
"""
return tesseract_samples_dir / "simple.tif"
@pytest.fixture(scope="session")
def single_page_mixed_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a single-page mixed PDF sample file.
Returns
-------
Path
Absolute path to ``tesseract/single-page-mixed.pdf``.
"""
return tesseract_samples_dir / "single-page-mixed.pdf"
@pytest.fixture(scope="session")
def with_form_pdf_file(tesseract_samples_dir: Path) -> Path:
"""Path to a PDF with form sample file.
Returns
-------
Path
Absolute path to ``tesseract/with-form.pdf``.
"""
return tesseract_samples_dir / "with-form.pdf"
# ------------------------------------------------------------------
# Tesseract parser instance and settings helpers
# ------------------------------------------------------------------
@@ -10,8 +10,8 @@ from imagehash import average_hash
from PIL import Image
from pytest_mock import MockerFixture
from documents.tests.utils import util_call_with_backoff
from paperless.parsers.mail import MailDocumentParser
from paperless_testing.retry import util_call_with_backoff
def extract_text(pdf_path: Path) -> str:
@@ -3,13 +3,13 @@ import json
from django.test import TestCase
from django.test import override_settings
from documents.tests.utils import FileSystemAssertsMixin
from paperless.models import ApplicationConfiguration
from paperless.models import CleanChoices
from paperless.models import ColorConvertChoices
from paperless.models import ModeChoices
from paperless.models import OutputTypeChoices
from paperless.parsers.tesseract import RasterisedDocumentParser
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
@@ -3,8 +3,8 @@ from pathlib import Path
import pytest
from documents.tests.utils import util_call_with_backoff
from paperless.parsers.tika import TikaDocumentParser
from paperless_testing.retry import util_call_with_backoff
@pytest.mark.skipif(
@@ -1,4 +1,4 @@
from documents.tests.utils import TestMigrations
from paperless_testing.migrations import TestMigrations
class TestMigrateSkipArchiveFile(TestMigrations):
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,375 @@
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() == []
+29 -8
View File
@@ -14,14 +14,16 @@ if TYPE_CHECKING:
from llama_index.llms.openai_like import OpenAILike
from paperless.config import AIConfig
from paperless.network import PinnedHostAsyncHTTPTransport
from paperless.network import PinnedHostHTTPTransport
from paperless.network import create_pinned_async_httpx_client
from paperless.network import create_pinned_httpx_client
from paperless.network import GuardedAsyncHTTPTransport
from paperless.network import GuardedHTTPTransport
from paperless.network import OutboundRequestBlockedError
from paperless.network import create_guarded_async_httpx_client
from paperless.network import create_guarded_httpx_client
from paperless.network import validate_outbound_http_url
from paperless_ai.base_model import ClassificationSuggestions
from paperless_ai.base_model import DocumentClassifierSchema
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 LLMTimeoutError
@@ -43,6 +45,19 @@ LLM_SYSTEM_PROMPT = (
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:
"""
A client for interacting with an LLM backend.
@@ -63,10 +78,10 @@ class AIClient:
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
)
transport = PinnedHostHTTPTransport(
transport = GuardedHTTPTransport(
allow_internal=self.settings.llm_allow_internal_endpoints,
)
async_transport = PinnedHostAsyncHTTPTransport(
async_transport = GuardedAsyncHTTPTransport(
allow_internal=self.settings.llm_allow_internal_endpoints,
)
return Ollama(
@@ -93,12 +108,12 @@ class AIClient:
http_client = None
async_http_client = None
if endpoint:
http_client = create_pinned_httpx_client(
http_client = create_guarded_httpx_client(
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
timeout=self.settings.llm_request_timeout,
)
async_http_client = create_pinned_async_httpx_client(
async_http_client = create_guarded_async_httpx_client(
endpoint,
allow_internal=self.settings.llm_allow_internal_endpoints,
timeout=self.settings.llm_request_timeout,
@@ -179,6 +194,12 @@ class AIClient:
except httpx.TimeoutException as exc:
raise LLMTimeoutError from 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):
raise LLMTimeoutError from exc
if self._is_provider_error(exc):
+8 -8
View File
@@ -9,10 +9,10 @@ if TYPE_CHECKING:
from documents.models import Document
from paperless.config import AIConfig
from paperless.models import LLMEmbeddingBackend
from paperless.network import PinnedHostAsyncHTTPTransport
from paperless.network import PinnedHostHTTPTransport
from paperless.network import create_pinned_async_httpx_client
from paperless.network import create_pinned_httpx_client
from paperless.network import GuardedAsyncHTTPTransport
from paperless.network import GuardedHTTPTransport
from paperless.network import create_guarded_async_httpx_client
from paperless.network import create_guarded_httpx_client
from paperless.network import validate_outbound_http_url
from paperless_ai.client import PLACEHOLDER_API_KEY
@@ -29,12 +29,12 @@ def get_embedding_model(config: AIConfig) -> "BaseEmbedding":
http_client = None
async_http_client = None
if endpoint:
http_client = create_pinned_httpx_client(
http_client = create_guarded_httpx_client(
endpoint,
allow_internal=config.llm_allow_internal_endpoints,
timeout=config.llm_request_timeout,
)
async_http_client = create_pinned_async_httpx_client(
async_http_client = create_guarded_async_httpx_client(
endpoint,
allow_internal=config.llm_allow_internal_endpoints,
timeout=config.llm_request_timeout,
@@ -77,14 +77,14 @@ def get_embedding_model(config: AIConfig) -> "BaseEmbedding":
embedding._client = Client(
host=endpoint,
timeout=config.llm_request_timeout,
transport=PinnedHostHTTPTransport(
transport=GuardedHTTPTransport(
allow_internal=config.llm_allow_internal_endpoints,
),
)
embedding._async_client = AsyncClient(
host=endpoint,
timeout=config.llm_request_timeout,
transport=PinnedHostAsyncHTTPTransport(
transport=GuardedAsyncHTTPTransport(
allow_internal=config.llm_allow_internal_endpoints,
),
)
+4
View File
@@ -4,3 +4,7 @@ class LLMTimeoutError(Exception):
class LLMProviderError(Exception):
"""The LLM backend rejected the request."""
class LLMBlockedError(Exception):
"""The outbound request policy refused the connection to the LLM backend."""
+7 -2
View File
@@ -1,6 +1,7 @@
import datetime
from collections.abc import Generator
from types import SimpleNamespace
from typing import TYPE_CHECKING
from unittest.mock import MagicMock
from unittest.mock import patch
@@ -28,6 +29,9 @@ from paperless_testing.factories import TagFactory
from paperless_testing.factories import UserFactory
from paperless_testing.permissions import grant_object
if TYPE_CHECKING:
from paperless_testing.dirs import PaperlessDirs
@pytest.fixture
def mock_document():
@@ -630,10 +634,11 @@ class TestFulltextSimilarDocuments:
def fulltext_backend(
self,
mocker: pytest_mock.MockerFixture,
paperless_dirs: "PaperlessDirs",
) -> Generator[TantivyBackend, None, None]:
"""An in-memory Tantivy backend, wired up as the module-level
"""An on-disk Tantivy backend, wired up as the module-level
singleton _fulltext_similar_documents resolves via get_backend()."""
backend = TantivyBackend(path=None)
backend = TantivyBackend(path=paperless_dirs.index_dir)
backend.open()
mocker.patch("documents.search.get_backend", return_value=backend)
try:
+144
View File
@@ -1,3 +1,4 @@
import ipaddress
import json
from unittest.mock import ANY
from unittest.mock import MagicMock
@@ -9,11 +10,15 @@ import openai
import pytest
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 PLACEHOLDER_API_KEY
from paperless_ai.client import AIClient
from paperless_ai.exceptions import LLMBlockedError
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError
from paperless_testing.outbound import guard_of
@pytest.fixture
@@ -277,3 +282,142 @@ def test_run_llm_query_httpx_timeout_raises_local_error(
with pytest.raises(LLMTimeoutError):
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")
+65
View File
@@ -1,9 +1,12 @@
from typing import TYPE_CHECKING
from typing import cast
from unittest.mock import ANY
from unittest.mock import MagicMock
from unittest.mock import patch
import pytest
from django.conf import settings
from pytest_mock import MockerFixture
from documents.models import Document
from paperless.models import LLMEmbeddingBackend
@@ -12,6 +15,10 @@ from paperless_ai.embedding import _normalize_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_embedding_model
from paperless_testing.outbound import guard_of
if TYPE_CHECKING:
from llama_index.embeddings.ollama import OllamaEmbedding
@pytest.fixture
@@ -283,3 +290,61 @@ def test_normalize_llm_index_text_collapses_ocr_leaders_without_joining_lines():
def test_normalize_llm_index_text_collapses_non_breaking_spaces():
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
+43 -21
View File
@@ -45,8 +45,11 @@ from documents.models import Correspondent
from documents.models import PaperlessTask
from documents.parsers import is_mime_type_supported
from documents.tasks import consume_file
from paperless.network import is_public_ip
from paperless.network import resolve_hostname_ips
from paperless.network import HostResolutionError
from paperless.network import IPAddress
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 MailRule
from paperless_mail.models import ProcessedMail
@@ -445,18 +448,34 @@ class PinnedIMAP4(imaplib.IMAP4):
Without pinned addresses, and with the ssl_context of the matching imaplib
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__(self, host, port, pinned_ips, ssl_context=None, timeout=None) -> None:
def __init__(
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.ssl_context = ssl_context
super().__init__(host, port, timeout=timeout)
def _connect_pinned(self, timeout):
def _connect_pinned(
self,
pinned_ips: tuple[IPAddress, ...],
timeout: float | None,
) -> socket.socket:
last_error: OSError | None = None
for ip_str in self._pinned_ips:
for ip in pinned_ips:
try:
address = (ip_str, self.port)
address = (str(ip), self.port)
if timeout is not None:
return socket.create_connection(address, timeout)
return socket.create_connection(address)
@@ -464,9 +483,9 @@ class PinnedIMAP4(imaplib.IMAP4):
last_error = e
raise last_error or OSError(f"Could not connect to {self.host}")
def _create_socket(self, timeout):
if self._pinned_ips:
sock = self._connect_pinned(timeout)
def _create_socket(self, timeout: float | None) -> socket.socket:
if self._pinned_ips is not None:
sock = self._connect_pinned(self._pinned_ips, timeout)
else:
sock = super()._create_socket(timeout)
if self.ssl_context is None:
@@ -477,7 +496,12 @@ class PinnedIMAP4(imaplib.IMAP4):
class PinnedClientMixin:
"""Builds the imaplib client against the pre-resolved addresses, if any."""
def __init__(self, *args, pinned_ips: list[str] | None, **kwargs) -> None:
def __init__(
self,
*args,
pinned_ips: tuple[IPAddress, ...] | None,
**kwargs,
) -> None:
self._pinned_ips = pinned_ips
super().__init__(*args, **kwargs)
@@ -515,22 +539,20 @@ class PinnedMailBoxStartTls(PinnedClientMixin, MailBoxStartTls):
return client
def get_mailbox(server, port, security) -> MailBox:
def get_mailbox(
server: str,
port: int | None,
security: int,
) -> MailBox:
"""
Returns the correct MailBox instance for the given configuration.
"""
pinned_ips: list[str] | None = None
pinned_ips: tuple[IPAddress, ...] | None = None
if not settings.EMAIL_ALLOW_INTERNAL_HOSTS:
try:
pinned_ips = resolve_hostname_ips(server)
except ValueError as 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",
)
pinned_ips = resolve_public_addresses(server, port)
except (OutboundRequestBlockedError, HostResolutionError) as e:
raise MailError(blocked_message(e)) from e
ssl_context = ssl.create_default_context()
if settings.EMAIL_CERTIFICATE_FILE is not None: # pragma: no cover
+262
View File
@@ -0,0 +1,262 @@
import dataclasses
import email.message
import uuid
from contextlib import AbstractContextManager
from imap_tools import MailboxFolderSelectError
from imap_tools import MailboxLoginError
from imap_tools import MailMessage
from imap_tools import MailMessageFlags
@dataclasses.dataclass
class _AttachmentDef:
filename: str = "a_file.pdf"
maintype: str = "application/pdf"
subtype: str = "pdf"
disposition: str = "attachment"
content: bytes = b"a PDF document"
class BogusFolderManager:
current_folder = "INBOX"
uidvalidity = "1"
def set(self, new_folder) -> None:
if new_folder not in ["INBOX", "spam"]:
raise MailboxFolderSelectError(None, "uhm")
self.current_folder = new_folder
def status(self, folder, options):
return {"UIDVALIDITY": self.uidvalidity}
class BogusClient:
def __init__(self, messages) -> None:
self.messages: list[MailMessage] = messages
self.capabilities: list[str] = []
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
pass
def authenticate(self, mechanism, authobject) -> None:
# authobject must be a callable object
auth_bytes = authobject(None)
if auth_bytes != b"\x00admin\x00w57\xc3\xa4\xc3\xb6\xc3\xbcw4b6huwb6nhu":
raise MailboxLoginError("BAD", "OK")
def uid(self, command, *args) -> None:
if command == "STORE":
for message in self.messages:
if message.uid == args[0]:
flag = args[2]
if flag == "processed":
message._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(message, "flags"):
del message.flags
class BogusMailBox(AbstractContextManager):
# Common values so tests don't need to remember an accepted login
USERNAME: str = "admin"
ASCII_PASSWORD: str = "secret"
# Note the non-ascii characters here
UTF_PASSWORD: str = "w57äöüw4b6huwb6nhu"
# A dummy access token
ACCESS_TOKEN = "ea7e075cd3acf2c54c48e600398d5d5a"
def __init__(self) -> None:
self.messages: list[MailMessage] = []
self.messages_spam: list[MailMessage] = []
self.folder = BogusFolderManager()
self.client = BogusClient(self.messages)
self._host = ""
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
pass
def updateClient(self) -> None:
self.client = BogusClient(self.messages)
def login(self, username, password) -> None:
# This will raise a UnicodeEncodeError if the password is not ASCII only
password.encode("ascii")
# Otherwise, check for correct values
if username != self.USERNAME or password != self.ASCII_PASSWORD:
raise MailboxLoginError("BAD", "OK")
def login_utf8(self, username, password) -> None:
# Expected to only be called with the UTF-8 password
if username != self.USERNAME or password != self.UTF_PASSWORD:
raise MailboxLoginError("BAD", "OK")
def xoauth2(self, username: str, access_token: str) -> None:
if username != self.USERNAME or access_token != self.ACCESS_TOKEN:
raise MailboxLoginError("BAD", "OK")
def fetch(
self,
criteria="ALL",
charset="",
*,
mark_seen=True,
bulk=True,
uid_list=None,
):
if uid_list is not None:
return [m for m in self.messages if m.uid in uid_list]
return self._filter_messages(criteria)
def uids(self, criteria, charset="") -> list[str]:
return [m.uid for m in self._filter_messages(criteria)]
def _filter_messages(self, criteria):
msg = self.messages
criteria = str(criteria).strip("()").split(" ")
if "UNSEEN" in criteria:
msg = filter(lambda m: not m.seen, msg)
if "SUBJECT" in criteria:
subject = criteria[criteria.index("SUBJECT") + 1].strip('"')
msg = filter(lambda m: subject in m.subject, msg)
if "BODY" in criteria:
body = criteria[criteria.index("BODY") + 1].strip('"')
msg = filter(lambda m: body in m.text, msg)
if "FROM" in criteria:
from_ = criteria[criteria.index("FROM") + 1].strip('"')
msg = filter(lambda m: from_ in m.from_, msg)
if "TO" in criteria:
to_ = criteria[criteria.index("TO") + 1].strip('"')
msg = filter(lambda m: any(to_ in to_addr for to_addr in m.to), msg)
if "UNFLAGGED" in criteria:
msg = filter(lambda m: not m.flagged, msg)
if "UNKEYWORD" in criteria:
tag = criteria[criteria.index("UNKEYWORD") + 1].strip("'")
msg = filter(lambda m: tag not in m.flags, msg)
if "(X-GM-LABELS" in criteria: # ['NOT', '(X-GM-LABELS', '"processed"']
msg = filter(lambda m: "processed" not in m.flags, msg)
if "UID" in criteria:
uid_list = criteria[criteria.index("UID") + 1].split(",")
msg = filter(lambda m: m.uid in uid_list, msg)
return list(msg)
def delete(self, uid_list) -> None:
self.messages = list(filter(lambda m: m.uid not in uid_list, self.messages))
def flag(self, uid_list, flag_set, value) -> None:
for message in self.messages:
if message.uid in uid_list:
for flag in flag_set:
if flag == MailMessageFlags.FLAGGED:
message.flagged = value
if flag == MailMessageFlags.SEEN:
message.seen = value
if flag == "processed":
message._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(message, "flags"):
del message.flags
def move(self, uid_list, folder) -> None:
if folder == "spam":
self.messages_spam += list(
filter(lambda m: m.uid in uid_list, self.messages),
)
self.messages = list(filter(lambda m: m.uid not in uid_list, self.messages))
else:
raise Exception
def fake_magic_from_buffer(buffer, *, mime=False):
if mime:
if "PDF" in str(buffer):
return "application/pdf"
else:
return "unknown/type"
else:
return "Some verbose file description"
class MessageBuilder:
def __init__(self) -> None:
self._next_uid = 1
def create_message(
self,
*,
attachments: int | list[_AttachmentDef] = 1,
body: str = "",
subject: str = "the subject",
from_: str = "no_one@mail.com",
to: list[str] | None = None,
seen: bool = False,
flagged: bool = False,
processed: bool = False,
) -> MailMessage:
if to is None:
to = ["tosomeone@somewhere.com"]
email_msg = email.message.EmailMessage()
# TODO: This does NOT set the UID
email_msg["Message-ID"] = str(uuid.uuid4())
email_msg["Subject"] = subject
email_msg["From"] = from_
email_msg["To"] = str(" ,".join(to))
email_msg.set_content(body)
# Either add some default number of attachments
# or the provided attachments
if isinstance(attachments, int):
for i in range(attachments):
attachment = _AttachmentDef(filename=f"file_{i}.pdf")
email_msg.add_attachment(
attachment.content,
maintype=attachment.maintype,
subtype=attachment.subtype,
disposition=attachment.disposition,
filename=attachment.filename,
)
else:
for attachment in attachments:
email_msg.add_attachment(
attachment.content,
maintype=attachment.maintype,
subtype=attachment.subtype,
disposition=attachment.disposition,
filename=attachment.filename,
)
# Convert the EmailMessage to an imap_tools MailMessage
imap_msg = MailMessage.from_bytes(email_msg.as_bytes())
# TODO: Unsure how to add a uid to the actual EmailMessage. This hacks it in,
# based on how imap_tools uses regex to extract it.
# This should be a large enough pool
uid = self._next_uid
self._next_uid += 1
imap_msg._raw_uid_data = f"UID {uid}".encode()
imap_msg.seen = seen
imap_msg.flagged = flagged
if processed:
imap_msg._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(imap_msg, "flags"):
del imap_msg.flags
return imap_msg
+1 -1
View File
@@ -11,7 +11,7 @@ from paperless_mail.models import ProcessedMail
from paperless_mail.tests.factories import MailAccountFactory
from paperless_mail.tests.factories import MailRuleFactory
from paperless_mail.tests.factories import ProcessedMailFactory
from paperless_mail.tests.test_mail import BogusMailBox
from paperless_mail.tests.helpers import BogusMailBox
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import CorrespondentFactory
from paperless_testing.factories import DocumentTypeFactory
+65 -271
View File
@@ -1,11 +1,12 @@
import dataclasses
import email.contentmanager
import ipaddress
import socket
import time
import uuid
from collections import namedtuple
from contextlib import AbstractContextManager
from datetime import timedelta
from unittest import mock
from unittest.mock import MagicMock
import pytest
from django.contrib.auth.models import Permission
@@ -18,19 +19,16 @@ from imap_tools import NOT
from imap_tools import EmailAddress
from imap_tools import FolderInfo
from imap_tools import MailboxFolderSelectError
from imap_tools import MailboxLoginError
from imap_tools import MailMessage
from imap_tools import MailMessageFlags
from imap_tools import errors
from rest_framework import status
from rest_framework.test import APITestCase
from documents.models import Correspondent
from documents.models import MatchingModel
from documents.tests.utils import FileSystemAssertsMixin
from paperless_mail import tasks
from paperless_mail.mail import MailAccountHandler
from paperless_mail.mail import MailError
from paperless_mail.mail import PinnedIMAP4
from paperless_mail.mail import TagMailAction
from paperless_mail.mail import apply_mail_action
from paperless_mail.mail import error_callback
@@ -40,265 +38,17 @@ from paperless_mail.models import MailRule
from paperless_mail.models import ProcessedMail
from paperless_mail.tests.factories import MailAccountFactory
from paperless_mail.tests.factories import MailRuleFactory
from paperless_mail.tests.helpers import BogusMailBox
from paperless_mail.tests.helpers import MessageBuilder
from paperless_mail.tests.helpers import _AttachmentDef
from paperless_mail.tests.helpers import fake_magic_from_buffer
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import CorrespondentFactory
from paperless_testing.factories import UserFactory
from paperless_testing.permissions import grant_global
@dataclasses.dataclass
class _AttachmentDef:
filename: str = "a_file.pdf"
maintype: str = "application/pdf"
subtype: str = "pdf"
disposition: str = "attachment"
content: bytes = b"a PDF document"
class BogusFolderManager:
current_folder = "INBOX"
uidvalidity = "1"
def set(self, new_folder) -> None:
if new_folder not in ["INBOX", "spam"]:
raise MailboxFolderSelectError(None, "uhm")
self.current_folder = new_folder
def status(self, folder, options):
return {"UIDVALIDITY": self.uidvalidity}
class BogusClient:
def __init__(self, messages) -> None:
self.messages: list[MailMessage] = messages
self.capabilities: list[str] = []
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
pass
def authenticate(self, mechanism, authobject) -> None:
# authobject must be a callable object
auth_bytes = authobject(None)
if auth_bytes != b"\x00admin\x00w57\xc3\xa4\xc3\xb6\xc3\xbcw4b6huwb6nhu":
raise MailboxLoginError("BAD", "OK")
def uid(self, command, *args) -> None:
if command == "STORE":
for message in self.messages:
if message.uid == args[0]:
flag = args[2]
if flag == "processed":
message._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(message, "flags"):
del message.flags
class BogusMailBox(AbstractContextManager):
# Common values so tests don't need to remember an accepted login
USERNAME: str = "admin"
ASCII_PASSWORD: str = "secret"
# Note the non-ascii characters here
UTF_PASSWORD: str = "w57äöüw4b6huwb6nhu"
# A dummy access token
ACCESS_TOKEN = "ea7e075cd3acf2c54c48e600398d5d5a"
def __init__(self) -> None:
self.messages: list[MailMessage] = []
self.messages_spam: list[MailMessage] = []
self.folder = BogusFolderManager()
self.client = BogusClient(self.messages)
self._host = ""
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
pass
def updateClient(self) -> None:
self.client = BogusClient(self.messages)
def login(self, username, password) -> None:
# This will raise a UnicodeEncodeError if the password is not ASCII only
password.encode("ascii")
# Otherwise, check for correct values
if username != self.USERNAME or password != self.ASCII_PASSWORD:
raise MailboxLoginError("BAD", "OK")
def login_utf8(self, username, password) -> None:
# Expected to only be called with the UTF-8 password
if username != self.USERNAME or password != self.UTF_PASSWORD:
raise MailboxLoginError("BAD", "OK")
def xoauth2(self, username: str, access_token: str) -> None:
if username != self.USERNAME or access_token != self.ACCESS_TOKEN:
raise MailboxLoginError("BAD", "OK")
def fetch(
self,
criteria="ALL",
charset="",
*,
mark_seen=True,
bulk=True,
uid_list=None,
):
if uid_list is not None:
return [m for m in self.messages if m.uid in uid_list]
return self._filter_messages(criteria)
def uids(self, criteria, charset="") -> list[str]:
return [m.uid for m in self._filter_messages(criteria)]
def _filter_messages(self, criteria):
msg = self.messages
criteria = str(criteria).strip("()").split(" ")
if "UNSEEN" in criteria:
msg = filter(lambda m: not m.seen, msg)
if "SUBJECT" in criteria:
subject = criteria[criteria.index("SUBJECT") + 1].strip('"')
msg = filter(lambda m: subject in m.subject, msg)
if "BODY" in criteria:
body = criteria[criteria.index("BODY") + 1].strip('"')
msg = filter(lambda m: body in m.text, msg)
if "FROM" in criteria:
from_ = criteria[criteria.index("FROM") + 1].strip('"')
msg = filter(lambda m: from_ in m.from_, msg)
if "TO" in criteria:
to_ = criteria[criteria.index("TO") + 1].strip('"')
msg = filter(lambda m: any(to_ in to_addr for to_addr in m.to), msg)
if "UNFLAGGED" in criteria:
msg = filter(lambda m: not m.flagged, msg)
if "UNKEYWORD" in criteria:
tag = criteria[criteria.index("UNKEYWORD") + 1].strip("'")
msg = filter(lambda m: tag not in m.flags, msg)
if "(X-GM-LABELS" in criteria: # ['NOT', '(X-GM-LABELS', '"processed"']
msg = filter(lambda m: "processed" not in m.flags, msg)
if "UID" in criteria:
uid_list = criteria[criteria.index("UID") + 1].split(",")
msg = filter(lambda m: m.uid in uid_list, msg)
return list(msg)
def delete(self, uid_list) -> None:
self.messages = list(filter(lambda m: m.uid not in uid_list, self.messages))
def flag(self, uid_list, flag_set, value) -> None:
for message in self.messages:
if message.uid in uid_list:
for flag in flag_set:
if flag == MailMessageFlags.FLAGGED:
message.flagged = value
if flag == MailMessageFlags.SEEN:
message.seen = value
if flag == "processed":
message._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(message, "flags"):
del message.flags
def move(self, uid_list, folder) -> None:
if folder == "spam":
self.messages_spam += list(
filter(lambda m: m.uid in uid_list, self.messages),
)
self.messages = list(filter(lambda m: m.uid not in uid_list, self.messages))
else:
raise Exception
def fake_magic_from_buffer(buffer, *, mime=False):
if mime:
if "PDF" in str(buffer):
return "application/pdf"
else:
return "unknown/type"
else:
return "Some verbose file description"
class MessageBuilder:
def __init__(self) -> None:
self._next_uid = 1
def create_message(
self,
*,
attachments: int | list[_AttachmentDef] = 1,
body: str = "",
subject: str = "the subject",
from_: str = "no_one@mail.com",
to: list[str] | None = None,
seen: bool = False,
flagged: bool = False,
processed: bool = False,
) -> MailMessage:
if to is None:
to = ["tosomeone@somewhere.com"]
email_msg = email.message.EmailMessage()
# TODO: This does NOT set the UID
email_msg["Message-ID"] = str(uuid.uuid4())
email_msg["Subject"] = subject
email_msg["From"] = from_
email_msg["To"] = str(" ,".join(to))
email_msg.set_content(body)
# Either add some default number of attachments
# or the provided attachments
if isinstance(attachments, int):
for i in range(attachments):
attachment = _AttachmentDef(filename=f"file_{i}.pdf")
email_msg.add_attachment(
attachment.content,
maintype=attachment.maintype,
subtype=attachment.subtype,
disposition=attachment.disposition,
filename=attachment.filename,
)
else:
for attachment in attachments:
email_msg.add_attachment(
attachment.content,
maintype=attachment.maintype,
subtype=attachment.subtype,
disposition=attachment.disposition,
filename=attachment.filename,
)
# Convert the EmailMessage to an imap_tools MailMessage
imap_msg = MailMessage.from_bytes(email_msg.as_bytes())
# TODO: Unsure how to add a uid to the actual EmailMessage. This hacks it in,
# based on how imap_tools uses regex to extract it.
# This should be a large enough pool
uid = self._next_uid
self._next_uid += 1
imap_msg._raw_uid_data = f"UID {uid}".encode()
imap_msg.seen = seen
imap_msg.flagged = flagged
if processed:
imap_msg._raw_flag_data.append(b"+FLAGS (processed)")
if hasattr(imap_msg, "flags"):
del imap_msg.flags
return imap_msg
def reset_bogus_mailbox(
bogus_mailbox: BogusMailBox,
message_builder: MessageBuilder,
@@ -2299,10 +2049,13 @@ class TestMailAccountTestView(APITestCase):
self.assertEqual(response.content.decode(), "Unable to connect to server")
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
@mock.patch("paperless_mail.mail.resolve_hostname_ips", return_value=["127.0.0.1"])
@mock.patch(
"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(
self,
_mock_resolve_hostname_ips,
_mock_getaddrinfo: MagicMock,
) -> None:
data = {
"imap_server": "internal.example",
@@ -2459,10 +2212,10 @@ class TestGetMailboxHostPinning(TestCase):
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
@mock.patch(
"paperless_mail.mail.resolve_hostname_ips",
return_value=["93.184.216.34"],
"paperless_mail.mail.resolve_public_addresses",
return_value=(ipaddress.ip_address("93.184.216.34"),),
)
def test_connects_to_validated_ip(self, _mock_resolve) -> None:
def test_connects_to_validated_ip(self, _mock_resolve: MagicMock) -> None:
with mock.patch(
"paperless_mail.mail.socket.create_connection",
side_effect=OSError("no connection in tests"),
@@ -2479,10 +2232,13 @@ class TestGetMailboxHostPinning(TestCase):
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
@mock.patch(
"paperless_mail.mail.resolve_hostname_ips",
return_value=["93.184.216.34"],
"paperless_mail.mail.resolve_public_addresses",
return_value=(ipaddress.ip_address("93.184.216.34"),),
)
def test_ssl_pins_ip_but_keeps_hostname_for_sni(self, _mock_resolve) -> None:
def test_ssl_pins_ip_but_keeps_hostname_for_sni(
self,
_mock_resolve: MagicMock,
) -> None:
ssl_context = mock.MagicMock()
ssl_context.wrap_socket.return_value.makefile.side_effect = OSError(
"no connection in tests",
@@ -2513,13 +2269,51 @@ class TestGetMailboxHostPinning(TestCase):
@override_settings(EMAIL_ALLOW_INTERNAL_HOSTS=False)
@mock.patch(
"paperless_mail.mail.resolve_hostname_ips",
return_value=["93.184.216.34", "127.0.0.1"],
"paperless.network._getaddrinfo",
return_value=[
(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(self, _mock_resolve) -> None:
with self.assertRaises(MailError):
def test_blocks_when_any_resolved_address_is_internal(
self,
_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)
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):
def setUp(self) -> None:
+3 -3
View File
@@ -15,9 +15,9 @@ import pytest
from paperless_mail.models import MailRule
from paperless_mail.tests.factories import MailAccountFactory
from paperless_mail.tests.test_mail import MessageBuilder
from paperless_mail.tests.test_mail import _AttachmentDef
from paperless_mail.tests.test_mail import fake_magic_from_buffer
from paperless_mail.tests.helpers import MessageBuilder
from paperless_mail.tests.helpers import _AttachmentDef
from paperless_mail.tests.helpers import fake_magic_from_buffer
@pytest.fixture()
@@ -16,8 +16,8 @@ from paperless_mail.mail import MailAccountHandler
from paperless_mail.models import MailRule
from paperless_mail.preprocessor import MailMessageDecryptor
from paperless_mail.tests.factories import MailAccountFactory
from paperless_mail.tests.helpers import _AttachmentDef
from paperless_mail.tests.test_mail import TestMail
from paperless_mail.tests.test_mail import _AttachmentDef
class MessageEncryptor:
+37
View File
@@ -0,0 +1,37 @@
"""Filesystem assertions for unittest-style tests."""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from os import PathLike
class FileSystemAssertsMixin:
def assertIsFile(self, path: PathLike[str] | str) -> None:
if not Path(path).resolve().is_file():
raise AssertionError(f"File does not exist: {path}")
def assertIsNotFile(self, path: PathLike[str] | str) -> None:
if Path(path).resolve().is_file():
raise AssertionError(f"File does exist: {path}")
def assertIsDir(self, path: PathLike[str] | str) -> None:
if not Path(path).resolve().is_dir():
raise AssertionError(f"Dir does not exist: {path}")
def assertIsNotDir(self, path: PathLike[str] | str) -> None:
if Path(path).resolve().is_dir():
raise AssertionError(f"Dir does exist: {path}")
def assertFileCountInDir(self, path: PathLike[str] | str, count: int) -> None:
path = Path(path).resolve()
if not path.is_dir():
raise AssertionError(f"Path {path} is not a directory")
found = len([x for x in path.iterdir() if x.is_file()])
if found != count:
raise AssertionError(
f"Path {path} contains {found} files instead of {count} files",
)
+31
View File
@@ -0,0 +1,31 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from documents.plugins.helpers import ProgressManager
if TYPE_CHECKING:
from documents.plugins.helpers import WebsocketPayload
class FakeProgressManager(ProgressManager):
"""
The real ProgressManager with the channel layer cut out: send_progress still
builds the payload, so it cannot drift, and the payloads are recorded instead
of being sent to Redis.
Use it through the `fake_progress_manager` fixture, or construct it directly.
"""
def __init__(self, filename: str | None = None, task_id: str | None = None) -> None:
super().__init__(filename, task_id)
self.payloads: list[WebsocketPayload] = []
def open(self) -> None:
pass
def close(self) -> None:
pass
def send(self, payload: WebsocketPayload) -> None:
self.payloads.append(payload)
+13
View File
@@ -0,0 +1,13 @@
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from django.http import StreamingHttpResponse
def read_streaming_response(response: StreamingHttpResponse) -> bytes:
"""Consume a StreamingHttpResponse/FileResponse and close it."""
content = b"".join(response.streaming_content)
response.close()
return content
+65
View File
@@ -0,0 +1,65 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import Any
from django.apps import apps
from django.db import connection
from django.db.migrations.executor import MigrationExecutor
from django.test import TransactionTestCase
if TYPE_CHECKING:
from django.apps.registry import Apps
class TestMigrations(TransactionTestCase):
@property
def app(self) -> str:
return apps.get_containing_app_config(type(self).__module__).name
migrate_from: Any = None
dependencies: list[tuple[str, str]] | None = None
migrate_to: Any = None
def setUp(self) -> None:
super().setUp()
assert self.migrate_from and self.migrate_to, (
f"TestCase '{type(self).__name__}' must define migrate_from and migrate_to properties"
)
self.migrate_from = [(self.app, self.migrate_from)]
if self.dependencies is not None:
self.migrate_from.extend(self.dependencies)
self.migrate_to = [(self.app, self.migrate_to)]
executor = MigrationExecutor(connection)
old_apps = executor.loader.project_state(self.migrate_from).apps
# Reverse to the original migration
executor.migrate(self.migrate_from)
self.setUpBeforeMigration(old_apps)
self.apps = old_apps
# Run the migration to test
executor = MigrationExecutor(connection)
executor.loader.build_graph() # reload.
executor.migrate(self.migrate_to)
self.apps = executor.loader.project_state(self.migrate_to).apps
def setUpBeforeMigration(self, apps: Apps) -> None:
pass
def tearDown(self) -> None:
"""
Ensure the database schema is restored to the latest migration after
each migration test, so subsequent tests run against HEAD.
"""
try:
executor = MigrationExecutor(connection)
executor.loader.build_graph()
targets = executor.loader.graph.leaf_nodes()
executor.migrate(targets)
finally:
super().tearDown()
+218
View File
@@ -0,0 +1,218 @@
"""
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
+76
View File
@@ -0,0 +1,76 @@
from __future__ import annotations
import time
import warnings
from typing import TYPE_CHECKING
from typing import Any
import httpx
import pytest
from documents.parsers import ParseError
if TYPE_CHECKING:
from collections.abc import Callable
def util_call_with_backoff(
method_or_callable: Callable,
args: list | tuple,
*,
skip_on_50x_err: bool = True,
) -> tuple[bool, Any]:
"""
For whatever reason, the images started during the test pipeline like to
segfault sometimes, crash and otherwise fail randomly, when run with the
exact files that usually pass.
So, this function will retry the given method/function up to 3 times, with larger backoff
periods between each attempt, in hopes the issue resolves itself during
one attempt to parse.
This will wait the following:
- Attempt 1 - 20s following failure
- Attempt 2 - 40s following failure
- Attempt 3 - 80s following failure
"""
result = None
succeeded = False
retry_time = 20.0
retry_count = 0
status_codes = []
max_retry_count = 3
while retry_count < max_retry_count and not succeeded:
try:
result = method_or_callable(*args)
succeeded = True
except ParseError as e: # pragma: no cover
cause_exec = e.__cause__
if cause_exec is not None and isinstance(cause_exec, httpx.HTTPStatusError):
status_codes.append(cause_exec.response.status_code)
warnings.warn(
f"HTTP Exception for {cause_exec.request.url} - {cause_exec}",
)
else:
warnings.warn(f"Unexpected error: {e}")
except Exception as e: # pragma: no cover
warnings.warn(f"Unexpected error: {e}")
retry_count = retry_count + 1
if not succeeded and retry_count < max_retry_count:
time.sleep(retry_time)
retry_time = retry_time * 2.0
if (
not succeeded
and status_codes
and skip_on_50x_err
and all(httpx.codes.is_server_error(code) for code in status_codes)
):
pytest.skip("Repeated HTTP 50x for service") # pragma: no cover
return succeeded, result
Generated
+6
View File
@@ -2874,6 +2874,7 @@ name = "paperless-ngx"
version = "3.2.1"
source = { virtual = "." }
dependencies = [
{ name = "anyio" },
{ name = "azure-ai-documentintelligence" },
{ name = "babel" },
{ name = "bleach" },
@@ -2902,6 +2903,8 @@ dependencies = [
{ name = "filelock" },
{ name = "flower" },
{ name = "gotenberg-client", extra = ["httpx"] },
{ name = "httpcore" },
{ name = "httpx" },
{ name = "httpx-oauth" },
{ name = "ijson" },
{ name = "imap-tools" },
@@ -3026,6 +3029,7 @@ typing = [
[package.metadata]
requires-dist = [
{ name = "anyio", specifier = ">=4.12" },
{ name = "azure-ai-documentintelligence", specifier = ">=1.0.2" },
{ name = "babel", specifier = ">=2.17" },
{ name = "bleach", specifier = "~=6.4.0" },
@@ -3055,6 +3059,8 @@ requires-dist = [
{ name = "flower", specifier = ">=2.0.1,<2.2" },
{ name = "gotenberg-client", extras = ["httpx"], specifier = "~=1.0" },
{ name = "granian", extras = ["uvloop"], marker = "extra == 'webserver'", specifier = ">=2.7,<2.9" },
{ name = "httpcore", specifier = "~=1.0.9" },
{ name = "httpx", specifier = "~=0.28.1" },
{ name = "httpx-oauth", specifier = "~=0.17" },
{ name = "ijson", specifier = ">=3.5.1" },
{ name = "imap-tools", specifier = ">=1.14,<1.16" },