mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-28 12:20:34 +00:00
Compare commits
31
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
296ddff37e | ||
|
|
de8edc31c5 | ||
|
|
46c150eb67 | ||
|
|
58777076c5 | ||
|
|
d414297c05 | ||
|
|
d53a930aed | ||
|
|
78ca920771 | ||
|
|
4c5255a8dd | ||
|
|
283e49b1ee | ||
|
|
3130fc3a7c | ||
|
|
8f56c6167b | ||
|
|
eb644f3fac | ||
|
|
6ad00ca55a | ||
|
|
526adbad1a | ||
|
|
d53453fb71 | ||
|
|
4eaf8032ad | ||
|
|
9e1f938d56 | ||
|
|
0da50ad348 | ||
|
|
5e971bc0ce | ||
|
|
17a91c385d | ||
|
|
15b73b890c | ||
|
|
02d355061f | ||
|
|
d8b5b4d447 | ||
|
|
03ac4aed7e | ||
|
|
90d23bad9c | ||
|
|
e4367b5648 | ||
|
|
cb85441c2f | ||
|
|
cceaa559d4 | ||
|
|
a748d4c64f | ||
|
|
452ed005bd | ||
|
|
40058ff7d5 |
@@ -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
@@ -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}
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
+1450
-73
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() == []
|
||||
@@ -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):
|
||||
|
||||
@@ -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,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."""
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user