Compare commits

...
Author SHA1 Message Date
stumpylog 0ef2c39c0b Assert permission assignment does not scale with selection size
The two batching tests only checked that permissions came out correct at
5 and 50 objects, so reverting to the old per-object loop would still
have passed.  Check sizes as well to prevent that
2026-09-09 12:21:38 -07:00
stumpylog 9e076e7cdf Resolve every permission action before applying any of them
This fixes the existing issue and resolves the Copilot comment
2026-09-09 12:21:12 -07:00
stumpylog e1df61a192 Perf: drop speculative row-chunking in bulk permission assignment, keep the query-batching fix 2026-09-09 11:11:23 -07:00
stumpylog 41524a6230 Mark empty-pks early-return in set_permissions_for_objects as no-cover
Defensive guard for an edge case (all requested pks already gone/invalid)
rather than a path normal usage exercises; matches the existing
pragma: no cover convention elsewhere in this file.
2026-09-09 11:11:23 -07:00
stumpylog e5c9b8facb Fix: use .distinct() for existing-grant lookup, drop flaky query-count invariant tests
.distinct() lets the database dedupe identity ids server-side instead of
transferring one row per (object, grantee) match and deduping in Python --
was the dominant cost on a large selection with existing grants.

Also replaced the two query-count-equality tests (bulk_edit and the
bulk_edit_objects API path) with plain functional-correctness checks at
both batch sizes.  Hopefully stops that flake.
2026-09-09 11:11:23 -07:00
stumpylog cf3a080694 Perf: avoid unnecessary full-row fetches in batch permission assignment
set_permissions_for_objects now takes a model + pks instead of instances,
and identity filtering resolves straight to ids, so bulk-editing
permissions no longer materializes full Document/User/Group rows just to
read their pk/id. Row construction for bulk_create is also chunked to
bound peak memory for very large "apply to all" operations.
2026-09-09 11:11:23 -07:00
stumpylog 018de125ff Perf: batch guardian permission assignment in bulk-edit
bulk_edit.set_permissions and BulkEditObjectPermissionsView both
looped documents/objects and called set_permissions_for_object per
object, which itself calls guardian's assign_perm/remove_perm once
per (object, user) pair -- ~10-20+ queries per object, scaling with
selection size.

Added set_permissions_for_objects, a bulk equivalent that resolves
existing permission holders once across the whole batch (not once per
object) and applies changes with a small, batch-size-independent
number of queries per action instead of one per (object, user) pair.
2026-09-09 11:11:22 -07:00
Trenton H 310628699d Fix: If a file remains stable but has zero size after the stability window has passed, drop it from tracking (#14047)
A slow writer might create a zero byte file, then take longer than the window to
finish the write.  We would then queue a zero byte file for consumption and race the writer
to most likely fail due to still being empty.  Instead, drop zero size files at yeild time.
The slow writer may or may not finish, but if it does, there will be a Change.modified
event fired again
2026-09-09 16:24:17 +00:00
GitHub Actions aff0f9cf41 Auto translate strings 2026-09-09 16:04:18 +00:00
shamoon bf716ebfd1 Fixhancement: better LLM errors (#14031) 2026-09-09 16:02:43 +00:00
GitHub Actions 7d67a10a35 Auto translate strings 2026-09-09 15:50:33 +00:00
shamoon 8d1bc5dd24 Fix: prevent orphaned versions from bulk delete (#14030) 2026-09-09 08:48:59 -07:00
Trenton H 43a8d7d412 Enhancement: add Tantivy full-text fallback adapter for taxonomy candidates (#13820)
This brings users without an embedding backend configured to closer
parity with those who do.  Reuse the search backend to locate similar
documents and use them to provide the LLM with the better suggestion pool
to draw from
2026-09-09 15:15:15 +00:00
Trenton H c40922440b Fix: prevent overlapping mail-account processing runs (#14046)
process_mail_accounts had no guard against a scheduled run still being
in progress when the next one fires. Skip a run outright
if another MAIL_FETCH task is already PENDING/STARTED.
2026-09-09 07:49:49 -07:00
23 changed files with 1799 additions and 560 deletions
+9 -5
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import logging import logging
import tempfile import tempfile
import uuid
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from typing import Literal from typing import Literal
@@ -27,7 +28,7 @@ from documents.models import DocumentType
from documents.models import PaperlessTask from documents.models import PaperlessTask
from documents.models import StoragePath from documents.models import StoragePath
from documents.models import Tag from documents.models import Tag
from documents.permissions import set_permissions_for_object from documents.permissions import set_permissions_for_objects
from documents.plugins.helpers import DocumentsStatusManager from documents.plugins.helpers import DocumentsStatusManager
from documents.tasks import bulk_update_documents from documents.tasks import bulk_update_documents
from documents.tasks import consume_file from documents.tasks import consume_file
@@ -379,7 +380,7 @@ def delete(doc_ids: list[int]) -> Literal["OK"]:
) )
delete_ids = list({*doc_ids, *version_ids}) delete_ids = list({*doc_ids, *version_ids})
Document.objects.filter(id__in=delete_ids).delete() Document.objects.filter(id__in=delete_ids).delete(transaction_id=uuid.uuid4())
from documents.search import get_backend from documents.search import get_backend
@@ -430,10 +431,13 @@ def set_permissions(
else: else:
qs.update(owner=owner) qs.update(owner=owner)
for doc in qs:
set_permissions_for_object(permissions=set_permissions, object=doc, merge=merge)
affected_docs = list(qs.values_list("pk", flat=True)) affected_docs = list(qs.values_list("pk", flat=True))
set_permissions_for_objects(
permissions=set_permissions,
model=Document,
pks=affected_docs,
merge=merge,
)
bulk_update_documents.apply_async( bulk_update_documents.apply_async(
kwargs={"document_ids": affected_docs}, kwargs={"document_ids": affected_docs},
@@ -156,6 +156,15 @@ class FileStabilityTracker:
logger.debug(f"File disappeared during stability check: {path}") logger.debug(f"File disappeared during stability check: {path}")
continue continue
# Stable, but empty: some scanners create a zero byte placeholder
# and only write the page some time later. Consuming it now can
# only fail so drop it and let the writer's next event
# (or the periodic rescan) bring it back once it has content
if not tracked.last_size:
to_remove.append(path)
logger.debug("Ignoring stable but empty file: %s", path)
continue
# File is stable, we can return it # File is stable, we can return it
to_yield.append(path) to_yield.append(path)
logger.info(f"File is stable: {path}") logger.info(f"File is stable: {path}")
+10 -2
View File
@@ -1,4 +1,5 @@
import datetime import datetime
import uuid
from pathlib import Path from pathlib import Path
from typing import Final from typing import Final
@@ -514,13 +515,20 @@ class Document(SoftDeleteModel, ModelWithOwner): # type: ignore[django-manager-
def delete( def delete(
self, self,
*args, *args,
transaction_id=None,
**kwargs, **kwargs,
): ):
# If deleting a root document, move all its versions to trash as well. # Versions must share the root's transaction ID so they are restored
# together by django-softdelete.
if transaction_id is None:
transaction_id = uuid.uuid4()
if self.root_document_id is None: if self.root_document_id is None:
Document.objects.filter(root_document=self).delete() Document.objects.filter(root_document=self).delete(
transaction_id=transaction_id,
)
return super().delete( return super().delete(
*args, *args,
transaction_id=transaction_id,
**kwargs, **kwargs,
) )
+173
View File
@@ -173,6 +173,179 @@ def set_permissions_for_object(
) )
def _resolve_permissions(codenames: set[str], ctype: ContentType) -> list[Permission]:
"""
Resolves `codenames` to Permission rows, raising like the single-object
assign_perm() this bulk path replaces does (via a `.get()` internally)
if any codename doesn't exist -- e.g. a client-supplied action name that
was never validated (BulkEditObjectsSerializer._validate_permissions
calls validate_set_permissions() only for its side-effecting id checks
and discards the filtered dict it returns, so an unrecognized action key
reaches this function as-is). A plain `.filter()` with no existence
check would otherwise silently build zero rows and no-op instead of
reporting the bad input.
"""
permission_objs = list(
Permission.objects.filter(content_type=ctype, codename__in=codenames),
)
missing = codenames - {p.codename for p in permission_objs}
if missing:
raise Permission.DoesNotExist(
f"Permission matching query does not exist for codename(s): "
f"{', '.join(sorted(missing))}",
)
return permission_objs
def _apply_bulk_permission_entry(
*,
perm_model: type[UserObjectPermission] | type[GroupObjectPermission],
identity_model: type[User] | type[Group],
identity_field: str,
ids: list[int],
codename: str,
permission_objs: list[Permission],
ctype: ContentType,
object_pks: list[str],
merge: bool,
) -> None:
# Only the ids are needed to build permission rows (via `<field>_id=`),
# so avoid fetching full User/Group rows for identities that may not
# even end up being granted anything new.
add_ids = set(
identity_model.objects.filter(id__in=ids).values_list("id", flat=True),
)
if not merge:
existing_ids = set(
perm_model.objects.filter(
content_type=ctype,
object_pk__in=object_pks,
permission__codename=codename,
)
.values_list(f"{identity_field}_id", flat=True)
.distinct(),
)
remove_ids = existing_ids - add_ids
if remove_ids:
perm_model.objects.filter(
content_type=ctype,
object_pk__in=object_pks,
permission__codename=codename,
**{f"{identity_field}_id__in": remove_ids},
).delete()
if not add_ids:
return
rows = [
perm_model(
content_type=ctype,
object_pk=pk,
permission=permission_obj,
**{f"{identity_field}_id": identity_id},
)
for permission_obj in permission_objs
for pk in object_pks
for identity_id in add_ids
]
# ignore_conflicts skips only rows that already exist as an exact
# (identity, permission, object) match -- the same de-dup the
# underlying (user|group, permission, object_pk) unique constraint
# already enforces for the single-object assign_perm() this replaces,
# so it doesn't change what counts as "already granted". batch_size
# caps how many rows go into a single INSERT statement.
perm_model.objects.bulk_create(rows, ignore_conflicts=True, batch_size=1000)
def set_permissions_for_objects(
permissions: dict,
model: type[Model],
pks: QuerySet | list,
*,
merge: bool = False,
) -> None:
"""
Bulk equivalent of set_permissions_for_object: applies the same
permission changes to every object identified by `pks` at once.
Takes a model + pks (rather than model instances) deliberately -- the
permission rows built below only ever need `pk`, `content_type`, and
identity ids, so callers shouldn't have to fetch full rows (with every
other field) just to hand them to this function.
Deliberately does not use guardian's queryset/list-aware assign_perm:
passing a list as the object routes to bulk_assign_perm, which skips
creating a direct permission row for anyone who already has the
permission via ANY group membership (it checks
ObjectPermissionChecker.has_perm, which is group-inheritance-aware) --
unlike the single-object assign_perm this replaces, which always
ensures a direct row via get_or_create regardless of group-derived
access. Losing that guarantee would mean a later revocation of the
group's grant silently strips access an admin explicitly asked to be
direct. Bulk-creating rows straight against the permission models
instead (see _apply_bulk_permission_entry) preserves the original
always-create-a-direct-row semantics while still batching every object
and every identity into one query per action, rather than one query per
(object, user) pair.
"""
object_pks = [str(pk) for pk in pks]
if not object_pks: # pragma: no cover
return
model_name = model.__name__.lower()
ctype = ContentType.objects.get_for_model(model)
# Every action is resolved up front, before anything is written, so an
# unrecognized action name (see _resolve_permissions) aborts the whole
# call instead of leaving the actions ahead of it already applied --
# BulkEditObjectsSerializer lets unknown keys through and its view turns
# the exception into a 400, so a half-applied change would otherwise be
# reported to the client as a failure.
permissions_by_action: dict[str, list[Permission]] = {}
for action, entry in permissions.items():
if "users" not in entry and "groups" not in entry:
continue
implied_codenames = {f"{action}_{model_name}"}
if action == "change":
# change gives view too
implied_codenames.add(f"view_{model_name}")
permissions_by_action[action] = _resolve_permissions(
implied_codenames,
ctype,
)
for action, entry in permissions.items():
codename = f"{action}_{model_name}"
permission_objs = permissions_by_action.get(action, [])
if "users" in entry:
_apply_bulk_permission_entry(
perm_model=UserObjectPermission,
identity_model=User,
identity_field="user",
ids=entry["users"],
codename=codename,
permission_objs=permission_objs,
ctype=ctype,
object_pks=object_pks,
merge=merge,
)
if "groups" in entry:
_apply_bulk_permission_entry(
perm_model=GroupObjectPermission,
identity_model=Group,
identity_field="group",
ids=entry["groups"],
codename=codename,
permission_objs=permission_objs,
ctype=ctype,
object_pks=object_pks,
merge=merge,
)
def permitted_object_ids( def permitted_object_ids(
user: User | None, user: User | None,
model: type[Model], model: type[Model],
+65
View File
@@ -2,10 +2,15 @@ import datetime
import json import json
from unittest import mock from unittest import mock
from django.contrib.auth.models import Group
from django.contrib.auth.models import Permission from django.contrib.auth.models import Permission
from django.contrib.auth.models import User from django.contrib.auth.models import User
from django.db import connection
from django.test import override_settings from django.test import override_settings
from django.test.utils import CaptureQueriesContext
from guardian.shortcuts import assign_perm from guardian.shortcuts import assign_perm
from guardian.shortcuts import get_groups_with_perms
from guardian.shortcuts import get_users_with_perms
from rest_framework import status from rest_framework import status
from rest_framework.test import APITestCase from rest_framework.test import APITestCase
@@ -842,6 +847,66 @@ class TestBulkEditObjects(APITestCase):
self.assertEqual(response.status_code, status.HTTP_200_OK) self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(StoragePath.objects.count(), 0) self.assertEqual(StoragePath.objects.count(), 0)
def test_bulk_objects_set_permissions_batched_across_object_count(
self,
) -> None:
"""
GIVEN:
- Many tags are being bulk-edited to set permissions at once
WHEN:
- bulk_edit_objects API endpoint is called with set_permissions
operation over a small batch vs. a much larger one
THEN:
- Permissions are applied correctly at both scales
- Query count does not grow with the number of tags, i.e. each
user/group is applied across all tags with one batched call
rather than one call per (tag, identity) pair
"""
group1 = Group.objects.create(name="perm-group")
permissions = {
"view": {"users": [self.user1.id, self.user2.id], "groups": [group1.id]},
"change": {"users": [self.user1.id], "groups": [group1.id]},
}
def run_with_n_tags(n: int) -> int:
tags = [Tag.objects.create(name=f"perm-tag-{n}-{i}") for i in range(n)]
with CaptureQueriesContext(connection) as ctx:
response = self.client.post(
"/api/bulk_edit_objects/",
json.dumps(
{
"objects": [t.id for t in tags],
"object_type": "tags",
"operation": "set_permissions",
"permissions": permissions,
"merge": False,
},
),
content_type="application/json",
)
self.assertEqual(response.status_code, status.HTTP_200_OK)
for tag in tags:
self.assertEqual(get_users_with_perms(tag).count(), 2)
self.assertEqual(get_groups_with_perms(tag).count(), 1)
return len(ctx.captured_queries)
small_batch_queries = run_with_n_tags(5)
large_batch_queries = run_with_n_tags(50)
# A tolerance rather than equality, matching the N+1 check in
# test_views.py: bulk_create's batch_size caps rows per INSERT, so a
# large enough selection does legitimately add statements, and the
# per-process ContentType cache makes the first run carry an extra
# query. Neither can hide a regression to per-object assignment,
# which would be ~10x the small-batch count here.
self.assertLessEqual(
large_batch_queries,
small_batch_queries + 5,
"Permission assignment appears to scale with object count: "
f"{small_batch_queries} queries for 5 tags vs. "
f"{large_batch_queries} for 50",
)
def test_bulk_objects_delete_all_filtered(self) -> None: def test_bulk_objects_delete_all_filtered(self) -> None:
""" """
GIVEN: GIVEN:
+62
View File
@@ -207,3 +207,65 @@ class TestTrashAPI(DirectoriesMixin, APITestCase):
) )
self.assertEqual(resp.status_code, status.HTTP_400_BAD_REQUEST) self.assertEqual(resp.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("have not yet been deleted", resp.data["documents"][0]) self.assertIn("have not yet been deleted", resp.data["documents"][0])
def _make_versioned_document(self) -> tuple[Document, list[Document]]:
root = Document.objects.create(
title="root",
content="root-content",
checksum="root",
mime_type="application/pdf",
)
versions = [
Document.objects.create(
title=f"v{index}",
content=f"v{index}-content",
checksum=f"v{index}",
mime_type="application/pdf",
root_document=root,
version_index=index,
)
for index in range(1, 3)
]
return root, versions
def test_api_trash_restore_document_restores_its_versions(self) -> None:
"""
GIVEN:
- Existing document with two versions
WHEN:
- API request to delete the document
- API request to restore it from the trash
THEN:
- Only the document itself is listed in the trash
- A version cannot be restored without its root
- The document is restored together with all of its versions
"""
root, versions = self._make_versioned_document()
self.client.force_login(user=self.user)
self.client.delete(f"/api/documents/{root.pk}/")
self.assertEqual(Document.deleted_objects.count(), 3)
resp = self.client.get("/api/trash/")
self.assertEqual(resp.status_code, status.HTTP_200_OK)
self.assertEqual(resp.data["count"], 1)
self.assertEqual(resp.data["results"][0]["id"], root.pk)
# A version cannot be restored while its root remains in the trash.
resp = self.client.post(
"/api/trash/",
{"action": "restore", "documents": [versions[0].pk]},
)
self.assertEqual(resp.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("Restore the root document", resp.data["documents"][0])
resp = self.client.post(
"/api/trash/",
{"action": "restore", "documents": [root.pk]},
)
self.assertEqual(resp.status_code, status.HTTP_200_OK)
self.assertEqual(Document.deleted_objects.count(), 0)
self.assertCountEqual(
Document.objects.filter(root_document=root).values_list("id", flat=True),
[version.pk for version in versions],
)
+181
View File
@@ -5,8 +5,11 @@ from unittest import mock
import pikepdf import pikepdf
from django.contrib.auth.models import Group from django.contrib.auth.models import Group
from django.contrib.auth.models import Permission
from django.contrib.auth.models import User from django.contrib.auth.models import User
from django.db import connection
from django.test import TestCase from django.test import TestCase
from django.test.utils import CaptureQueriesContext
from guardian.shortcuts import assign_perm from guardian.shortcuts import assign_perm
from guardian.shortcuts import get_groups_with_perms from guardian.shortcuts import get_groups_with_perms
from guardian.shortcuts import get_users_with_perms from guardian.shortcuts import get_users_with_perms
@@ -19,6 +22,7 @@ from documents.models import Document
from documents.models import DocumentType from documents.models import DocumentType
from documents.models import StoragePath from documents.models import StoragePath
from documents.models import Tag from documents.models import Tag
from documents.permissions import set_permissions_for_objects
from documents.tests.utils import DirectoriesMixin from documents.tests.utils import DirectoriesMixin
@@ -392,6 +396,11 @@ class TestBulkEdit(DirectoriesMixin, TestCase):
self.assertFalse(Document.objects.filter(id=self.doc1.id).exists()) self.assertFalse(Document.objects.filter(id=self.doc1.id).exists())
self.assertFalse(Document.objects.filter(id=version.id).exists()) self.assertFalse(Document.objects.filter(id=version.id).exists())
Document.deleted_objects.get(id=self.doc1.id).restore(strict=False)
self.assertTrue(Document.objects.filter(id=self.doc1.id).exists())
self.assertTrue(Document.objects.filter(id=version.id).exists())
def test_delete_version_document_keeps_root(self) -> None: def test_delete_version_document_keeps_root(self) -> None:
version = Document.objects.create( version = Document.objects.create(
checksum="A-v1", checksum="A-v1",
@@ -510,6 +519,178 @@ class TestBulkEdit(DirectoriesMixin, TestCase):
) )
self.assertEqual(groups_with_perms.count(), 2) self.assertEqual(groups_with_perms.count(), 2)
@mock.patch("documents.tasks.bulk_update_documents.apply_async")
def test_set_permissions_batched_across_document_count(
self,
m,
) -> None:
"""
GIVEN:
- Many documents are being bulk-edited to set permissions at once
WHEN:
- set_permissions runs over a small batch vs. a much larger one
THEN:
- Permissions are applied correctly at both scales
- Query count does not grow with the number of documents, i.e.
each user/group is applied across all documents with one
batched call rather than one call per (document, identity)
pair
"""
permissions = {
"view": {
"users": [self.user1.id, self.user2.id],
"groups": [self.group2.id],
},
"change": {
"users": [self.user1.id],
"groups": [self.group2.id],
},
}
def run_with_n_documents(n: int) -> int:
docs = [
Document.objects.create(checksum=f"perm-{n}-{i}", title=f"perm-{n}-{i}")
for i in range(n)
]
with CaptureQueriesContext(connection) as ctx:
bulk_edit.set_permissions(
[doc.id for doc in docs],
set_permissions=permissions,
owner=self.owner,
merge=False,
)
for doc in docs:
self.assertEqual(get_users_with_perms(doc).count(), 2)
self.assertEqual(get_groups_with_perms(doc).count(), 1)
return len(ctx.captured_queries)
small_batch_queries = run_with_n_documents(5)
large_batch_queries = run_with_n_documents(50)
# A tolerance rather than equality, matching the N+1 check in
# test_views.py: bulk_create's batch_size caps rows per INSERT, so a
# large enough selection does legitimately add statements, and the
# per-process ContentType cache makes the first run carry an extra
# query. Neither can hide a regression to per-document assignment,
# which would be ~10x the small-batch count here.
self.assertLessEqual(
large_batch_queries,
small_batch_queries + 5,
"Permission assignment appears to scale with document count: "
f"{small_batch_queries} queries for 5 documents vs. "
f"{large_batch_queries} for 50",
)
@mock.patch("documents.tasks.bulk_update_documents.apply_async")
def test_set_permissions_grants_direct_perm_even_if_already_granted_via_group(
self,
m,
) -> None:
"""
GIVEN:
- A user already has view access to a document via group
membership, with no direct grant of their own
WHEN:
- set_permissions explicitly grants that same user direct view
access via bulk_edit
THEN:
- A direct permission grant is created for the user, not skipped
because they already have equivalent access via the group
Regression test: guardian's queryset-aware assign_perm() (routed to
when the target is a list/queryset) skips creating a direct row for
anyone whose ObjectPermissionChecker.has_perm() already returns True
-- which includes group-derived access. The single-object assign_perm
this bulk path replaces has no such check; it always ensures a
direct row via get_or_create. Losing that guarantee would mean
revoking the group's grant later silently strips access that was
supposed to be explicit.
"""
self.doc1.owner = self.user1
self.doc1.save()
self.user1.groups.add(self.group1)
assign_perm("view_document", self.group1, self.doc1)
bulk_edit.set_permissions(
[self.doc1.id],
set_permissions={
"view": {"users": [self.user1.id], "groups": []},
},
merge=True,
)
direct_users = get_users_with_perms(
self.doc1,
only_with_perms_in=["view_document"],
with_group_users=False,
)
self.assertIn(self.user1, direct_users)
def test_set_permissions_for_objects_raises_for_unknown_action(self) -> None:
"""
GIVEN:
- An unrecognized permission action name with users to grant it
to
WHEN:
- set_permissions_for_objects is called
THEN:
- Permission.DoesNotExist is raised, not a silent no-op
Regression test: the endpoint that calls this
(BulkEditObjectPermissionsView) never actually validates action
names against the raw client-supplied permissions dict --
BulkEditObjectsSerializer._validate_permissions calls
validate_set_permissions() only for its side-effecting user/group id
checks and discards the filtered dict it returns -- so a bogus
action key reaches this function as-is. Resolving the Permission via
a bare `.filter()` (which returns empty instead of raising) would
silently drop the grant and report success.
"""
with self.assertRaises(Permission.DoesNotExist):
set_permissions_for_objects(
{"not_a_real_action": {"users": [self.user1.id], "groups": []}},
Document,
[self.doc1.pk],
)
def test_set_permissions_for_objects_unknown_action_applies_nothing(
self,
) -> None:
"""
GIVEN:
- A permissions dict with a valid action ordered ahead of an
unrecognized one
WHEN:
- set_permissions_for_objects is called
THEN:
- Permission.DoesNotExist is raised
- The valid action ahead of it is not applied either
Every action is resolved before any row is written, so a bad action
name cannot leave a half-applied change behind. That matters because
BulkEditObjectsView turns this exception into a 400: without the
up-front resolution the client would be told the request failed
while the leading action had already been committed.
"""
with self.assertRaises(Permission.DoesNotExist):
set_permissions_for_objects(
{
"view": {"users": [self.user1.id], "groups": []},
"not_a_real_action": {"users": [self.user1.id], "groups": []},
},
Document,
[self.doc1.pk],
)
self.assertNotIn(
self.user1,
get_users_with_perms(
self.doc1,
only_with_perms_in=["view_document"],
with_group_users=False,
),
)
@mock.patch("documents.models.Document.delete") @mock.patch("documents.models.Document.delete")
def test_delete_documents_old_uuid_field(self, m) -> None: def test_delete_documents_old_uuid_field(self, m) -> None:
m.side_effect = Exception("Data too long for column 'transaction_id' at row 1") m.side_effect = Exception("Data too long for column 'transaction_id' at row 1")
+5 -1
View File
@@ -110,7 +110,7 @@ class TestDocument(TestCase):
checksum="checksum", checksum="checksum",
mime_type="application/pdf", mime_type="application/pdf",
) )
Document.objects.create( version = Document.objects.create(
root_document=root, root_document=root,
correspondent=root.correspondent, correspondent=root.correspondent,
title="Version", title="Version",
@@ -124,6 +124,10 @@ class TestDocument(TestCase):
self.assertEqual(Document.objects.count(), 0) self.assertEqual(Document.objects.count(), 0)
self.assertEqual(Document.deleted_objects.count(), 2) self.assertEqual(Document.deleted_objects.count(), 2)
root.restore(strict=False)
self.assertTrue(Document.objects.filter(pk=version.pk).exists())
def test_file_name(self) -> None: def test_file_name(self) -> None:
doc = Document( doc = Document(
mime_type="application/pdf", mime_type="application/pdf",
@@ -136,6 +136,23 @@ def wait_for_mock_call(
return False return False
def sleep_past_stability(
owner: FileStabilityTracker | ConsumerThread,
*,
windows: float = 1.5,
) -> None:
"""
Block until a tracked file's stability window has certainly elapsed.
Args:
owner: The tracker, or the consumer thread running one, whose
configured stability delay sets the wait.
windows: How many stability windows to wait, giving slop for a slow
or loaded test runner.
"""
sleep(owner.stability_delay * windows)
class TestTrackedFile: class TestTrackedFile:
"""Tests for the TrackedFile dataclass.""" """Tests for the TrackedFile dataclass."""
@@ -261,6 +278,56 @@ class TestFileStabilityTracker:
assert len(stable) == 0 assert len(stable) == 0
assert stability_tracker.pending_count == 1 assert stability_tracker.pending_count == 1
def test_get_stable_files_skips_empty_file(
self,
stability_tracker: FileStabilityTracker,
tmp_path: Path,
) -> None:
"""
GIVEN:
- A zero byte file, tracked and past its stability delay
WHEN:
- Stable files are collected
THEN:
- The file is not yielded for consumption
- The file is dropped from tracking rather than held, so an
abandoned placeholder does not keep the watch loop awake
"""
empty = tmp_path / "scan.pdf"
empty.write_bytes(b"")
stability_tracker.track(empty, Change.added)
sleep_past_stability(stability_tracker)
stable = list(stability_tracker.get_stable_files())
assert stable == []
assert stability_tracker.pending_count == 0
def test_empty_file_is_yielded_once_content_arrives(
self,
stability_tracker: FileStabilityTracker,
tmp_path: Path,
) -> None:
"""
GIVEN:
- A zero byte file which was dropped from tracking while empty
WHEN:
- The writer fills the file and a new event re-tracks it
THEN:
- The file is yielded for consumption once it is stable
"""
target = tmp_path / "scan.pdf"
target.write_bytes(b"")
stability_tracker.track(target, Change.added)
sleep_past_stability(stability_tracker)
assert list(stability_tracker.get_stable_files()) == []
target.write_bytes(b"%PDF-1.4 content")
stability_tracker.track(target, Change.modified)
sleep_past_stability(stability_tracker)
assert list(stability_tracker.get_stable_files()) == [target]
def test_get_stable_files_deleted_during_check(self, temp_file: Path) -> None: def test_get_stable_files_deleted_during_check(self, temp_file: Path) -> None:
"""Test deleted file is not returned during stability check.""" """Test deleted file is not returned during stability check."""
tracker = FileStabilityTracker(stability_delay=0.1) tracker = FileStabilityTracker(stability_delay=0.1)
@@ -879,6 +946,51 @@ class TestCommandWatch:
mock_consume_file_delay.apply_async.assert_called() mock_consume_file_delay.apply_async.assert_called()
def test_scanner_placeholder_is_not_consumed_while_empty(
self,
consumption_dir: Path,
sample_pdf: Path,
mock_consume_file_delay: MagicMock,
start_consumer: Callable[..., ConsumerThread],
) -> None:
"""
GIVEN:
- A scanner which creates a zero byte placeholder and only writes
the page some time later (GH discussion #13969)
WHEN:
- The placeholder sits untouched well past the stability delay
- The scanner then writes the real content
THEN:
- The empty placeholder is never queued, as it could only fail
with "Unsupported mime type inode/x-empty"
- The file is queued exactly once, when the content lands
"""
thread = start_consumer(stability_delay=0.2)
target = consumption_dir / "scan.pdf"
target.write_bytes(b"") # the scanner's placeholder
# Well past the stability delay: the old behaviour queued it here.
sleep_past_stability(thread, windows=5)
if thread.exception:
raise thread.exception
assert mock_consume_file_delay.apply_async.call_count == 0
shutil.copy(sample_pdf, target) # the scanner finishes the page
assert wait_for_mock_call(
mock_consume_file_delay.apply_async,
timeout_s=5.0,
)
if thread.exception:
raise thread.exception
assert mock_consume_file_delay.apply_async.call_count == 1
queued_doc = mock_consume_file_delay.apply_async.call_args.kwargs["kwargs"][
"input_doc"
]
assert queued_doc.original_file.name == "scan.pdf"
def test_ignores_macos_files( def test_ignores_macos_files(
self, self,
consumption_dir: Path, consumption_dir: Path,
+33
View File
@@ -32,6 +32,7 @@ from documents.signals.handlers import update_llm_suggestions_cache
from documents.tests.utils import DirectoriesMixin from documents.tests.utils import DirectoriesMixin
from documents.tests.utils import read_streaming_response from documents.tests.utils import read_streaming_response
from paperless.models import ApplicationConfiguration from paperless.models import ApplicationConfiguration
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError from paperless_ai.exceptions import LLMTimeoutError
@@ -737,6 +738,38 @@ class TestAISuggestions(DirectoriesMixin, TestCase):
get_llm_suggestion_cache(self.document.pk, backend="openai-like"), get_llm_suggestion_cache(self.document.pk, backend="openai-like"),
) )
@patch("documents.views.get_ai_document_classification")
@override_settings(
AI_ENABLED=True,
LLM_BACKEND="openai-like",
)
def test_ai_suggestions_with_llm_provider_error(
self,
mock_get_ai_classification,
) -> None:
mock_get_ai_classification.side_effect = LLMProviderError(
"confidential provider response",
)
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 rejected the request. Check logs for details.",
],
},
)
self.assertNotIn("confidential provider response", response.content.decode())
self.assertIsNone(
get_llm_suggestion_cache(self.document.pk, backend="openai-like"),
)
@patch("documents.views.get_ai_document_classification") @patch("documents.views.get_ai_document_classification")
@override_settings( @override_settings(
AI_ENABLED=True, AI_ENABLED=True,
+37 -9
View File
@@ -179,7 +179,7 @@ from documents.permissions import has_perms_owner_aware
from documents.permissions import has_system_status_permission from documents.permissions import has_system_status_permission
from documents.permissions import permitted_document_ids from documents.permissions import permitted_document_ids
from documents.permissions import permitted_object_ids from documents.permissions import permitted_object_ids
from documents.permissions import set_permissions_for_object from documents.permissions import set_permissions_for_objects
from documents.permissions import user_is_unrestricted from documents.permissions import user_is_unrestricted
from documents.plugins.date_parsing import get_date_parser from documents.plugins.date_parsing import get_date_parser
from documents.schema import generate_object_with_permissions_schema from documents.schema import generate_object_with_permissions_schema
@@ -252,6 +252,7 @@ from paperless.views import StandardPagination
from paperless_ai.ai_classifier import get_ai_document_classification from paperless_ai.ai_classifier import get_ai_document_classification
from paperless_ai.ai_classifier import get_llm_output_language from paperless_ai.ai_classifier import get_llm_output_language
from paperless_ai.chat import stream_chat_with_documents from paperless_ai.chat import stream_chat_with_documents
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError from paperless_ai.exceptions import LLMTimeoutError
from paperless_ai.matching import extract_unmatched_names from paperless_ai.matching import extract_unmatched_names
from paperless_ai.matching import match_correspondents_by_name from paperless_ai.matching import match_correspondents_by_name
@@ -1603,6 +1604,22 @@ class DocumentViewSet(
{"ai": [_("AI backend request timed out.")]}, {"ai": [_("AI backend request timed out.")]},
status=status.HTTP_503_SERVICE_UNAVAILABLE, status=status.HTTP_503_SERVICE_UNAVAILABLE,
) )
except LLMProviderError:
logger.exception(
"AI backend rejected the request for document %s",
doc.pk,
)
return Response(
{
"ai": [
_(
"AI backend rejected the request. "
"Check logs for details.",
),
],
},
status=status.HTTP_502_BAD_GATEWAY,
)
set_llm_suggestions_cache( set_llm_suggestions_cache(
doc.pk, doc.pk,
llm_suggestions, llm_suggestions,
@@ -4950,12 +4967,12 @@ class BulkEditObjectsView(PassUserMixin):
qs_owner_update.update(owner=owner) qs_owner_update.update(owner=owner)
if "permissions" in serializer.validated_data: if "permissions" in serializer.validated_data:
for obj in qs: set_permissions_for_objects(
set_permissions_for_object( permissions=permissions,
permissions=permissions, model=object_class,
object=obj, pks=qs.values_list("pk", flat=True),
merge=merge, merge=merge,
) )
except Exception as e: except Exception as e:
logger.warning( logger.warning(
@@ -5437,7 +5454,10 @@ class TrashView(ListModelMixin, PassUserMixin):
model = Document model = Document
queryset = Document.deleted_objects.all() # A version is listed separately only when its root is not in the trash.
queryset = Document.deleted_objects.exclude(
root_document_id__in=Document.deleted_objects.values("id"),
)
def get(self, request: Request, format: str | None = None) -> Response: def get(self, request: Request, format: str | None = None) -> Response:
self.serializer_class = DocumentSerializer self.serializer_class = DocumentSerializer
@@ -5468,7 +5488,15 @@ class TrashView(ListModelMixin, PassUserMixin):
return HttpResponseForbidden("Insufficient permissions") return HttpResponseForbidden("Insufficient permissions")
action = serializer.validated_data.get("action") action = serializer.validated_data.get("action")
if action == "restore": if action == "restore":
restored = list(Document.deleted_objects.filter(id__in=doc_ids)) restored = list(self.get_queryset().filter(id__in=doc_ids))
if len(restored) != len(doc_ids):
raise ValidationError(
{
"documents": [
"Restore the root document instead of one of its versions.",
],
},
)
for doc in restored: for doc in restored:
doc.restore(strict=False) doc.restore(strict=False)
if restored: if restored:
File diff suppressed because it is too large Load Diff
+97 -47
View File
@@ -4,21 +4,24 @@ from django.conf import settings
from django.contrib.auth.models import User from django.contrib.auth.models import User
from documents.models import Document from documents.models import Document
from documents.permissions import get_objects_for_user_owner_aware from documents.permissions import permitted_object_ids
from documents.permissions import restrict_queryset_to_visible
from documents.permissions import user_is_unrestricted
from paperless.config import AIConfig from paperless.config import AIConfig
from paperless_ai.base_model import ClassificationSuggestions from paperless_ai.base_model import ClassificationSuggestions
from paperless_ai.base_model import TaxonomyChoiceDict from paperless_ai.base_model import TaxonomyChoiceDict
from paperless_ai.base_model import classification_suggestions_to_model from paperless_ai.base_model import classification_suggestions_to_model
from paperless_ai.client import AIClient from paperless_ai.client import AIClient
from paperless_ai.db import db_connection_released from paperless_ai.db import db_connection_released
from paperless_ai.indexing import _node_document_ids
from paperless_ai.indexing import retrieve_similar_nodes from paperless_ai.indexing import retrieve_similar_nodes
from paperless_ai.indexing import truncate_content from paperless_ai.indexing import truncate_content
from paperless_ai.prompts.context import ClassificationPromptContext from paperless_ai.prompts.context import ClassificationPromptContext
from paperless_ai.prompts.context import LocalizationPromptContext from paperless_ai.prompts.context import LocalizationPromptContext
from paperless_ai.prompts.context import RagContextPromptContext from paperless_ai.prompts.context import RagContextPromptContext
from paperless_ai.prompts.render import render_prompt from paperless_ai.prompts.render import render_prompt
from paperless_ai.taxonomy import SimilarDocument
from paperless_ai.taxonomy import TaxonomyCandidates from paperless_ai.taxonomy import TaxonomyCandidates
from paperless_ai.taxonomy import _node_document_weights
from paperless_ai.taxonomy import build_taxonomy_candidates from paperless_ai.taxonomy import build_taxonomy_candidates
from paperless_ai.taxonomy import empty_taxonomy_candidates from paperless_ai.taxonomy import empty_taxonomy_candidates
from paperless_ai.taxonomy import format_taxonomy_for_prompt from paperless_ai.taxonomy import format_taxonomy_for_prompt
@@ -37,6 +40,48 @@ logger = logging.getLogger("paperless_ai.rag_classifier")
TAXONOMY_CANDIDATE_TOP_K = 15 TAXONOMY_CANDIDATE_TOP_K = 15
def _fulltext_similar_documents(
document: Document,
user: User | None,
top_k: int,
) -> list[SimilarDocument]:
"""Rank-based fallback when no embedding backend is configured. Uses
Tantivy's "More Like This" (term-overlap similarity) instead of vector
similarity - cruder, but far better than no candidates at all.
more_like_this_ids returns only a ranked ID list, no scores, so weight is
synthesized from rank (descending from top_k) rather than claiming a
similarity magnitude that doesn't exist. An unrestricted user (none, or an
active superuser - see user_is_unrestricted) is normalized to ``None``
before calling, since the backend's permission filter has no superuser
short-circuit of its own. Results are re-checked with
restrict_queryset_to_visible() since Tantivy's indexed permission fields
lag the DB via async reindexing.
"""
from documents.search import get_backend
unrestricted = user_is_unrestricted(user)
search_user = None if unrestricted else user
backend = get_backend()
similar_ids = backend.more_like_this_ids(
document.pk,
user=search_user,
limit=top_k,
)
if not unrestricted:
allowed_ids = set(
restrict_queryset_to_visible(
Document.objects.filter(pk__in=similar_ids),
user,
"view_document",
).values_list("pk", flat=True),
)
similar_ids = [doc_id for doc_id in similar_ids if doc_id in allowed_ids]
return [
SimilarDocument(document_id=doc_id, weight=float(top_k - rank))
for rank, doc_id in enumerate(similar_ids)
]
def get_language_name(language_code: str) -> str: def get_language_name(language_code: str) -> str:
normalized_language_code = language_code.lower() normalized_language_code = language_code.lower()
for code, name in settings.LANGUAGES: for code, name in settings.LANGUAGES:
@@ -136,43 +181,52 @@ def get_taxonomy_context(
user: User | None = None, user: User | None = None,
max_docs: int = 5, max_docs: int = 5,
) -> tuple[TaxonomyCandidates, str]: ) -> tuple[TaxonomyCandidates, str]:
"""One retrieval feeds both taxonomy candidates and RAG text context. """One retrieval feeds both taxonomy candidates and RAG text context. Uses
On any retrieval failure, degrades to empty candidates/context rather than vector similarity when an embedding backend is configured, otherwise
propagating the exception - a vector-store outage should not block falls back to Tantivy full-text "More Like This" similarity - see
classification, only its RAG-assisted enrichment. _fulltext_similar_documents. On any retrieval failure, degrades to empty
candidates/context rather than propagating the exception - neither a
vector-store outage nor a search-index issue should block classification,
only its context-assisted enrichment.
""" """
ai_config = AIConfig()
try: try:
# None means "no restriction" to retrieve_similar_nodes. A superuser if ai_config.llm_embedding_backend:
# (like no user at all) can see every document, so skip materializing # None means "no restriction" to retrieve_similar_nodes. An
# every visible pk into a Python list and passing it through as an IN # unrestricted user (no user at all, or an active superuser -- see
# filter: for a large library that is a wasted quadratic scan in the # user_is_unrestricted) can see every document, so skip
# vector store at best, and past ~32,763 documents a hard # materializing every visible pk into a Python list and passing it
# sqlite3.OperationalError (SQLite's bound-parameter limit) at worst. # through as an IN filter: for a large library that is a wasted
# get_objects_for_user_owner_aware() would return every Document for a # quadratic scan in the vector store at best, and past ~32,763
# superuser anyway (guardian's own with_superuser shortcut), so this # documents a hard sqlite3.OperationalError (SQLite's
# changes nothing about which documents are considered -- only how we # bound-parameter limit) at worst.
# get there. # permitted_object_ids() has its own superuser shortcut that would
visible_document_ids = ( # return every Document's id anyway, so this changes nothing about
None # which documents are considered -- only how we get there.
if user is None or user.is_superuser visible_document_ids = (
else list( None
get_objects_for_user_owner_aware( if user_is_unrestricted(user)
user, else list(permitted_object_ids(user, Document, "view_document"))
"view_document", )
Document, nodes = retrieve_similar_nodes(
).values_list("pk", flat=True), document,
top_k=TAXONOMY_CANDIDATE_TOP_K,
document_ids=visible_document_ids,
)
similar_documents = _node_document_weights(nodes)
else:
# See _fulltext_similar_documents: it applies its own permission
# filter via `user`, so no visible-document-id list is needed here.
similar_documents = _fulltext_similar_documents(
document,
user,
top_k=TAXONOMY_CANDIDATE_TOP_K,
) )
)
nodes = retrieve_similar_nodes(
document,
top_k=TAXONOMY_CANDIDATE_TOP_K,
document_ids=visible_document_ids,
)
candidates = build_taxonomy_candidates(nodes, user) candidates = build_taxonomy_candidates(similar_documents, user)
# ``nodes`` are already ordered by descending vector similarity; don't lose it. # similar_documents is already ordered by descending weight; don't lose it.
similar_document_ids = list(dict.fromkeys(_node_document_ids(nodes))) similar_document_ids = [s["document_id"] for s in similar_documents]
similar_documents_by_id = Document.objects.in_bulk(similar_document_ids) similar_documents_by_id = Document.objects.in_bulk(similar_document_ids)
similar_docs = [ similar_docs = [
similar_documents_by_id[document_id] similar_documents_by_id[document_id]
@@ -186,8 +240,8 @@ def get_taxonomy_context(
context_blocks.append(f"TITLE: {title}\n{text}") context_blocks.append(f"TITLE: {title}\n{text}")
except Exception: except Exception:
logger.exception( logger.exception(
"Failed to retrieve RAG neighbours for document %s; continuing " "Failed to retrieve similar-document context for document %s; "
"without taxonomy candidates or similar-document context.", "continuing without taxonomy candidates or similar-document context.",
document.pk, document.pk,
) )
return empty_taxonomy_candidates(), "" return empty_taxonomy_candidates(), ""
@@ -241,17 +295,13 @@ def get_ai_document_classification(
) -> ClassificationSuggestions: ) -> ClassificationSuggestions:
ai_config = AIConfig() ai_config = AIConfig()
if ai_config.llm_embedding_backend: candidates, context = get_taxonomy_context(document, user)
candidates, context = get_taxonomy_context(document, user) prompt = build_prompt_with_rag(
prompt = build_prompt_with_rag( document,
document, ai_config,
ai_config, candidates=candidates,
candidates=candidates, context=context,
context=context, )
)
else:
candidates = empty_taxonomy_candidates()
prompt = build_prompt_without_rag(document, ai_config, candidates=candidates)
client = AIClient() client = AIClient()
# Hand the pooled DB connection back while the (slow) LLM query runs so it # Hand the pooled DB connection back while the (slow) LLM query runs so it
+19 -3
View File
@@ -22,6 +22,7 @@ from paperless.network import validate_outbound_http_url
from paperless_ai.base_model import ClassificationSuggestions from paperless_ai.base_model import ClassificationSuggestions
from paperless_ai.base_model import DocumentClassifierSchema from paperless_ai.base_model import DocumentClassifierSchema
from paperless_ai.base_model import model_to_classification_suggestions from paperless_ai.base_model import model_to_classification_suggestions
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError from paperless_ai.exceptions import LLMTimeoutError
logger = logging.getLogger("paperless_ai.client") logger = logging.getLogger("paperless_ai.client")
@@ -132,7 +133,7 @@ class AIClient:
from llama_index.core.llms import ChatMessage from llama_index.core.llms import ChatMessage
if self.settings.llm_backend == LLMBackend.OLLAMA: if self.settings.llm_backend == LLMBackend.OLLAMA:
with self._normalize_timeouts(): with self._normalize_errors():
result = self.llm.chat( result = self.llm.chat(
[ChatMessage(role="user", content=prompt)], [ChatMessage(role="user", content=prompt)],
format=DocumentClassifierSchema.model_json_schema(), format=DocumentClassifierSchema.model_json_schema(),
@@ -153,7 +154,7 @@ class AIClient:
content=f"{prompt}\n\n" content=f"{prompt}\n\n"
f"Answer by calling the {tool.metadata.name} tool. Do not write the answer as text.", f"Answer by calling the {tool.metadata.name} tool. Do not write the answer as text.",
) )
with self._normalize_timeouts(): with self._normalize_errors():
result = self.llm.chat_with_tools( result = self.llm.chat_with_tools(
tools=[tool], tools=[tool],
user_msg=user_msg, user_msg=user_msg,
@@ -173,7 +174,7 @@ class AIClient:
) )
@contextmanager @contextmanager
def _normalize_timeouts(self) -> Iterator[None]: def _normalize_errors(self) -> Iterator[None]:
try: try:
yield yield
except httpx.TimeoutException as exc: except httpx.TimeoutException as exc:
@@ -181,8 +182,23 @@ class AIClient:
except Exception as exc: except Exception as exc:
if self._is_openai_timeout(exc): if self._is_openai_timeout(exc):
raise LLMTimeoutError from exc raise LLMTimeoutError from exc
if self._is_provider_error(exc):
raise LLMProviderError from exc
raise raise
def _is_provider_error(self, exc: Exception) -> bool:
if self.settings.llm_backend == LLMBackend.OLLAMA:
from ollama import ResponseError
return isinstance(exc, ResponseError)
if self.settings.llm_backend == LLMBackend.OPENAI_LIKE:
from openai import APIStatusError
return isinstance(exc, APIStatusError)
return False
def _is_openai_timeout(self, exc: Exception) -> bool: def _is_openai_timeout(self, exc: Exception) -> bool:
if self.settings.llm_backend != LLMBackend.OPENAI_LIKE: if self.settings.llm_backend != LLMBackend.OPENAI_LIKE:
return False return False
+4
View File
@@ -1,2 +1,6 @@
class LLMTimeoutError(Exception): class LLMTimeoutError(Exception):
pass pass
class LLMProviderError(Exception):
"""The LLM backend rejected the request."""
-17
View File
@@ -721,20 +721,3 @@ def retrieve_similar_nodes(
continue continue
filtered.append(node) filtered.append(node)
return filtered return filtered
def _node_document_ids(nodes: list["NodeWithScore"]) -> list[int]:
document_ids: list[int] = []
for node in nodes:
document_id = node.metadata.get("document_id")
if document_id is None: # pragma: no cover
# See the matching guard in retrieve_similar_nodes() above.
continue
try:
document_ids.append(int(document_id))
except ValueError: # pragma: no cover
logger.warning(
"Skipping LLM index result with invalid document_id %r.",
document_id,
)
return document_ids
+31 -14
View File
@@ -31,6 +31,11 @@ class TaxonomyCandidate(TypedDict):
weight: float weight: float
class SimilarDocument(TypedDict):
document_id: int
weight: float
class TaxonomyCandidates(TypedDict): class TaxonomyCandidates(TypedDict):
tags: list[TaxonomyCandidate] tags: list[TaxonomyCandidate]
document_types: list[TaxonomyCandidate] document_types: list[TaxonomyCandidate]
@@ -49,10 +54,10 @@ def empty_taxonomy_candidates() -> TaxonomyCandidates:
) )
def _node_document_weights(nodes: list["NodeWithScore"]) -> dict[int, float]: def _node_document_weights(nodes: list["NodeWithScore"]) -> list[SimilarDocument]:
"""document_id -> that node's similarity score, summed if a document_id """Sum each node's similarity score into its document_id (a document can
appears more than once across the retrieved nodes (e.g. multiple chunks appear via multiple chunks/nodes) and return one SimilarDocument per
of the same source document).""" distinct document_id."""
weights: dict[int, float] = defaultdict(float) weights: dict[int, float] = defaultdict(float)
for node in nodes: for node in nodes:
document_id = node.metadata.get("document_id") document_id = node.metadata.get("document_id")
@@ -65,7 +70,14 @@ def _node_document_weights(nodes: list["NodeWithScore"]) -> dict[int, float]:
weights[int(document_id)] += float(node.score or 0.0) weights[int(document_id)] += float(node.score or 0.0)
except (TypeError, ValueError): # pragma: no cover except (TypeError, ValueError): # pragma: no cover
continue continue
return weights return sorted(
(
SimilarDocument(document_id=document_id, weight=weight)
for document_id, weight in weights.items()
),
key=lambda similar: similar["weight"],
reverse=True,
)
def _visible_ranked_candidates( def _visible_ranked_candidates(
@@ -101,21 +113,26 @@ def _visible_ranked_candidates(
def build_taxonomy_candidates( def build_taxonomy_candidates(
nodes: list["NodeWithScore"], similar_documents: list[SimilarDocument],
user: User | None, user: User | None,
) -> TaxonomyCandidates: ) -> TaxonomyCandidates:
"""Resolve each neighbour node's document_id to a live Document, read its """Resolve each similar document's id to a live Document, read its
*current* tags/type/correspondent/storage_path via the ORM (never the *current* tags/type/correspondent/storage_path via the ORM (never any
possibly-stale names cached in vector-index node metadata), weight each possibly-stale names an adapter's source might have cached), weight each
distinct taxonomy object by aggregate neighbour similarity, permission-filter distinct taxonomy object by aggregate similarity weight, permission-filter
against what ``user`` can see, and return each category ranked by weight against what ``user`` can see, and return each category ranked by weight
and capped. and capped. ``similar_documents`` may come from either the vector-RAG
adapter or the full-text fallback adapter - both produce this same shape.
""" """
if not similar_documents:
document_weights = _node_document_weights(nodes)
if not document_weights:
return empty_taxonomy_candidates() return empty_taxonomy_candidates()
# Both adapters guarantee at most one SimilarDocument per document_id, so
# this never silently drops a duplicate's weight.
document_weights: dict[int, float] = {
s["document_id"]: s["weight"] for s in similar_documents
}
# Only .tags.all() needs prefetching (a reverse M2M, one extra query for # Only .tags.all() needs prefetching (a reverse M2M, one extra query for
# the whole batch). document_type/correspondent/storage_path are read # the whole batch). document_type/correspondent/storage_path are read
# below via their *_id columns (neighbour.document_type_id, etc.), which # below via their *_id columns (neighbour.document_type_id, etc.), which
+306 -22
View File
@@ -1,4 +1,5 @@
import datetime import datetime
from collections.abc import Generator
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock
from unittest.mock import patch from unittest.mock import patch
@@ -6,18 +7,24 @@ from unittest.mock import patch
import pytest import pytest
import pytest_mock import pytest_mock
from django.test import override_settings from django.test import override_settings
from guardian.shortcuts import assign_perm
from guardian.shortcuts import remove_perm
from documents.models import Document from documents.models import Document
from documents.search import TantivyBackend
from documents.tests.factories import DocumentFactory from documents.tests.factories import DocumentFactory
from documents.tests.factories import TagFactory from documents.tests.factories import TagFactory
from documents.tests.factories import UserFactory from documents.tests.factories import UserFactory
from paperless.config import AIConfig from paperless.config import AIConfig
from paperless_ai.ai_classifier import TAXONOMY_CANDIDATE_TOP_K
from paperless_ai.ai_classifier import _fulltext_similar_documents
from paperless_ai.ai_classifier import build_localization_prompt from paperless_ai.ai_classifier import build_localization_prompt
from paperless_ai.ai_classifier import build_prompt_with_rag from paperless_ai.ai_classifier import build_prompt_with_rag
from paperless_ai.ai_classifier import build_prompt_without_rag from paperless_ai.ai_classifier import build_prompt_without_rag
from paperless_ai.ai_classifier import get_ai_document_classification from paperless_ai.ai_classifier import get_ai_document_classification
from paperless_ai.ai_classifier import get_language_name from paperless_ai.ai_classifier import get_language_name
from paperless_ai.ai_classifier import get_taxonomy_context from paperless_ai.ai_classifier import get_taxonomy_context
from paperless_ai.taxonomy import SimilarDocument
from paperless_ai.taxonomy import TaxonomyCandidate from paperless_ai.taxonomy import TaxonomyCandidate
from paperless_ai.taxonomy import TaxonomyCandidates from paperless_ai.taxonomy import TaxonomyCandidates
@@ -220,12 +227,10 @@ def test_use_rag_if_configured(
@pytest.mark.django_db @pytest.mark.django_db
@patch("paperless_ai.client.AIClient.run_llm_query") @patch("paperless_ai.client.AIClient.run_llm_query")
@patch("paperless_ai.ai_classifier.build_prompt_without_rag") @patch("paperless_ai.ai_classifier.build_prompt_with_rag")
@patch("paperless_ai.ai_classifier.AIConfig")
@override_settings(LLM_BACKEND="ollama", LLM_MODEL="some_model") @override_settings(LLM_BACKEND="ollama", LLM_MODEL="some_model")
def test_use_without_rag_if_not_configured( def test_use_rag_prompt_even_without_embedding_backend(
mock_ai_config, mock_build_prompt_with_rag,
mock_build_prompt_without_rag,
mock_run_llm_query, mock_run_llm_query,
mock_document, mock_document,
): ):
@@ -235,13 +240,13 @@ def test_use_without_rag_if_not_configured(
WHEN: WHEN:
- get_ai_document_classification() is called - get_ai_document_classification() is called
THEN: THEN:
- The non-RAG prompt builder is used - The RAG-context prompt builder is still used (fed by the full-text
fallback's context/candidates instead of the vector store's)
""" """
mock_ai_config.return_value.llm_embedding_backend = None mock_build_prompt_with_rag.return_value = "Prompt with RAG"
mock_build_prompt_without_rag.return_value = "Prompt without RAG"
mock_run_llm_query.return_value = NESTED_SUGGESTIONS mock_run_llm_query.return_value = NESTED_SUGGESTIONS
get_ai_document_classification(mock_document) get_ai_document_classification(mock_document)
mock_build_prompt_without_rag.assert_called_once() mock_build_prompt_with_rag.assert_called_once()
@pytest.mark.django_db @pytest.mark.django_db
@@ -320,6 +325,7 @@ def test_build_localization_prompt_preserves_unicode_characters():
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_get_taxonomy_context_assembles_rag_text_and_candidates(): def test_get_taxonomy_context_assembles_rag_text_and_candidates():
""" """
GIVEN: GIVEN:
@@ -354,6 +360,7 @@ def test_get_taxonomy_context_assembles_rag_text_and_candidates():
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_get_taxonomy_context_preserves_similarity_order_and_distinct_documents(): def test_get_taxonomy_context_preserves_similarity_order_and_distinct_documents():
""" """
GIVEN: GIVEN:
@@ -424,6 +431,7 @@ def test_get_taxonomy_context_preserves_similarity_order_and_distinct_documents(
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_get_taxonomy_context_no_similar_docs(): def test_get_taxonomy_context_no_similar_docs():
""" """
GIVEN: GIVEN:
@@ -447,6 +455,67 @@ def test_get_taxonomy_context_no_similar_docs():
} }
@pytest.mark.django_db
def test_get_taxonomy_context_uses_fulltext_fallback_when_no_embedding_backend(
mocker: pytest_mock.MockerFixture,
) -> None:
"""
GIVEN:
- No LLM embedding backend is configured (the default test settings)
WHEN:
- get_taxonomy_context() is called
THEN:
- _fulltext_similar_documents() is called with the document, the user
and TAXONOMY_CANDIDATE_TOP_K
- retrieve_similar_nodes() (the vector path) is never called
"""
document = DocumentFactory.create(content="Some content")
mock_fulltext = mocker.patch(
"paperless_ai.ai_classifier._fulltext_similar_documents",
return_value=[],
)
mock_retrieve = mocker.patch("paperless_ai.ai_classifier.retrieve_similar_nodes")
get_taxonomy_context(document, user=None)
mock_fulltext.assert_called_once_with(
document,
None,
top_k=TAXONOMY_CANDIDATE_TOP_K,
)
mock_retrieve.assert_not_called()
@pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_get_taxonomy_context_uses_vector_path_when_embedding_backend_configured(
mocker: pytest_mock.MockerFixture,
) -> None:
"""
GIVEN:
- An LLM embedding backend is configured
WHEN:
- get_taxonomy_context() is called
THEN:
- retrieve_similar_nodes() (the vector path) is called
- _fulltext_similar_documents() (the no-embedding-backend fallback)
is never called
"""
document = DocumentFactory.create(content="Some content")
mock_retrieve = mocker.patch(
"paperless_ai.ai_classifier.retrieve_similar_nodes",
return_value=[],
)
mock_fulltext = mocker.patch(
"paperless_ai.ai_classifier._fulltext_similar_documents",
)
get_taxonomy_context(document, user=None)
mock_retrieve.assert_called_once()
mock_fulltext.assert_not_called()
class TestGetTaxonomyContextVisibility: class TestGetTaxonomyContextVisibility:
"""get_taxonomy_context must not materialize every visible document id """get_taxonomy_context must not materialize every visible document id
for a user who can already see the whole library: a superuser (like no for a user who can already see the whole library: a superuser (like no
@@ -459,6 +528,7 @@ class TestGetTaxonomyContextVisibility:
""" """
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_skips_permission_lookup_for_superuser( def test_skips_permission_lookup_for_superuser(
self, self,
mocker: pytest_mock.MockerFixture, mocker: pytest_mock.MockerFixture,
@@ -477,17 +547,18 @@ class TestGetTaxonomyContextVisibility:
"paperless_ai.ai_classifier.retrieve_similar_nodes", "paperless_ai.ai_classifier.retrieve_similar_nodes",
return_value=[], return_value=[],
) )
mock_get_objects = mocker.patch( mock_permitted = mocker.patch(
"paperless_ai.ai_classifier.get_objects_for_user_owner_aware", "paperless_ai.ai_classifier.permitted_object_ids",
) )
user = UserFactory.create(is_superuser=True) user = UserFactory.create(is_superuser=True)
get_taxonomy_context(document, user) get_taxonomy_context(document, user)
mock_get_objects.assert_not_called() mock_permitted.assert_not_called()
assert mock_retrieve.call_args.kwargs["document_ids"] is None assert mock_retrieve.call_args.kwargs["document_ids"] is None
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_skips_permission_lookup_when_no_user( def test_skips_permission_lookup_when_no_user(
self, self,
mocker: pytest_mock.MockerFixture, mocker: pytest_mock.MockerFixture,
@@ -506,16 +577,17 @@ class TestGetTaxonomyContextVisibility:
"paperless_ai.ai_classifier.retrieve_similar_nodes", "paperless_ai.ai_classifier.retrieve_similar_nodes",
return_value=[], return_value=[],
) )
mock_get_objects = mocker.patch( mock_permitted = mocker.patch(
"paperless_ai.ai_classifier.get_objects_for_user_owner_aware", "paperless_ai.ai_classifier.permitted_object_ids",
) )
get_taxonomy_context(document, None) get_taxonomy_context(document, None)
mock_get_objects.assert_not_called() mock_permitted.assert_not_called()
assert mock_retrieve.call_args.kwargs["document_ids"] is None assert mock_retrieve.call_args.kwargs["document_ids"] is None
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_restricts_to_visible_documents_for_non_superuser( def test_restricts_to_visible_documents_for_non_superuser(
self, self,
mocker: pytest_mock.MockerFixture, mocker: pytest_mock.MockerFixture,
@@ -526,7 +598,7 @@ class TestGetTaxonomyContextVisibility:
WHEN: WHEN:
- get_taxonomy_context() is called - get_taxonomy_context() is called
THEN: THEN:
- The user's visible document ids are looked up and passed to - The user's permitted document ids are looked up and passed to
retrieve_similar_nodes() as a restriction retrieve_similar_nodes() as a restriction
""" """
document = DocumentFactory.create(content="Some content") document = DocumentFactory.create(content="Some content")
@@ -534,21 +606,232 @@ class TestGetTaxonomyContextVisibility:
"paperless_ai.ai_classifier.retrieve_similar_nodes", "paperless_ai.ai_classifier.retrieve_similar_nodes",
return_value=[], return_value=[],
) )
mock_queryset = mocker.MagicMock() mock_permitted = mocker.patch(
mock_queryset.values_list.return_value = [1, 2, 3] "paperless_ai.ai_classifier.permitted_object_ids",
mock_get_objects = mocker.patch( return_value=[1, 2, 3],
"paperless_ai.ai_classifier.get_objects_for_user_owner_aware",
return_value=mock_queryset,
) )
user = UserFactory.create(is_superuser=False) user = UserFactory.create(is_superuser=False)
get_taxonomy_context(document, user) get_taxonomy_context(document, user)
mock_get_objects.assert_called_once_with(user, "view_document", Document) mock_permitted.assert_called_once_with(user, Document, "view_document")
assert mock_retrieve.call_args.kwargs["document_ids"] == [1, 2, 3] assert mock_retrieve.call_args.kwargs["document_ids"] == [1, 2, 3]
@pytest.mark.django_db @pytest.mark.django_db
class TestFulltextSimilarDocuments:
"""_fulltext_similar_documents is the no-embedding-backend fallback: it
asks the Tantivy full-text index for "More Like This" neighbours instead
of the vector store, and synthesizes a rank-based weight since Tantivy's
more_like_this_ids returns only an ordered id list, no scores.
"""
@pytest.fixture
def fulltext_backend(
self,
mocker: pytest_mock.MockerFixture,
) -> Generator[TantivyBackend, None, None]:
"""An in-memory Tantivy backend, wired up as the module-level
singleton _fulltext_similar_documents resolves via get_backend()."""
backend = TantivyBackend(path=None)
backend.open()
mocker.patch("documents.search.get_backend", return_value=backend)
try:
yield backend
finally:
backend.close()
def test_ranks_by_rank_based_weight_descending(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A source document and two similar documents indexed in Tantivy
WHEN:
- _fulltext_similar_documents() is called
THEN:
- Each result's weight reflects its rank (first result weighted
higher than the second), not a raw similarity score
"""
source = DocumentFactory.create(content="quarterly financial report details")
first = DocumentFactory.create(content="quarterly financial report details")
second = DocumentFactory.create(content="financial report")
for doc in (source, first, second):
fulltext_backend.add_or_update(doc)
result = _fulltext_similar_documents(source, user=None, top_k=5)
assert len(result) == 2
weight_by_id = {s["document_id"]: s["weight"] for s in result}
assert weight_by_id[first.pk] > weight_by_id[second.pk]
def test_excludes_source_document(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A source document indexed in Tantivy with no other documents
WHEN:
- _fulltext_similar_documents() is called
THEN:
- An empty list is returned - the source document is never its
own similar document
"""
source = DocumentFactory.create(content="unique unrelated content")
fulltext_backend.add_or_update(source)
result = _fulltext_similar_documents(source, user=None, top_k=5)
assert result == []
def test_empty_index_returns_empty_list(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A document that has never been indexed (fresh/empty Tantivy index)
WHEN:
- _fulltext_similar_documents() is called
THEN:
- An empty list is returned rather than raising
"""
source = DocumentFactory.create(content="never indexed")
result = _fulltext_similar_documents(source, user=None, top_k=5)
assert result == []
def test_respects_top_k_limit(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A source document and four similar documents indexed
WHEN:
- _fulltext_similar_documents() is called with top_k=2
THEN:
- At most 2 results are returned
"""
source = DocumentFactory.create(content="shared overlapping keyword text")
fulltext_backend.add_or_update(source)
for _ in range(4):
fulltext_backend.add_or_update(
DocumentFactory.create(content="shared overlapping keyword text"),
)
result = _fulltext_similar_documents(source, user=None, top_k=2)
assert len(result) == 2
def test_result_shape_is_similar_document(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A source document and one similar document indexed
WHEN:
- _fulltext_similar_documents() is called
THEN:
- Each result is a SimilarDocument (document_id + weight only)
"""
source = DocumentFactory.create(content="shared content phrase")
other = DocumentFactory.create(content="shared content phrase")
fulltext_backend.add_or_update(source)
fulltext_backend.add_or_update(other)
result = _fulltext_similar_documents(source, user=None, top_k=5)
# rank 0 (the only/best result) with top_k=5 -> weight = top_k - rank = 5.0,
# per the "first result gets top_k, the last gets 1" formula.
assert result == [SimilarDocument(document_id=other.pk, weight=5.0)]
def test_superuser_sees_other_users_documents(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A source document owned by one user and a similar document
owned by a different user, with no sharing between them
WHEN:
- _fulltext_similar_documents() is called with a superuser
THEN:
- The other user's document is still returned as a similar
document - a superuser must not be narrowed by the backend's
owner-based permission filter
"""
owner = UserFactory.create()
other_owner = UserFactory.create()
superuser = UserFactory.create(is_superuser=True)
source = DocumentFactory.create(
content="shared content phrase",
owner=owner,
)
other = DocumentFactory.create(
content="shared content phrase",
owner=other_owner,
)
fulltext_backend.add_or_update(source)
fulltext_backend.add_or_update(other)
result = _fulltext_similar_documents(source, user=superuser, top_k=5)
assert [s["document_id"] for s in result] == [other.pk]
def test_excludes_stale_permitted_document_for_regular_user(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A regular (non-superuser) user
- A similar document the user is permitted to view, and another
similar document indexed while the user still had view
permission but which has since had that permission revoked in
the database, i.e. the Tantivy index has stale permission data
WHEN:
- _fulltext_similar_documents() is called with that user
THEN:
- Only the still-permitted document is returned - the DB
re-check via restrict_queryset_to_visible() must catch the
document Tantivy's stale index still thinks is visible
"""
owner = UserFactory.create()
viewer = UserFactory.create(is_superuser=False)
source = DocumentFactory.create(
content="shared content phrase",
owner=owner,
)
permitted = DocumentFactory.create(
content="shared content phrase",
owner=owner,
)
now_private = DocumentFactory.create(
content="shared content phrase",
owner=owner,
)
assign_perm("view_document", viewer, permitted)
assign_perm("view_document", viewer, now_private)
fulltext_backend.add_or_update(source)
fulltext_backend.add_or_update(permitted)
fulltext_backend.add_or_update(now_private)
# Revoke access after indexing, without reindexing: the index still
# carries viewer as a permitted viewer for `now_private`.
remove_perm("view_document", viewer, now_private)
result = _fulltext_similar_documents(source, user=viewer, top_k=5)
assert [s["document_id"] for s in result] == [permitted.pk]
@pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
@patch("paperless_ai.ai_classifier.retrieve_similar_nodes") @patch("paperless_ai.ai_classifier.retrieve_similar_nodes")
def test_get_taxonomy_context_retrieval_failure_degrades_to_no_hints(mock_retrieve): def test_get_taxonomy_context_retrieval_failure_degrades_to_no_hints(mock_retrieve):
""" """
@@ -575,6 +858,7 @@ def test_get_taxonomy_context_retrieval_failure_degrades_to_no_hints(mock_retrie
@pytest.mark.django_db @pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
@patch("paperless_ai.ai_classifier.build_taxonomy_candidates") @patch("paperless_ai.ai_classifier.build_taxonomy_candidates")
@patch("paperless_ai.ai_classifier.retrieve_similar_nodes") @patch("paperless_ai.ai_classifier.retrieve_similar_nodes")
def test_get_taxonomy_context_candidate_building_failure_degrades_to_no_hints( def test_get_taxonomy_context_candidate_building_failure_degrades_to_no_hints(
+3 -5
View File
@@ -1188,9 +1188,7 @@ class TestRetrieveSimilarNodesAgainstRealIndex:
nodes = indexing.retrieve_similar_nodes(a, document_ids=[b.id]) nodes = indexing.retrieve_similar_nodes(a, document_ids=[b.id])
assert all( assert all(int(node.metadata["document_id"]) == b.id for node in nodes)
document_id == b.id for document_id in indexing._node_document_ids(nodes)
)
def test_excludes_self( def test_excludes_self(
self, self,
@@ -1212,7 +1210,7 @@ class TestRetrieveSimilarNodesAgainstRealIndex:
nodes = indexing.retrieve_similar_nodes(a, top_k=5) nodes = indexing.retrieve_similar_nodes(a, top_k=5)
assert set(indexing._node_document_ids(nodes)) == {b.id} assert {int(node.metadata["document_id"]) for node in nodes} == {b.id}
def test_excludes_self_with_multiple_chunks( def test_excludes_self_with_multiple_chunks(
self, self,
@@ -1235,4 +1233,4 @@ class TestRetrieveSimilarNodesAgainstRealIndex:
nodes = indexing.retrieve_similar_nodes(a, top_k=3) nodes = indexing.retrieve_similar_nodes(a, top_k=3)
assert set(indexing._node_document_ids(nodes)) == {b.id} assert {int(node.metadata["document_id"]) for node in nodes} == {b.id}
+48
View File
@@ -4,6 +4,7 @@ from unittest.mock import MagicMock
from unittest.mock import patch from unittest.mock import patch
import httpx import httpx
import ollama
import openai import openai
import pytest import pytest
from llama_index.core.llms.llm import ToolSelection from llama_index.core.llms.llm import ToolSelection
@@ -11,6 +12,7 @@ from llama_index.core.llms.llm import ToolSelection
from paperless_ai.client import LLM_SYSTEM_PROMPT from paperless_ai.client import LLM_SYSTEM_PROMPT
from paperless_ai.client import PLACEHOLDER_API_KEY from paperless_ai.client import PLACEHOLDER_API_KEY
from paperless_ai.client import AIClient from paperless_ai.client import AIClient
from paperless_ai.exceptions import LLMProviderError
from paperless_ai.exceptions import LLMTimeoutError from paperless_ai.exceptions import LLMTimeoutError
@@ -214,6 +216,52 @@ def test_run_llm_query_openai_timeout_raises_local_error(
client.run_llm_query("test_prompt") client.run_llm_query("test_prompt")
def test_run_llm_query_openai_status_error_raises_provider_error(
mock_ai_config,
mock_openai_llm,
):
mock_ai_config.llm_backend = "openai-like"
mock_ai_config.llm_model = "test_model"
mock_ai_config.llm_endpoint = "http://test-url"
request = httpx.Request("POST", "http://test-url/v1/chat/completions")
body = {"error": {"message": "Thinking mode does not support this tool_choice"}}
mock_openai_llm.return_value.chat_with_tools.side_effect = openai.BadRequestError(
"Error code: 400",
response=httpx.Response(400, request=request, json=body),
body=body,
)
client = AIClient()
with pytest.raises(LLMProviderError) as exc_info:
client.run_llm_query("test_prompt")
assert str(exc_info.value) == ""
assert isinstance(exc_info.value.__cause__, openai.BadRequestError)
def test_run_llm_query_ollama_response_error_raises_provider_error(
mock_ai_config,
mock_ollama_llm,
):
mock_ai_config.llm_backend = "ollama"
mock_ai_config.llm_model = "test_model"
mock_ai_config.llm_endpoint = "http://test-url"
response_error = ollama.ResponseError(
"confidential provider response",
status_code=400,
)
mock_ollama_llm.return_value.chat.side_effect = response_error
client = AIClient()
with pytest.raises(LLMProviderError) as exc_info:
client.run_llm_query("test_prompt")
assert str(exc_info.value) == ""
assert exc_info.value.__cause__ is response_error
def test_run_llm_query_httpx_timeout_raises_local_error( def test_run_llm_query_httpx_timeout_raises_local_error(
mock_ai_config, mock_ai_config,
mock_ollama_llm, mock_ollama_llm,
+33 -31
View File
@@ -1,5 +1,4 @@
import json import json
from types import SimpleNamespace
import pytest import pytest
import pytest_mock import pytest_mock
@@ -10,14 +9,14 @@ from documents.tests.factories import DocumentTypeFactory
from documents.tests.factories import StoragePathFactory from documents.tests.factories import StoragePathFactory
from documents.tests.factories import TagFactory from documents.tests.factories import TagFactory
from documents.tests.factories import UserFactory from documents.tests.factories import UserFactory
from paperless_ai.taxonomy import SimilarDocument
from paperless_ai.taxonomy import TaxonomyCandidates from paperless_ai.taxonomy import TaxonomyCandidates
from paperless_ai.taxonomy import build_taxonomy_candidates from paperless_ai.taxonomy import build_taxonomy_candidates
from paperless_ai.taxonomy import format_taxonomy_for_prompt from paperless_ai.taxonomy import format_taxonomy_for_prompt
def make_node(document_id: int, score: float) -> SimpleNamespace: def make_similar(document_id: int, weight: float) -> SimilarDocument:
"""A stand-in for NodeWithScore: only ``.metadata``/``.score`` are read.""" return SimilarDocument(document_id=document_id, weight=weight)
return SimpleNamespace(metadata={"document_id": str(document_id)}, score=score)
@pytest.mark.django_db @pytest.mark.django_db
@@ -53,9 +52,9 @@ class TestBuildTaxonomyCandidates:
doc_a.tags.add(tag) doc_a.tags.add(tag)
doc_b = DocumentFactory.create() doc_b = DocumentFactory.create()
doc_b.tags.add(tag) doc_b.tags.add(tag)
nodes = [make_node(doc_a.pk, 0.9), make_node(doc_b.pk, 0.4)] similar_documents = [make_similar(doc_a.pk, 0.9), make_similar(doc_b.pk, 0.4)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["tags"]) == 1 assert len(result["tags"]) == 1
assert result["tags"][0]["id"] == tag.pk assert result["tags"][0]["id"] == tag.pk
@@ -80,9 +79,9 @@ class TestBuildTaxonomyCandidates:
document.tags.add(tag) document.tags.add(tag)
tag.name = "New Name" tag.name = "New Name"
tag.save() tag.save()
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert result["tags"][0]["name"] == "New Name" assert result["tags"][0]["name"] == "New Name"
@@ -102,9 +101,9 @@ class TestBuildTaxonomyCandidates:
document = DocumentFactory.create() document = DocumentFactory.create()
document.tags.add(tag) document.tags.add(tag)
tag.delete() tag.delete()
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert result["tags"] == [] assert result["tags"] == []
@@ -123,9 +122,12 @@ class TestBuildTaxonomyCandidates:
strong_doc.tags.add(strong_tag) strong_doc.tags.add(strong_tag)
weak_doc = DocumentFactory.create() weak_doc = DocumentFactory.create()
weak_doc.tags.add(weak_tag) weak_doc.tags.add(weak_tag)
nodes = [make_node(strong_doc.pk, 0.9), make_node(weak_doc.pk, 0.1)] similar_documents = [
make_similar(strong_doc.pk, 0.9),
make_similar(weak_doc.pk, 0.1),
]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert [c["name"] for c in result["tags"]] == ["Strong", "Weak"] assert [c["name"] for c in result["tags"]] == ["Strong", "Weak"]
@@ -141,9 +143,9 @@ class TestBuildTaxonomyCandidates:
document = DocumentFactory.create() document = DocumentFactory.create()
for i in range(15): for i in range(15):
document.tags.add(TagFactory.create(name=f"Tag{i}")) document.tags.add(TagFactory.create(name=f"Tag{i}"))
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["tags"]) == 10 assert len(result["tags"]) == 10
@@ -157,12 +159,12 @@ class TestBuildTaxonomyCandidates:
- Only 5 correspondents are returned - Only 5 correspondents are returned
""" """
correspondents = CorrespondentFactory.create_batch(7) correspondents = CorrespondentFactory.create_batch(7)
nodes = [ similar_documents = [
make_node(DocumentFactory.create(correspondent=c).pk, 0.5) make_similar(DocumentFactory.create(correspondent=c).pk, 0.5)
for c in correspondents for c in correspondents
] ]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["correspondents"]) == 5 assert len(result["correspondents"]) == 5
@@ -177,9 +179,9 @@ class TestBuildTaxonomyCandidates:
""" """
document_type = DocumentTypeFactory.create(name="Invoice") document_type = DocumentTypeFactory.create(name="Invoice")
document = DocumentFactory.create(document_type=document_type) document = DocumentFactory.create(document_type=document_type)
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["document_types"]) == 1 assert len(result["document_types"]) == 1
assert result["document_types"][0]["id"] == document_type.pk assert result["document_types"][0]["id"] == document_type.pk
@@ -195,12 +197,12 @@ class TestBuildTaxonomyCandidates:
- Only 5 document_types are returned - Only 5 document_types are returned
""" """
document_types = DocumentTypeFactory.create_batch(7) document_types = DocumentTypeFactory.create_batch(7)
nodes = [ similar_documents = [
make_node(DocumentFactory.create(document_type=dt).pk, 0.5) make_similar(DocumentFactory.create(document_type=dt).pk, 0.5)
for dt in document_types for dt in document_types
] ]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["document_types"]) == 5 assert len(result["document_types"]) == 5
@@ -215,9 +217,9 @@ class TestBuildTaxonomyCandidates:
""" """
storage_path = StoragePathFactory.create(name="Invoices") storage_path = StoragePathFactory.create(name="Invoices")
document = DocumentFactory.create(storage_path=storage_path) document = DocumentFactory.create(storage_path=storage_path)
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["storage_paths"]) == 1 assert len(result["storage_paths"]) == 1
assert result["storage_paths"][0]["id"] == storage_path.pk assert result["storage_paths"][0]["id"] == storage_path.pk
@@ -233,12 +235,12 @@ class TestBuildTaxonomyCandidates:
- Only 5 storage_paths are returned - Only 5 storage_paths are returned
""" """
storage_paths = StoragePathFactory.create_batch(7) storage_paths = StoragePathFactory.create_batch(7)
nodes = [ similar_documents = [
make_node(DocumentFactory.create(storage_path=sp).pk, 0.5) make_similar(DocumentFactory.create(storage_path=sp).pk, 0.5)
for sp in storage_paths for sp in storage_paths
] ]
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert len(result["storage_paths"]) == 5 assert len(result["storage_paths"]) == 5
@@ -258,14 +260,14 @@ class TestBuildTaxonomyCandidates:
tag = TagFactory.create(name="Restricted") tag = TagFactory.create(name="Restricted")
document = DocumentFactory.create() document = DocumentFactory.create()
document.tags.add(tag) document.tags.add(tag)
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
user = UserFactory.create() user = UserFactory.create()
mocker.patch( mocker.patch(
"documents.permissions.permitted_object_ids", "documents.permissions.permitted_object_ids",
return_value=[], # user cannot see this tag return_value=[], # user cannot see this tag
) )
result = build_taxonomy_candidates(nodes, user=user) result = build_taxonomy_candidates(similar_documents, user=user)
assert result["tags"] == [] assert result["tags"] == []
@@ -295,10 +297,10 @@ class TestBuildTaxonomyCandidates:
tag.save() tag.save()
document = DocumentFactory.create() document = DocumentFactory.create()
document.tags.add(tag) document.tags.add(tag)
nodes = [make_node(document.pk, 0.5)] similar_documents = [make_similar(document.pk, 0.5)]
spy = mocker.patch("documents.permissions.permitted_object_ids") spy = mocker.patch("documents.permissions.permitted_object_ids")
result = build_taxonomy_candidates(nodes, user=None) result = build_taxonomy_candidates(similar_documents, user=None)
assert result["tags"][0]["name"] == "Owned" assert result["tags"][0]["name"] == "Owned"
spy.assert_not_called() spy.assert_not_called()
+22 -2
View File
@@ -1,7 +1,9 @@
import logging import logging
from celery import Task
from celery import shared_task from celery import shared_task
from documents.models import PaperlessTask
from paperless_mail.mail import MailAccountHandler from paperless_mail.mail import MailAccountHandler
from paperless_mail.mail import MailError from paperless_mail.mail import MailError
from paperless_mail.models import MailAccount from paperless_mail.models import MailAccount
@@ -10,8 +12,26 @@ from paperless_mail.models import MailRule
logger = logging.getLogger("paperless.mail.tasks") logger = logging.getLogger("paperless.mail.tasks")
@shared_task @shared_task(bind=True)
def process_mail_accounts(account_ids: list[int] | None = None) -> str: def process_mail_accounts(self: Task, account_ids: list[int] | None = None) -> str:
# A scheduled check can still be running (or queued) when the next one
# ProcessedMail dedup only records a message once its
# handling has finished, so an overlapping run can still pick up the same
# not-yet-recorded message. Skip outright rather than race it.
other_mail_fetch_running = (
PaperlessTask.objects.filter(
task_type=PaperlessTask.TaskType.MAIL_FETCH,
status__in=[PaperlessTask.Status.PENDING, PaperlessTask.Status.STARTED],
)
.exclude(task_id=self.request.id)
.exists()
)
if other_mail_fetch_running:
logger.info(
"Mail account processing is already running; skipping this run.",
)
return "Skipped: mail account processing already in progress."
total_new_documents = 0 total_new_documents = 0
accounts = ( accounts = (
MailAccount.objects.filter(pk__in=account_ids) MailAccount.objects.filter(pk__in=account_ids)
@@ -0,0 +1,134 @@
from typing import Final
import pytest
import pytest_mock
from documents.models import PaperlessTask
from documents.tests.factories import PaperlessTaskFactory
from paperless_mail import tasks
from paperless_mail.tests.factories import MailAccountFactory
from paperless_mail.tests.factories import MailRuleFactory
NO_DOCUMENTS_ADDED: Final = "No new documents were added."
SKIPPED: Final = "Skipped: mail account processing already in progress."
@pytest.mark.django_db
@pytest.mark.usefixtures("account_with_rule")
class TestProcessMailAccountsOverlap:
@pytest.fixture
def account_with_rule(self) -> None:
"""An enabled mail account with a single enabled rule."""
account = MailAccountFactory.create()
MailRuleFactory.create(account=account, enabled=True)
@pytest.mark.parametrize(
("status", "expected_result", "expected_call_count"),
[
pytest.param(
PaperlessTask.Status.PENDING,
SKIPPED,
0,
id="pending-task-blocks",
),
pytest.param(
PaperlessTask.Status.STARTED,
SKIPPED,
0,
id="started-task-blocks",
),
pytest.param(
PaperlessTask.Status.SUCCESS,
NO_DOCUMENTS_ADDED,
1,
id="finished-task-does-not-block",
),
],
)
def test_skips_only_while_another_mail_fetch_task_runs(
self,
mocker: pytest_mock.MockerFixture,
status: PaperlessTask.Status,
expected_result: str,
expected_call_count: int,
) -> None:
"""
GIVEN:
- An enabled mail account with a rule
- Another mail fetch task row in the given status
WHEN:
- Mail accounts are processed
THEN:
- Processing is skipped only if that other task is pending or running
"""
PaperlessTaskFactory.create(
task_type=PaperlessTask.TaskType.MAIL_FETCH,
trigger_source=PaperlessTask.TriggerSource.SCHEDULED,
status=status,
)
mocked_handle = mocker.patch.object(
tasks.MailAccountHandler,
"handle_mail_account",
return_value=0,
)
result = tasks.process_mail_accounts()
assert mocked_handle.call_count == expected_call_count
assert result == expected_result
def test_runs_when_no_other_mail_fetch_task_exists(
self,
mocker: pytest_mock.MockerFixture,
) -> None:
"""
GIVEN:
- An enabled mail account with a rule
- No other mail fetch task rows
WHEN:
- Mail accounts are processed
THEN:
- The account is handled
"""
mocked_handle = mocker.patch.object(
tasks.MailAccountHandler,
"handle_mail_account",
return_value=0,
)
result = tasks.process_mail_accounts()
mocked_handle.assert_called_once()
assert result == NO_DOCUMENTS_ADDED
def test_does_not_skip_due_to_its_own_task_row(
self,
mocker: pytest_mock.MockerFixture,
) -> None:
"""
GIVEN:
- An enabled mail account with a rule
- A running mail fetch task row belonging to this very task
WHEN:
- Mail accounts are processed under that task id
THEN:
- The task does not skip itself and handles the account
"""
PaperlessTaskFactory.create(
task_id="self-task-id",
task_type=PaperlessTask.TaskType.MAIL_FETCH,
trigger_source=PaperlessTask.TriggerSource.SCHEDULED,
status=PaperlessTask.Status.STARTED,
)
mocked_handle = mocker.patch.object(
tasks.MailAccountHandler,
"handle_mail_account",
return_value=0,
)
result = tasks.process_mail_accounts.apply(task_id="self-task-id").result
mocked_handle.assert_called_once()
assert result == NO_DOCUMENTS_ADDED