Compare commits

..
Author SHA1 Message Date
shamoon f3bc34c99f Enhancement: skip local 2FA for trusted OIDC AMR 2026-10-10 12:27:39 -07:00
22 changed files with 268 additions and 903 deletions

No files matched your search

+11
View File
@@ -794,6 +794,17 @@ system. See the corresponding
Defaults to None
#### [`PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR=<comma-separated-list>`](#PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR) {#PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR}
: A list of authentication method values which, if present in the `amr` claim of the ID token sent by an OpenID Connect provider, cause Paperless-ngx to skip its own two-factor authentication prompt for that login. This avoids users having to enter a second factor twice when the identity provider already performed multi-factor authentication. Logins with a username and password, or via a provider which does not send a matching `amr` value, still require the Paperless-ngx code.
Which values to trust depends entirely on your identity provider, e.g. Authelia sends `mfa` and Pocket ID sends `phr` for passkey logins (but `otp` for single-factor email codes, which should _not_ be trusted). Some providers, such as Keycloak, do not send an `amr` claim by default and need to be configured to do so. You can see what your provider sends from the Django admin as a superuser by checking `id_token` → `amr` in the "Extra data" of the social account.
!!! warning
Only list values that mean a second factor was actually used for the login.
Defaults to `mfa`, the standard value for multi-factor authentication ([RFC 8176](https://www.rfc-editor.org/rfc/rfc8176.html)).
#### [`PAPERLESS_SOCIAL_ACCOUNT_DEFAULT_GROUPS=<comma-separated-list>`](#PAPERLESS_SOCIAL_ACCOUNT_DEFAULT_GROUPS) {#PAPERLESS_SOCIAL_ACCOUNT_DEFAULT_GROUPS}
: A list of group names that users who signup via social accounts will be added to upon signup. Groups listed here must already exist.
+2
View File
@@ -473,6 +473,8 @@ Users can enable two-factor authentication (2FA) for their accounts from the 'My
Should a user lose access to their 2FA device and all recovery codes, a superuser can disable 2FA for the user from the 'Users & Groups' management screen.
If users log in via an OpenID Connect provider which already enforces multi-factor authentication, see [`PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR`](configuration.md#PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR) to skip the Paperless-ngx prompt for those logins.
## Workflows
!!! note
@@ -67,22 +67,18 @@ class Command(PaperlessCommand):
if options.get("recreate"):
wipe_index(settings.INDEX_DIR)
documents = (
Document.objects.filter(root_document__isnull=True)
.select_related(
"correspondent",
"document_type",
"storage_path",
"owner",
)
.prefetch_related(
"tags",
"notes__user",
"custom_fields__field",
"versions",
"barcodes",
"versions__barcodes",
)
documents = Document.objects.select_related(
"correspondent",
"document_type",
"storage_path",
"owner",
).prefetch_related(
"tags",
"notes__user",
"custom_fields__field",
"versions",
"barcodes",
"versions__barcodes",
)
total = documents.count()
rebuild_kwargs = {}
+38 -71
View File
@@ -6,9 +6,7 @@ from django.contrib.auth.models import Permission
from django.contrib.auth.models import User
from django.contrib.contenttypes.models import ContentType
from django.db.models import Case
from django.db.models import CharField
from django.db.models import Count
from django.db.models import F
from django.db.models import IntegerField
from django.db.models import Model
from django.db.models import Q
@@ -28,7 +26,6 @@ from rest_framework.permissions import BasePermission
from rest_framework.permissions import DjangoObjectPermissions
from documents.models import Document
from documents.versioning import get_root_document
class PaperlessObjectPermissions(DjangoObjectPermissions):
@@ -352,7 +349,6 @@ def permitted_object_ids(
perm: str,
*,
include_deleted: bool = False,
parent_field: str | None = None,
) -> QuerySet[int]:
"""
Generic version of ``permitted_document_ids`` for any model with an
@@ -361,24 +357,6 @@ def permitted_object_ids(
soft-delete pattern (currently only ``Document``); for every other model
it is accepted but has no effect, since those models have no soft-delete
concept.
``parent_field`` names a self-referencing foreign key whose target
authorizes the row (``Document.root_document``). A row with a parent is
visible exactly when its parent is, judged by the parent's owner and
grants, so the row's own owner and grants are ignored.
Guardian stores ``object_pk`` as a string, so the row key is cast to a
string and tested against the user's and groups' grants with a single
uncorrelated ``IN``. Postgres and SQLite build that set once. MariaDB
evaluates it as an index probe per row, which is cheap because the
lookups use guardian's unique indexes. Casting every ``object_pk`` to an
integer instead cannot use an index, and MariaDB cannot materialize it
inside the owner ``OR``, so it re-scans the user's grants for every row.
A correlated ``EXISTS`` per grant fixes MariaDB too, but Postgres and
SQLite re-run it for every row and end up slower than the original. The
user's groups are matched with an ``IN`` subquery rather than a join
through the membership table, which SQLite plans badly once the grant
tables grow.
"""
has_soft_delete = hasattr(model, "global_objects")
manager = (
@@ -386,21 +364,8 @@ def permitted_object_ids(
)
base_qs = manager.all().only("id", "owner")
owner_field, key_field = "owner", "pk"
if parent_field is not None:
owner_field, key_field = "authorizing_owner", "authorizing_id"
base_qs = base_qs.annotate(
authorizing_id=Coalesce(f"{parent_field}_id", "id"),
authorizing_owner=Case(
When(**{f"{parent_field}_id__isnull": True}, then=F("owner_id")),
default=F(f"{parent_field}__owner_id"),
output_field=IntegerField(),
),
)
unowned = Q(**{f"{owner_field}__isnull": True})
if user is None or not getattr(user, "is_authenticated", False):
return base_qs.filter(unowned).values_list("id", flat=True)
return base_qs.filter(owner__isnull=True).values_list("id", flat=True)
# Deactivated users get nothing, deactivated superusers included, so this
# has to come before the superuser shortcut. guardian's
@@ -424,26 +389,21 @@ def permitted_object_ids(
"permission__content_type": content_type,
}
# Both grant sets are compared to the row key as strings, exactly as
# guardian stores them, and are uncorrelated, so each engine can build the
# set once instead of probing per row.
user_keys = UserObjectPermission.objects.filter(
user=user,
**perm_filter,
).values_list("object_pk", flat=True)
group_keys = GroupObjectPermission.objects.filter(
group_id__in=user.groups.values("id"),
**perm_filter,
).values_list("object_pk", flat=True)
permitted_keys = user_keys.union(group_keys, all=True)
return (
base_qs.annotate(permitted_key=Cast(key_field, CharField(max_length=64)))
.filter(
Q(**{owner_field: user.pk}) | unowned | Q(permitted_key__in=permitted_keys),
)
.values_list("id", flat=True)
user_perm_ids = (
UserObjectPermission.objects.filter(user=user, **perm_filter)
.annotate(object_pk_int=Cast("object_pk", IntegerField()))
.values_list("object_pk_int", flat=True)
)
group_perm_ids = (
GroupObjectPermission.objects.filter(group__user=user, **perm_filter)
.annotate(object_pk_int=Cast("object_pk", IntegerField()))
.values_list("object_pk_int", flat=True)
)
permitted_ids = user_perm_ids.union(group_perm_ids)
return base_qs.filter(
Q(owner=user) | Q(owner__isnull=True) | Q(id__in=permitted_ids),
).values_list("id", flat=True)
ModelT = TypeVar("ModelT", bound=Model)
@@ -511,16 +471,30 @@ def permitted_document_ids(
``include_deleted=True`` for callers that need to check permission on
soft-deleted documents (e.g. trash restore). This intentionally avoids
``get_objects_for_user`` to keep the subquery small and index-friendly.
A version is authorized by its root document, so a version's own owner and
grants never matter.
"""
return permitted_object_ids(
user,
Document,
perm,
include_deleted=include_deleted,
parent_field="root_document",
return permitted_object_ids(user, Document, perm, include_deleted=include_deleted)
def documents_without_permitted_root(
documents: QuerySet[Document],
user: User | None,
*,
perm: str = "view_document",
include_deleted: bool = False,
) -> QuerySet[Document]:
"""
The documents the user lacks ``perm`` on. Versions are authorized by their
root document, so a version's own owner is ignored. A single query, without
loading the documents or joining the root.
"""
return documents.annotate(
root_id=Coalesce("root_document_id", "id"),
).exclude(
root_id__in=permitted_document_ids(
user,
perm=perm,
include_deleted=include_deleted,
),
)
@@ -680,14 +654,7 @@ def has_perms_owner_aware(user, perms, obj):
single-object check still has many production callers. Several callers
remain across ``documents/``, ``paperless_mail/``, and ``paperless_ai/``
-- grep for this function name before removing it.
A document version is authorized by its root document, like in
``permitted_document_ids``, so a version's own owner and grants never
matter. Fetch the root with ``select_related("root_document__owner")`` to
avoid extra queries.
"""
if isinstance(obj, Document):
obj = get_root_document(obj)
checker = ObjectPermissionChecker(user)
return obj.owner is None or obj.owner == user or checker.has_perm(perms, obj)
+3 -15
View File
@@ -284,14 +284,9 @@ class WriteBatch:
and adding the new version. This ensures stale document data (e.g., after
permission changes) doesn't persist in the index.
Only root documents are indexed, with their effective content, so a
version is indexed as its root document.
Args:
document: Django Document instance to index
"""
if document.root_document_id is not None:
document = document.root_document
self.remove(document.pk)
doc = self._backend._build_tantivy_doc(document)
self._writer.add_document(doc)
@@ -316,22 +311,20 @@ class WriteBatch:
An id with no matching document (e.g. deleted between the caller
collecting ids and the batch running) is silently skipped, matching
``add_or_update()``'s existing single-document deferred-task behavior
rather than erroring or leaving a stale index entry. The id of a
version stands for its root document.
rather than erroring or leaving a stale index entry.
Args:
ids: Primary keys of Document instances to index
"""
from documents.models import Document
from documents.versioning import annotate_effective_content
from documents.versioning import root_document_ids
ids = list(ids)
if not ids:
return
queryset = annotate_effective_content(
Document.objects.filter(pk__in=root_document_ids(ids))
Document.objects.filter(pk__in=ids)
.select_related("correspondent", "document_type", "storage_path", "owner")
.prefetch_related(
"tags",
@@ -1043,21 +1036,16 @@ class TantivyBackend:
excluded from results.
Args:
doc_id: Primary key of the reference document, or of one of its
versions
doc_id: Primary key of the reference document
user: User for permission filtering (None for no filtering)
limit: Maximum number of IDs to return (None = all matching docs)
Returns:
List of similar document IDs (excluding the original)
"""
from documents.versioning import root_document_ids
self._ensure_open()
searcher = self._index.searcher()
# Only root documents are indexed, so a version stands for its root
doc_id = next(iter(root_document_ids([doc_id])), doc_id)
id_query = tantivy.Query.term_query(self._schema, "id", doc_id)
results = searcher.search(id_query, limit=1)
+1 -3
View File
@@ -28,9 +28,7 @@ logger = logging.getLogger("paperless.search")
# columns dropped. tantivy compares schemas by ordered field list, so an
# index built by v1 rejects every write against the v2 schema.
# v3 - barcodes JSON field for stored barcode contents
# v4 - root documents only. Earlier indexes may hold document versions under
# their own id and metadata, so they are rebuilt without them.
SCHEMA_VERSION: Final[int] = 4
SCHEMA_VERSION: Final[int] = 3
# Present in the index directory from the moment a full rebuild starts until it
# finishes. If a rebuild is interrupted it is left behind, so the half-built
+2 -1
View File
@@ -90,6 +90,7 @@ from documents.templating.utils import convert_format_str_to_template_format
from documents.templating.workflows import validate_workflow_template
from documents.validators import uri_validator
from documents.validators import url_validator
from documents.versioning import get_root_document
from documents.versioning import has_prefetched_effective_content
from documents.versioning import sort_versions_newest_first
@@ -2894,7 +2895,7 @@ class ShareLinkSerializer(OwnedObjectSerializer):
and has_perms_owner_aware(
self.user,
"view_document",
document,
get_root_document(document),
)
):
return document
+8 -5
View File
@@ -490,20 +490,23 @@ def update_document_content_maybe_archive_file(
shutil.move(thumbnail, document.thumbnail_path)
document.refresh_from_db()
root_document = (
document.root_document if document.root_document_id else document
)
logger.info(
f"Updating index for document {document_id} ({document.archive_checksum})",
f"Updating index for document {root_document.pk} ({document.archive_checksum})",
)
from documents.search import get_backend
get_backend().add_or_update(document)
get_backend().add_or_update(root_document)
ai_config = AIConfig()
if ai_config.llm_index_enabled:
llm_index_add_or_update_document(document)
llm_index_add_or_update_document(root_document)
clear_document_caches(document.pk)
if document.root_document_id is not None:
clear_document_caches(document.root_document_id)
if root_document.pk != document.pk:
clear_document_caches(root_document.pk)
except Exception:
logger.exception(
-102
View File
@@ -293,80 +293,6 @@ class TestAddOrUpdateIds:
assert backend.search_ids("updated", user=None) == [doc.pk]
class TestVersionsAreIndexedAsTheirRoot:
"""Only root documents are indexed, with their effective content, so
every write path that is handed a version indexes its root instead."""
@staticmethod
def _root_with_version() -> tuple[Document, Document]:
root = DocumentFactory(title="Statement", content="stale text")
version = DocumentFactory(
title="Statement",
content="latest text",
root_document=root,
version_index=1,
)
return root, version
def test_add_or_update_indexes_the_root_of_a_version(
self,
backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- The version is passed to add_or_update
THEN:
- The root is indexed with the version's text, and the version is not
"""
root, version = self._root_with_version()
backend.add_or_update(version)
assert backend.search_ids("latest", user=None) == [root.pk]
assert backend.search_ids("stale", user=None) == []
def test_add_or_update_ids_indexes_each_root_once(
self,
backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- Both ids are passed to add_or_update_ids
THEN:
- The root is indexed once and the version is not indexed
"""
root, version = self._root_with_version()
with backend.batch_update() as batch:
batch.add_or_update_ids([version.pk, root.pk])
assert backend.search_ids("Statement", user=None) == [root.pk]
assert backend.search_ids("latest", user=None) == [root.pk]
def test_add_or_update_ids_resolves_a_lone_version(
self,
backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- Only the version's id is passed to add_or_update_ids
THEN:
- The root is indexed
"""
root, version = self._root_with_version()
with backend.batch_update() as batch:
batch.add_or_update_ids([version.pk])
assert backend.search_ids("latest", user=None) == [root.pk]
class TestSearch:
"""Test search query parsing and matching via search_ids."""
@@ -1136,34 +1062,6 @@ class TestMoreLikeThis:
assert 150 not in ids
assert 151 in ids
def test_more_like_this_ids_seeded_by_version_uses_root(
self,
backend: TantivyBackend,
) -> None:
"""A version is not indexed, so it must be looked up as its root."""
root = DocumentFactory.create(
title="Important document",
content="financial information report",
)
version = DocumentFactory.create(
title="Important document",
content="financial information report",
root_document=root,
version_index=1,
)
other = DocumentFactory.create(
title="Another document",
content="financial information report",
)
backend.add_or_update(root)
backend.add_or_update(other)
ids = backend.more_like_this_ids(doc_id=version.pk, user=None)
assert other.pk in ids
assert root.pk not in ids
assert version.pk not in ids
class TestSingleton:
"""Test get_backend() and reset_backend() singleton lifecycle."""
-53
View File
@@ -1196,59 +1196,6 @@ class TestDocumentSearchApi(DirectoriesMixin, APITestCase):
self.assertIn(d3.id, result_ids)
self.assertNotIn(d4.id, result_ids)
def test_search_more_like_version_uses_its_root(self) -> None:
"""
GIVEN:
- A document similar in content to a root document, and one that is not
- A version of the root document, which is never indexed
WHEN:
- API request for more like the version
THEN:
- The documents similar to the root are returned, not the version
"""
indexed = {}
for name, title, content, day in (
("root", "bank statement 1", "things i paid for in august", (2019, 3, 4)),
(
"similar",
"bank statement 3",
"things i paid for in september",
(2020, 7, 9),
),
(
"other",
"Quarterly Report",
"quarterly revenue profit margin",
(2021, 11, 30),
),
):
with time_machine.travel(
timezone.make_aware(datetime.datetime(*day)),
tick=False,
):
indexed[name] = DocumentFactory(
title=title,
content=content,
created=datetime.date(*day),
added=timezone.make_aware(datetime.datetime(*day)),
)
version = DocumentFactory(
root_document=indexed["root"],
version_index=1,
content="things i paid for in august",
)
backend = get_backend()
for document in indexed.values():
backend.add_or_update(document)
response = self.client.get(f"/api/documents/?more_like_id={version.id}")
self.assertEqual(response.status_code, status.HTTP_200_OK)
result_ids = [r["id"] for r in response.data["results"]]
self.assertIn(indexed["similar"].id, result_ids)
self.assertNotIn(indexed["other"].id, result_ids)
self.assertNotIn(version.id, result_ids)
def test_more_like_requires_id_of_existing_document(self) -> None:
"""
GIVEN:
-22
View File
@@ -22,7 +22,6 @@ from documents.models import Document
from documents.tasks import update_document_content_maybe_archive_file
from paperless_testing.assertions import FileSystemAssertsMixin
from paperless_testing.dirs import DirectoriesMixin
from paperless_testing.factories import DocumentFactory
sample_file: Path = Path(__file__).parent / "samples" / "simple.pdf"
@@ -117,27 +116,6 @@ class TestMakeIndex:
call_command("document_index", "reindex", skip_checks=True)
mock_get_backend.return_value.rebuild.assert_called_once()
def test_reindex_skips_versions(self, mocker: MockerFixture) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- The reindex command runs
THEN:
- Only the root document is handed to the rebuild, since a version
is indexed as its root
"""
root = DocumentFactory()
DocumentFactory(root_document=root, version_index=1)
mock_get_backend = mocker.patch(
"documents.management.commands.document_index.get_backend",
)
call_command("document_index", "reindex", skip_checks=True)
documents = mock_get_backend.return_value.rebuild.call_args.args[0]
assert list(documents.values_list("pk", flat=True)) == [root.pk]
def test_optimize(self) -> None:
"""Optimize command must execute without error (Tantivy handles optimization automatically)."""
call_command("document_index", "optimize", skip_checks=True)
@@ -18,7 +18,6 @@ from documents.models import Correspondent
from documents.models import DocumentType
from documents.models import StoragePath
from documents.models import Tag
from documents.permissions import has_perms_owner_aware
from documents.permissions import permitted_document_ids
from documents.permissions import permitted_object_ids
from documents.permissions import restrict_queryset_to_visible
@@ -33,8 +32,6 @@ from paperless_testing.permissions import grant_global
from paperless_testing.permissions import grant_object
if TYPE_CHECKING:
from django.contrib.auth.models import User
from paperless_testing.dirs import PaperlessDirs
@@ -181,382 +178,6 @@ class TestPermittedDocumentIdsIncludeDeleted:
)
@pytest.mark.django_db
class TestPermittedDocumentIdsVersions:
"""
A version is authorized by its root document: the version's own owner and
grants never matter.
"""
@pytest.mark.parametrize(
("root_owner", "version_owner", "expected_visible"),
[
pytest.param(
"other",
"nobody",
False,
id="unowned-version-of-private-root",
),
pytest.param("other", "user", False, id="own-version-of-private-root"),
pytest.param("user", "other", True, id="foreign-version-of-own-root"),
pytest.param("user", "nobody", True, id="unowned-version-of-own-root"),
pytest.param("nobody", "other", True, id="private-version-of-unowned-root"),
],
)
def test_version_follows_root_owner(
self,
root_owner: str,
version_owner: str,
*,
expected_visible: bool,
) -> None:
"""
GIVEN:
- A root document and a version with differing owners
WHEN:
- The permitted document ids are resolved for the user
THEN:
- The version is visible exactly when its root is
"""
user = UserFactory()
owners = {"user": user, "other": UserFactory(), "nobody": None}
root = DocumentFactory(owner=owners[root_owner])
version = DocumentFactory(root_document=root, owner=owners[version_owner])
visible = set(permitted_document_ids(user))
assert (version.pk in visible) is expected_visible
assert (root.pk in visible) is expected_visible
@staticmethod
def grantee(user: User, kind: str) -> User | Group:
"""The user itself, or a new group the user belongs to."""
if kind == "user":
return user
group = Group.objects.create(name="shared")
user.groups.add(group)
return group
@pytest.mark.parametrize(
"grantee_kind",
[pytest.param("user", id="user"), pytest.param("group", id="group")],
)
def test_grant_on_root_applies_to_version(self, grantee_kind: str) -> None:
"""
GIVEN:
- A private root document shared with a user or one of their groups
- A version of it owned by someone else
WHEN:
- The permitted document ids are resolved for the user
THEN:
- Both the root and the version are visible
- A user without the grant sees neither
"""
user = UserFactory()
stranger = UserFactory()
root = DocumentFactory(owner=UserFactory())
version = DocumentFactory(root_document=root, owner=UserFactory())
grant_object(self.grantee(user, grantee_kind), root, "view_document")
assert_visible_document_ids(
permitted_document_ids(user),
expected_visible=[root.pk, version.pk],
expected_hidden=[],
)
assert_visible_document_ids(
permitted_document_ids(stranger),
expected_visible=[],
expected_hidden=[root.pk, version.pk],
)
@pytest.mark.parametrize(
"grantee_kind",
[pytest.param("user", id="user"), pytest.param("group", id="group")],
)
def test_grant_on_version_is_ignored(self, grantee_kind: str) -> None:
"""
GIVEN:
- A private root document
- A version with an explicit grant for the user or one of their groups
WHEN:
- The permitted document ids are resolved for the user
THEN:
- Neither the root nor the version is visible
"""
user = UserFactory()
root = DocumentFactory(owner=UserFactory())
version = DocumentFactory(root_document=root, owner=UserFactory())
grant_object(self.grantee(user, grantee_kind), version, "view_document")
assert_visible_document_ids(
permitted_document_ids(user),
expected_visible=[],
expected_hidden=[root.pk, version.pk],
)
def test_grant_on_one_root_does_not_reach_another_roots_version(self) -> None:
"""
GIVEN:
- Two private roots, each with a version
- The user may view only the first root
WHEN:
- The permitted document ids are resolved for the user
THEN:
- Only the first root and its version are visible
"""
user = UserFactory()
first = DocumentFactory(owner=UserFactory())
first_version = DocumentFactory(root_document=first, owner=UserFactory())
second = DocumentFactory(owner=UserFactory())
second_version = DocumentFactory(root_document=second, owner=user)
grant_object(user, first, "view_document")
assert_visible_document_ids(
permitted_document_ids(user),
expected_visible=[first.pk, first_version.pk],
expected_hidden=[second.pk, second_version.pk],
)
def test_user_in_several_groups(self) -> None:
"""
GIVEN:
- A user in two groups
- Two private roots shared with one group each, and a third shared with nobody
- A version of each root
WHEN:
- The permitted document ids are resolved for the user
THEN:
- The two shared roots and their versions are visible
- The third root and its version are not
"""
user = UserFactory()
groups = [Group.objects.create(name=f"group{i}") for i in range(2)]
user.groups.add(*groups)
shared = [DocumentFactory(owner=UserFactory()) for _ in groups]
for root, group in zip(shared, groups, strict=True):
grant_object(group, root, "view_document")
unshared = DocumentFactory(owner=UserFactory())
shared_versions = [
DocumentFactory(root_document=root, owner=UserFactory()) for root in shared
]
unshared_version = DocumentFactory(root_document=unshared, owner=None)
assert_visible_document_ids(
permitted_document_ids(user),
expected_visible=[
*(root.pk for root in shared),
*(version.pk for version in shared_versions),
],
expected_hidden=[unshared.pk, unshared_version.pk],
)
def test_permission_is_resolved_through_the_root(self) -> None:
"""
GIVEN:
- A private root document where the user may view and change
WHEN:
- The permitted ids are resolved for view, change and delete
THEN:
- The version is visible for view and change only
"""
user = UserFactory()
root = DocumentFactory(owner=UserFactory())
version = DocumentFactory(root_document=root, owner=UserFactory())
grant_object(user, root, "view_document", "change_document")
assert version.pk in set(permitted_document_ids(user))
assert version.pk in set(permitted_document_ids(user, perm="change_document"))
assert version.pk in set(
permitted_document_ids(user, perm="documents.change_document"),
)
assert version.pk not in set(
permitted_document_ids(user, perm="delete_document"),
)
def test_anonymous_sees_versions_of_unowned_roots_only(self) -> None:
"""
GIVEN:
- A version owned by nobody under a private root
- A version owned by someone under an unowned root
WHEN:
- The permitted document ids are resolved for an anonymous user
THEN:
- Only the version of the unowned root is visible
"""
private_root = DocumentFactory(owner=UserFactory())
private_version = DocumentFactory(root_document=private_root, owner=None)
open_root = DocumentFactory(owner=None)
open_version = DocumentFactory(root_document=open_root, owner=UserFactory())
assert_visible_document_ids(
permitted_document_ids(AnonymousUser()),
expected_visible=[open_root.pk, open_version.pk],
expected_hidden=[private_root.pk, private_version.pk],
)
def test_deleted_versions_follow_their_deleted_root(self) -> None:
"""
GIVEN:
- A soft-deleted root document and its version, which deleting the
root soft-deletes too; the version is owned by someone else
WHEN:
- The permitted document ids are resolved with and without deleted
documents
THEN:
- Nothing is visible by default
- With deleted documents included, the version is visible to the
root's owner and not to the version's own owner
"""
owner = UserFactory()
version_owner = UserFactory()
root = DocumentFactory(owner=owner)
version = DocumentFactory(root_document=root, owner=version_owner)
root.delete()
assert not {root.pk, version.pk} & set(permitted_document_ids(owner))
assert_visible_document_ids(
permitted_document_ids(owner, include_deleted=True),
expected_visible=[root.pk, version.pk],
expected_hidden=[],
)
assert_visible_document_ids(
permitted_document_ids(version_owner, include_deleted=True),
expected_visible=[],
expected_hidden=[root.pk, version.pk],
)
@pytest.mark.parametrize(
"is_superuser",
[
pytest.param(False, id="regular-user"),
pytest.param(True, id="superuser"),
],
)
def test_inactive_user_sees_no_versions(self, *, is_superuser: bool) -> None:
"""
GIVEN:
- An inactive user, possibly a superuser, who owns a root and its version
WHEN:
- The permitted document ids are resolved for them
THEN:
- Nothing is visible
"""
user = UserFactory(is_active=False, is_superuser=is_superuser)
root = DocumentFactory(owner=user)
version = DocumentFactory(root_document=root, owner=user)
assert_visible_document_ids(
permitted_document_ids(user),
expected_visible=[],
expected_hidden=[root.pk, version.pk],
)
def test_superuser_sees_all_versions(self) -> None:
"""
GIVEN:
- A private root owned by someone else, with a version
WHEN:
- The permitted document ids are resolved for a superuser
THEN:
- Both the root and the version are visible
"""
superuser = UserFactory(superuser=True)
root = DocumentFactory(owner=UserFactory())
version = DocumentFactory(root_document=root, owner=UserFactory())
assert_visible_document_ids(
permitted_document_ids(superuser),
expected_visible=[root.pk, version.pk],
expected_hidden=[],
)
@pytest.mark.django_db
class TestHasPermsOwnerAwareVersions:
"""
The single-object check agrees with permitted_document_ids: a version is
authorized by its root document.
"""
@pytest.mark.parametrize(
("root_owner", "version_owner", "expected"),
[
pytest.param(
"other",
"nobody",
False,
id="unowned-version-of-private-root",
),
pytest.param("other", "user", False, id="own-version-of-private-root"),
pytest.param("user", "other", True, id="foreign-version-of-own-root"),
pytest.param("nobody", "other", True, id="private-version-of-unowned-root"),
],
)
def test_version_follows_root_owner(
self,
root_owner: str,
version_owner: str,
*,
expected: bool,
) -> None:
"""
GIVEN:
- A root document and a version with differing owners
WHEN:
- The single-object check runs for the version
THEN:
- The version is allowed exactly when its root is
"""
user = UserFactory()
owners = {"user": user, "other": UserFactory(), "nobody": None}
root = DocumentFactory(owner=owners[root_owner])
version = DocumentFactory(root_document=root, owner=owners[version_owner])
assert has_perms_owner_aware(user, "view_document", version) is expected
assert has_perms_owner_aware(user, "view_document", root) is expected
def test_grant_on_root_applies_and_grant_on_version_does_not(self) -> None:
"""
GIVEN:
- A private root with a version, and a second private root with a version
- The user may change only the first root, and was granted the second
root's version directly
WHEN:
- The single-object check runs for each version
THEN:
- Only the first root's version is allowed
"""
user = UserFactory()
shared_root = DocumentFactory(owner=UserFactory())
shared_version = DocumentFactory(root_document=shared_root, owner=UserFactory())
private_root = DocumentFactory(owner=UserFactory())
private_version = DocumentFactory(
root_document=private_root,
owner=UserFactory(),
)
grant_object(user, shared_root, "change_document")
grant_object(user, private_version, "change_document")
assert has_perms_owner_aware(user, "change_document", shared_version)
assert not has_perms_owner_aware(user, "change_document", private_version)
def test_other_models_use_their_own_owner(self) -> None:
"""
GIVEN:
- A tag owned by someone else, and one owned by the user
WHEN:
- The single-object check runs for each
THEN:
- Only the user's own tag is allowed without a grant
"""
user = UserFactory()
mine = TagFactory(owner=user)
theirs = TagFactory(owner=UserFactory())
assert has_perms_owner_aware(user, "view_tag", mine)
assert not has_perms_owner_aware(user, "view_tag", theirs)
@pytest.mark.django_db
class TestAiChatAllDocumentsPermissionBoundary:
"""
+9 -7
View File
@@ -310,7 +310,7 @@ class TestUpdateContent(DirectoriesMixin, TestCase):
@mock.patch("documents.tasks.clear_document_caches")
@mock.patch("documents.search.get_backend")
def test_update_content_version_clears_caches_for_root(
def test_update_content_version_indexes_root(
self,
mock_get_backend: mock.Mock,
mock_clear_caches: mock.Mock,
@@ -321,8 +321,8 @@ class TestUpdateContent(DirectoriesMixin, TestCase):
WHEN:
- Update content task is called for the version
THEN:
- The version's content is updated, not the root's
- The document is indexed
- The version's content is updated
- The root document is indexed rather than the version
- Caches are cleared for both
"""
root, version = self._create_root_with_version()
@@ -334,7 +334,8 @@ class TestUpdateContent(DirectoriesMixin, TestCase):
"my document",
)
self.assertEqual(Document.objects.get(pk=root.pk).content, "root content")
mock_get_backend.return_value.add_or_update.assert_called_once()
indexed = mock_get_backend.return_value.add_or_update.call_args.args[0]
self.assertEqual(indexed.pk, root.pk)
mock_clear_caches.assert_has_calls(
[mock.call(version.pk), mock.call(root.pk)],
)
@@ -342,7 +343,7 @@ class TestUpdateContent(DirectoriesMixin, TestCase):
@override_settings(AI_ENABLED=True, LLM_EMBEDDING_BACKEND="huggingface")
@mock.patch("documents.tasks.llm_index_add_or_update_document")
@mock.patch("documents.search.get_backend")
def test_update_content_version_updates_llm_index(
def test_update_content_version_updates_llm_index_for_root(
self,
mock_get_backend: mock.Mock,
mock_llm_index: mock.Mock,
@@ -354,13 +355,14 @@ class TestUpdateContent(DirectoriesMixin, TestCase):
WHEN:
- Update content task is called for the version
THEN:
- The LLM index is updated
- The LLM index is updated for the root document, not the version
"""
_, version = self._create_root_with_version()
root, version = self._create_root_with_version()
tasks.update_document_content_maybe_archive_file(version.pk)
mock_llm_index.assert_called_once()
self.assertEqual(mock_llm_index.call_args.args[0].pk, root.pk)
class TestUpdateContentRemoteOCR(DirectoriesMixin, TestCase):
-17
View File
@@ -17,26 +17,9 @@ from django.db.models.functions import RowNumber
from documents.models import Document
if TYPE_CHECKING:
from collections.abc import Iterable
from rest_framework.request import Request
def root_document_ids(ids: Iterable[int]) -> QuerySet[int]:
"""
The ids of the root documents of the given documents: a root stands for
itself and a version for its root. Only the indexes' bookkeeping needs
this, since they hold root documents only.
"""
return (
Document.objects.filter(pk__in=ids)
.annotate(root_id=Coalesce("root_document_id", "id"))
.order_by()
.values_list("root_id", flat=True)
.distinct()
)
def versions_newest_first(documents: QuerySet[Document]) -> QuerySet[Document]:
"""
Sorts versions so the newest one comes first using version_index and not on id,
+37 -30
View File
@@ -174,6 +174,7 @@ from documents.permissions import TrashPermissions
from documents.permissions import ViewDocumentsPermissions
from documents.permissions import annotate_document_count_by_ids
from documents.permissions import annotate_document_count_for_related_queryset
from documents.permissions import documents_without_permitted_root
from documents.permissions import get_document_count_filter_for_user
from documents.permissions import get_objects_for_user_owner_aware
from documents.permissions import has_global_statistics_permission
@@ -329,10 +330,9 @@ def _get_tantivy_query_and_mode(params):
def _get_more_like_id(query_params: dict[str, Any], user: User | None) -> int:
try:
more_like_doc_id = int(query_params["more_like_id"])
more_like_doc = Document.objects.select_related(
"owner",
"root_document__owner",
).get(pk=more_like_doc_id)
more_like_doc = Document.objects.select_related("owner").get(
pk=more_like_doc_id,
)
except (TypeError, ValueError, Document.DoesNotExist):
raise PermissionDenied(_("Invalid more_like_id"))
@@ -343,8 +343,7 @@ def _get_more_like_id(query_params: dict[str, Any], user: User | None) -> int:
):
raise PermissionDenied(_("Insufficient permissions."))
# Only root documents are indexed, a version stands for its root
return more_like_doc.root_document_id or more_like_doc.pk
return more_like_doc_id
class SearchParams(NamedTuple):
@@ -1561,7 +1560,7 @@ class DocumentViewSet(
if request.user is not None and not has_perms_owner_aware(
request.user,
"change_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden("Insufficient permissions")
@@ -1624,7 +1623,7 @@ class DocumentViewSet(
if request.user is not None and not has_perms_owner_aware(
request.user,
"change_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden("Insufficient permissions")
@@ -1877,7 +1876,7 @@ class DocumentViewSet(
if currentUser is not None and not has_perms_owner_aware(
currentUser,
"view_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden("Insufficient permissions to view notes")
except Document.DoesNotExist:
@@ -1899,7 +1898,7 @@ class DocumentViewSet(
if currentUser is not None and not has_perms_owner_aware(
currentUser,
"change_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden(
"Insufficient permissions to create notes",
@@ -1942,7 +1941,7 @@ class DocumentViewSet(
if currentUser is not None and not has_perms_owner_aware(
currentUser,
"change_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden("Insufficient permissions to delete notes")
@@ -1992,7 +1991,7 @@ class DocumentViewSet(
if currentUser is not None and not has_perms_owner_aware(
currentUser,
"change_document",
doc,
get_root_document(doc),
):
return HttpResponseForbidden(
"Insufficient permissions to add share link",
@@ -2118,7 +2117,7 @@ class DocumentViewSet(
documents = Document.objects.filter(pk__in=document_ids)
if (
request.user is not None
and documents.exclude(id__in=permitted_document_ids(request.user)).exists()
and documents_without_permitted_root(documents, request.user).exists()
):
return HttpResponseForbidden("Insufficient permissions")
@@ -2453,7 +2452,7 @@ class ChatStreamingView(GenericAPIView[Any]):
if not has_perms_owner_aware(
request.user,
"view_document",
document,
get_root_document(document),
):
return HttpResponseForbidden("Insufficient permissions")
@@ -3010,8 +3009,12 @@ class DocumentOperationPermissionMixin(PassUserMixin, DocumentSelectionMixin):
user.has_perm(
"documents.change_document",
)
and not Document.global_objects.filter(pk__in=documents)
.exclude(pk__in=permitted_document_ids(user, perm="change_document"))
and not Document.global_objects.filter(
pk__in=[doc.pk for doc in root_docs],
)
.exclude(
pk__in=permitted_document_ids(user, perm="change_document"),
)
.exists()
)
@@ -3624,7 +3627,7 @@ class SelectionDataView(DocumentSelectionMixin, GenericAPIView[Any]):
documents = Document.objects.filter(pk__in=ids)
if (
documents.count() != len(ids)
or documents.exclude(id__in=permitted_document_ids(request.user)).exists()
or documents_without_permitted_root(documents, request.user).exists()
):
return HttpResponseForbidden("Insufficient permissions")
@@ -4121,16 +4124,21 @@ class BulkDownloadView(DocumentSelectionMixin, GenericAPIView[Any]):
validated_data=serializer.validated_data,
)
documents = Document.objects.filter(pk__in=ids)
versioned_documents = []
compression = serializer.validated_data.get("compression")
content = serializer.validated_data.get("content")
follow_filename_format = serializer.validated_data.get("follow_formatting")
if documents.exclude(id__in=permitted_document_ids(request.user)).exists():
return HttpResponseForbidden("Insufficient permissions")
versioned_documents = [
get_latest_version_for_root(get_root_document(document))
for document in documents
]
permitted_ids = set(permitted_document_ids(request.user))
for document in documents:
root_doc = get_root_document(document)
if root_doc.pk not in permitted_ids:
return HttpResponseForbidden("Insufficient permissions")
versioned_documents.append(
get_latest_version_for_root(
root_doc,
),
)
if content == "both":
strategy_class = OriginalAndArchiveStrategy
@@ -4803,7 +4811,7 @@ class ShareLinkBundleViewSet(PassUserMixin, ModelViewSet[ShareLinkBundle]):
)
denied_id = (
documents_qs.exclude(id__in=permitted_document_ids(request.user))
documents_without_permitted_root(documents_qs, request.user)
.order_by("pk")
.values_list("pk", flat=True)
.first()
@@ -5649,12 +5657,11 @@ class TrashView(ListModelMixin, PassUserMixin):
if doc_ids is not None
else self.filter_queryset(self.get_queryset()).all()
)
if docs.exclude(
id__in=permitted_document_ids(
request.user,
perm="delete_document",
include_deleted=True,
),
if documents_without_permitted_root(
docs,
request.user,
perm="delete_document",
include_deleted=True,
).exists():
return HttpResponseForbidden("Insufficient permissions")
action = serializer.validated_data.get("action")
+28
View File
@@ -4,6 +4,7 @@ from urllib.parse import quote
from allauth.account.adapter import DefaultAccountAdapter
from allauth.core import context
from allauth.headless.tokens.strategies.sessions import SessionTokenStrategy
from allauth.mfa.stages import AuthenticateStage
from allauth.socialaccount.adapter import DefaultSocialAccountAdapter
from django.conf import settings
from django.contrib.auth.models import Group
@@ -19,7 +20,34 @@ from paperless.signals import handle_social_account_updated
logger = logging.getLogger("paperless.auth")
class SocialAccountAuthenticateStage(AuthenticateStage):
"""
Skips the 2FA prompt for a social login if the ID token's `amr` claim shows the
identity provider already used one of the SOCIAL_ACCOUNT_MFA_TRUSTED_AMR methods.
"""
def _should_handle(self, request: HttpRequest) -> bool:
sociallogin = (self.login.signal_kwargs or {}).get("sociallogin")
if sociallogin is not None and settings.SOCIAL_ACCOUNT_MFA_TRUSTED_AMR:
id_token = (sociallogin.account.extra_data or {}).get("id_token") or {}
amr = id_token.get("amr") or []
if set(amr) & set(settings.SOCIAL_ACCOUNT_MFA_TRUSTED_AMR):
logger.debug(
f"Skipping 2FA for `{self.login.user}`, provider `{sociallogin.account.provider}` reported amr {amr}",
)
return False
return super()._should_handle(request)
class CustomAccountAdapter(DefaultAccountAdapter):
def get_login_stages(self) -> list[str]:
return [
"paperless.adapter.SocialAccountAuthenticateStage"
if stage == "allauth.mfa.stages.AuthenticateStage"
else stage
for stage in super().get_login_stages()
]
def is_open_for_signup(self, request):
"""
Check whether the site is open for signups, which can be
+5
View File
@@ -348,6 +348,11 @@ SOCIAL_ACCOUNT_SYNC_SUPERUSER_GROUP: Final[str | None] = os.getenv(
SOCIAL_ACCOUNT_SYNC_STAFF_GROUP: Final[str | None] = os.getenv(
"PAPERLESS_SOCIAL_ACCOUNT_SYNC_STAFF_GROUP",
)
SOCIAL_ACCOUNT_MFA_TRUSTED_AMR = [
value.strip()
for value in os.getenv("PAPERLESS_SOCIAL_ACCOUNT_MFA_TRUSTED_AMR", "mfa").split(",")
if value.strip()
]
HEADLESS_TOKEN_STRATEGY = "paperless.adapter.DrfTokenStrategy"
+90
View File
@@ -2,19 +2,26 @@ import logging
import pytest
from allauth.account.adapter import get_adapter
from allauth.account.models import Login
from allauth.core import context
from allauth.mfa.totp.internal.auth import TOTP
from allauth.mfa.totp.internal.auth import generate_totp_secret
from allauth.socialaccount.adapter import get_adapter as get_social_adapter
from allauth.socialaccount.models import SocialAccount
from allauth.socialaccount.models import SocialLogin
from django.contrib.auth.models import AnonymousUser
from django.contrib.auth.models import Group
from django.contrib.auth.models import User
from django.forms import ValidationError
from django.http import HttpRequest
from django.test import RequestFactory
from django.urls import reverse
from pytest_django.fixtures import Settings
from pytest_mock import MockerFixture
from rest_framework.authtoken.models import Token
from paperless.adapter import DrfTokenStrategy
from paperless.adapter import SocialAccountAuthenticateStage
from paperless_testing.factories import UserFactory
@@ -129,6 +136,89 @@ class TestCustomAccountAdapter:
assert not user2.is_superuser
@pytest.mark.django_db
class TestSocialAccountAuthenticateStage:
@staticmethod
def _should_handle(user: User, extra_data: dict | None) -> bool:
signal_kwargs = None
if extra_data is not None:
account = SocialAccount(
provider="openid_connect",
uid="uid",
extra_data=extra_data,
)
signal_kwargs = {"sociallogin": SocialLogin(user=user, account=account)}
request = RequestFactory().get("/")
request.session = {}
stage = SocialAccountAuthenticateStage(
None,
request,
Login(user=user, signal_kwargs=signal_kwargs),
)
return stage._should_handle(request)
def test_stage_replaces_allauth_stage(self) -> None:
stages = get_adapter().get_login_stages()
assert "paperless.adapter.SocialAccountAuthenticateStage" in stages
assert "allauth.mfa.stages.AuthenticateStage" not in stages
@pytest.mark.parametrize(
("trusted_amr", "extra_data", "expected"),
[
pytest.param(
["mfa"],
{"id_token": {"amr": ["pwd", "mfa"]}},
False,
id="trusted-amr",
),
pytest.param(
["phr"],
{"id_token": {"amr": ["otp"]}},
True,
id="untrusted-amr",
),
pytest.param([], {"id_token": {"amr": ["mfa"]}}, True, id="setting-unset"),
pytest.param(["mfa"], {"id_token": {}}, True, id="no-amr-claim"),
pytest.param(
["mfa"],
{"userinfo": {"amr": ["mfa"]}},
True,
id="userinfo-only",
),
pytest.param(["mfa"], None, True, id="regular-login"),
],
)
def test_mfa_skipped_only_for_trusted_amr(
self,
settings: Settings,
trusted_amr: list[str],
extra_data: dict | None,
*,
expected: bool,
) -> None:
"""
GIVEN:
- A user with TOTP enabled
WHEN:
- The user logs in, via a social account or not
THEN:
- The 2FA prompt is only skipped if the ID token's amr claim contains
one of SOCIAL_ACCOUNT_MFA_TRUSTED_AMR
"""
settings.SOCIAL_ACCOUNT_MFA_TRUSTED_AMR = trusted_amr
user = UserFactory(username="testuser")
TOTP.activate(user, generate_totp_secret())
assert self._should_handle(user, extra_data) is expected
def test_no_mfa_for_user_without_authenticator(self, settings: Settings) -> None:
settings.SOCIAL_ACCOUNT_MFA_TRUSTED_AMR = ["mfa"]
user = UserFactory(username="testuser")
assert not self._should_handle(user, {"id_token": {"amr": ["pwd"]}})
class TestCustomSocialAccountAdapter:
@pytest.mark.django_db
def test_is_open_for_signup(self, settings: Settings) -> None:
+10 -8
View File
@@ -4,7 +4,8 @@ from django.conf import settings
from django.contrib.auth.models import User
from documents.models import Document
from documents.permissions import permitted_document_ids
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_ai.base_model import ClassificationSuggestions
@@ -53,8 +54,8 @@ def _fulltext_similar_documents(
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
permitted_document_ids() since Tantivy's indexed permission fields lag the
DB via async reindexing and judge a version by its own owner.
restrict_queryset_to_visible() since Tantivy's indexed permission fields
lag the DB via async reindexing.
"""
from documents.search import get_backend
@@ -68,9 +69,10 @@ def _fulltext_similar_documents(
)
if not unrestricted:
allowed_ids = set(
Document.objects.filter(
pk__in=similar_ids,
id__in=permitted_document_ids(user),
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]
@@ -198,13 +200,13 @@ def get_taxonomy_context(
# quadratic scan in the vector store at best, and past ~32,763
# documents a hard sqlite3.OperationalError (SQLite's
# bound-parameter limit) at worst.
# permitted_document_ids() has its own superuser shortcut that would
# 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 = (
None
if user_is_unrestricted(user)
else list(permitted_document_ids(user))
else list(permitted_object_ids(user, Document, "view_document"))
)
nodes = retrieve_similar_nodes(
document,
+7 -13
View File
@@ -17,7 +17,6 @@ from documents.models import PaperlessTask
from documents.utils import IterWrapper
from documents.utils import QuerySetStream
from documents.utils import identity
from documents.versioning import root_document_ids
from paperless.config import AIConfig
from paperless_ai.db import db_connection_released
from paperless_ai.embedding import build_llm_index_text
@@ -444,11 +443,11 @@ def update_llm_index(
"Skipping LLM index update: migration check deferred; "
"will retry next run."
)
documents = (
Document.objects.filter(root_document__isnull=True)
.select_related("correspondent", "document_type", "storage_path")
.prefetch_related("tags", "notes", "custom_fields__field")
)
documents = Document.objects.select_related(
"correspondent",
"document_type",
"storage_path",
).prefetch_related("tags", "notes", "custom_fields__field")
no_documents = not documents.exists()
# Fast exit before touching config: nothing to index and no existing index.
@@ -484,7 +483,7 @@ def update_llm_index(
msg = "LLM index rebuilt successfully."
else:
scoped_documents = (
documents.filter(id__in=root_document_ids(document_ids))
documents.filter(id__in=document_ids)
if document_ids is not None
else documents
)
@@ -511,12 +510,7 @@ def update_llm_index(
def llm_index_add_or_update_document(document: Document):
"""
Add or atomically replace a document's chunks in the index. Only root
documents are indexed, so a version is indexed as its root document.
"""
if document.root_document_id is not None:
document = document.root_document
"""Add or atomically replace a document's chunks in the index."""
config = AIConfig()
new_nodes = build_document_node(
document,
+5 -83
View File
@@ -552,7 +552,7 @@ class TestGetTaxonomyContextVisibility:
return_value=[],
)
mock_permitted = mocker.patch(
"paperless_ai.ai_classifier.permitted_document_ids",
"paperless_ai.ai_classifier.permitted_object_ids",
)
user = UserFactory.create(is_superuser=True)
@@ -582,7 +582,7 @@ class TestGetTaxonomyContextVisibility:
return_value=[],
)
mock_permitted = mocker.patch(
"paperless_ai.ai_classifier.permitted_document_ids",
"paperless_ai.ai_classifier.permitted_object_ids",
)
get_taxonomy_context(document, None)
@@ -611,55 +611,16 @@ class TestGetTaxonomyContextVisibility:
return_value=[],
)
mock_permitted = mocker.patch(
"paperless_ai.ai_classifier.permitted_document_ids",
"paperless_ai.ai_classifier.permitted_object_ids",
return_value=[1, 2, 3],
)
user = UserFactory.create(is_superuser=False)
get_taxonomy_context(document, user)
mock_permitted.assert_called_once_with(user)
mock_permitted.assert_called_once_with(user, Document, "view_document")
assert mock_retrieve.call_args.kwargs["document_ids"] == [1, 2, 3]
@pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
def test_version_of_private_root_is_not_visible(
self,
mocker: pytest_mock.MockerFixture,
) -> None:
"""
GIVEN:
- A private root document owned by someone else
- A version of it whose own owner is unset, as when the root
changed hands after the version was created
WHEN:
- get_taxonomy_context() is called for a non-superuser
THEN:
- Neither the root nor the version is in the visible ids passed to
retrieve_similar_nodes(), since a version follows its root
"""
owner = UserFactory.create()
viewer = UserFactory.create(is_superuser=False)
root = DocumentFactory.create(content="private", owner=owner)
version = DocumentFactory.create(
content="private",
owner=None,
root_document=root,
version_index=1,
)
source = DocumentFactory.create(content="Some content", owner=viewer)
mock_retrieve = mocker.patch(
"paperless_ai.ai_classifier.retrieve_similar_nodes",
return_value=[],
)
get_taxonomy_context(source, viewer)
visible = mock_retrieve.call_args.kwargs["document_ids"]
assert source.pk in visible
assert root.pk not in visible
assert version.pk not in visible
@pytest.mark.django_db
class TestFulltextSimilarDocuments:
@@ -842,7 +803,7 @@ class TestFulltextSimilarDocuments:
- _fulltext_similar_documents() is called with that user
THEN:
- Only the still-permitted document is returned - the DB
re-check via permitted_document_ids() must catch the
re-check via restrict_queryset_to_visible() must catch the
document Tantivy's stale index still thinks is visible
"""
owner = UserFactory.create()
@@ -873,45 +834,6 @@ class TestFulltextSimilarDocuments:
assert [s["document_id"] for s in result] == [permitted.pk]
def test_excludes_version_of_private_root_for_regular_user(
self,
fulltext_backend: TantivyBackend,
) -> None:
"""
GIVEN:
- A regular user and a private root owned by someone else
- A version of that root with no owner of its own, which the
Tantivy index therefore treats as visible to everyone
WHEN:
- _fulltext_similar_documents() is called with that user
THEN:
- The version is not returned, since the DB re-check judges it by
its root
"""
owner = UserFactory.create()
viewer = UserFactory.create(is_superuser=False)
source = DocumentFactory.create(
content="shared content phrase",
owner=viewer,
)
root = DocumentFactory.create(
content="shared content phrase",
owner=owner,
)
version = DocumentFactory.create(
content="shared content phrase",
owner=None,
root_document=root,
version_index=1,
)
fulltext_backend.add_or_update(source)
fulltext_backend.add_or_update(root)
fulltext_backend.add_or_update(version)
result = _fulltext_similar_documents(source, user=viewer, top_k=5)
assert version.pk not in [s["document_id"] for s in result]
@pytest.mark.django_db
@override_settings(LLM_EMBEDDING_BACKEND="huggingface")
@@ -388,83 +388,6 @@ def test_update_llm_index_partial_update(
assert after[str(doc2.pk)] == before[str(doc2.pk)]
@pytest.mark.django_db
class TestLlmIndexVersions:
"""The LLM index holds root documents only: a version is indexed as its root."""
def test_add_or_update_document_indexes_a_version_as_its_root(
self,
temp_llm_index_dir: Path,
mock_embed_model: FakeEmbedding,
) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- The version is passed to llm_index_add_or_update_document
THEN:
- Only the root document is in the index
"""
root = DocumentFactory(content="root content")
version = DocumentFactory(root_document=root, version_index=1)
indexing.llm_index_add_or_update_document(version)
with indexing.get_vector_store() as store:
indexed = store.get_modified_times()
assert set(indexed) == {str(root.pk)}
def test_rebuild_skips_versions(
self,
temp_llm_index_dir: Path,
mock_embed_model: FakeEmbedding,
) -> None:
"""
GIVEN:
- A root document with a version
WHEN:
- The LLM index is rebuilt
THEN:
- Only the root document is in the index
"""
root = DocumentFactory()
DocumentFactory(root_document=root, version_index=1)
indexing.update_llm_index(rebuild=True)
with indexing.get_vector_store() as store:
indexed = store.get_modified_times()
assert set(indexed) == {str(root.pk)}
def test_incremental_update_by_version_id_refreshes_the_root(
self,
temp_llm_index_dir: Path,
mock_embed_model: FakeEmbedding,
) -> None:
"""
GIVEN:
- An indexed root document with a version whose root was modified since
WHEN:
- An incremental update is scoped to the version's id
THEN:
- The root's entry is refreshed and no entry exists for the version
"""
root = DocumentFactory()
version = DocumentFactory(root_document=root, version_index=1)
indexing.update_llm_index(rebuild=True)
Document.objects.filter(pk=root.pk).update(modified=timezone.now())
root.refresh_from_db()
indexing.update_llm_index(document_ids=[version.pk])
with indexing.get_vector_store() as store:
indexed = store.get_modified_times()
assert indexed == {str(root.pk): root.modified.isoformat()}
@pytest.mark.django_db
def test_add_or_update_document_updates_existing_entry(
temp_llm_index_dir: Path,
@@ -714,7 +637,6 @@ class TestLlmIndexAddOrUpdateDocumentEmptyContent:
doc = MagicMock(spec=Document)
doc.id = 42
doc.root_document_id = None
# Must not raise
indexing.llm_index_add_or_update_document(doc)