mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-10 11:48:00 +00:00
Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0ef2c39c0b | ||
|
|
9e076e7cdf | ||
|
|
e1df61a192 | ||
|
|
41524a6230 | ||
|
|
e5c9b8facb | ||
|
|
cf3a080694 | ||
|
|
018de125ff | ||
|
|
310628699d | ||
|
|
aff0f9cf41 | ||
|
|
bf716ebfd1 | ||
|
|
7d67a10a35 | ||
|
|
8d1bc5dd24 | ||
|
|
43a8d7d412 | ||
|
|
c40922440b |
@@ -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
@@ -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,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],
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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],
|
||||||
|
)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
+34
-6
@@ -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,10 +4967,10 @@ 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,
|
||||||
object=obj,
|
model=object_class,
|
||||||
|
pks=qs.values_list("pk", flat=True),
|
||||||
merge=merge,
|
merge=merge,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -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
@@ -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
|
||||||
|
# return every Document's id anyway, so this changes nothing about
|
||||||
|
# which documents are considered -- only how we get there.
|
||||||
visible_document_ids = (
|
visible_document_ids = (
|
||||||
None
|
None
|
||||||
if user is None or user.is_superuser
|
if user_is_unrestricted(user)
|
||||||
else list(
|
else list(permitted_object_ids(user, Document, "view_document"))
|
||||||
get_objects_for_user_owner_aware(
|
|
||||||
user,
|
|
||||||
"view_document",
|
|
||||||
Document,
|
|
||||||
).values_list("pk", flat=True),
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
nodes = retrieve_similar_nodes(
|
nodes = retrieve_similar_nodes(
|
||||||
document,
|
document,
|
||||||
top_k=TAXONOMY_CANDIDATE_TOP_K,
|
top_k=TAXONOMY_CANDIDATE_TOP_K,
|
||||||
document_ids=visible_document_ids,
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
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,7 +295,6 @@ 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,
|
||||||
@@ -249,9 +302,6 @@ def get_ai_document_classification(
|
|||||||
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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -1,2 +1,6 @@
|
|||||||
class LLMTimeoutError(Exception):
|
class LLMTimeoutError(Exception):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LLMProviderError(Exception):
|
||||||
|
"""The LLM backend rejected the request."""
|
||||||
|
|||||||
@@ -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,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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user