diff --git a/src/documents/serialisers.py b/src/documents/serialisers.py index 02c3866c6..4bb1ff8c1 100644 --- a/src/documents/serialisers.py +++ b/src/documents/serialisers.py @@ -1079,6 +1079,16 @@ class DocumentSerializer( ) def get_page_count(self, obj) -> int | None: + # Like content versions get their own page count from the newest version, + # use the prefetched versions cache to avoid an extra query + prefetched_cache = getattr(obj, "_prefetched_objects_cache", None) + prefetched_versions = ( + prefetched_cache.get("versions") + if isinstance(prefetched_cache, dict) + else None + ) + if obj.root_document_id is None and prefetched_versions: + return sort_versions_newest_first(prefetched_versions)[0].page_count return obj.page_count @extend_schema_field(DuplicateDocumentSummarySerializer(many=True)) diff --git a/src/documents/tests/test_api_document_versions.py b/src/documents/tests/test_api_document_versions.py index 92c8ede99..6fce8fa49 100644 --- a/src/documents/tests/test_api_document_versions.py +++ b/src/documents/tests/test_api_document_versions.py @@ -19,6 +19,7 @@ from documents.models import Document from documents.versioning import annotate_effective_content from documents.views import DocumentSelectionMixin from paperless_testing.dirs import DirectoriesMixin +from paperless_testing.factories import DocumentFactory from paperless_testing.factories import UserFactory from paperless_testing.http import read_streaming_response from paperless_testing.permissions import grant_global @@ -821,6 +822,26 @@ class TestDocumentVersioningApi(DirectoriesMixin, APITestCase): self.assertEqual(resp.status_code, status.HTTP_200_OK) self.assertEqual(resp.data["content"], "v1-content") + def test_page_count_resolves_to_latest_version(self) -> None: + root = DocumentFactory(page_count=2) + DocumentFactory(root_document=root, version_index=1, page_count=1) + unversioned = DocumentFactory(page_count=5) + + resp = self.client.get("/api/documents/?fields=id,page_count") + self.assertEqual(resp.status_code, status.HTTP_200_OK) + self.assertEqual( + {doc["id"]: doc["page_count"] for doc in resp.data["results"]}, + {root.id: 1, unversioned.id: 5}, + ) + + resp = self.client.get(f"/api/documents/{root.id}/") + self.assertEqual(resp.status_code, status.HTTP_200_OK) + self.assertEqual(resp.data["page_count"], 1) + + resp = self.client.get(f"/api/documents/{root.id}/?version={root.id}") + self.assertEqual(resp.status_code, status.HTTP_200_OK) + self.assertEqual(resp.data["page_count"], 2) + def _make_root_with_out_of_order_versions(self) -> tuple[Document, ...]: """ A root whose newest version has a *lower* id than an older one, which is diff --git a/src/documents/views.py b/src/documents/views.py index 7a9b90406..3bbd1c0df 100644 --- a/src/documents/views.py +++ b/src/documents/views.py @@ -1192,6 +1192,7 @@ class DocumentViewSet( "version_label", "root_document_id", "version_index", + "page_count", ), ), "tags", @@ -1271,13 +1272,16 @@ class DocumentViewSet( if ( "version" not in request.query_params or not isinstance(response.data, dict) - or "content" not in response.data + or not ({"content", "page_count"} & response.data.keys()) ): return response root_doc = self.get_object() content_doc = self._resolve_file_doc(root_doc, request) - response.data["content"] = content_doc.content or "" + if "content" in response.data: + response.data["content"] = content_doc.content or "" + if "page_count" in response.data: + response.data["page_count"] = content_doc.page_count return response def update(self, request, *args, **kwargs):