From dfebab8fa3bbb4be9bcda9703797e762d106332b Mon Sep 17 00:00:00 2001 From: Arnav Bule <94260402+GODOSTROYER@users.noreply.github.com> Date: Tue, 6 Oct 2026 15:55:31 +0530 Subject: [PATCH] perf(api): avoid model hydration for pagination counts Keep supplied count querysets lazy and count projected pages without loading their models. Preserve raw querysets, custom lengths, grouping, and row locks, with unit and PostgreSQL/API regression coverage. --- .../test_work_item_pagination_evaluation.py | 229 ++++++++++++++ .../contract/test_paginator_evaluation.py | 284 ++++++++++++++++++ .../unit/utils/test_paginator_evaluation.py | 259 ++++++++++++++++ apps/api/plane/utils/paginator.py | 27 +- 4 files changed, 796 insertions(+), 3 deletions(-) create mode 100644 apps/api/plane/tests/contract/api/test_work_item_pagination_evaluation.py create mode 100644 apps/api/plane/tests/contract/test_paginator_evaluation.py create mode 100644 apps/api/plane/tests/unit/utils/test_paginator_evaluation.py diff --git a/apps/api/plane/tests/contract/api/test_work_item_pagination_evaluation.py b/apps/api/plane/tests/contract/api/test_work_item_pagination_evaluation.py new file mode 100644 index 000000000000..c431efb43fef --- /dev/null +++ b/apps/api/plane/tests/contract/api/test_work_item_pagination_evaluation.py @@ -0,0 +1,229 @@ +# Copyright (c) 2023-present Plane Software, Inc. and contributors +# SPDX-License-Identifier: AGPL-3.0-only +# See the LICENSE file for details. + +"""Real request-stack checks: counting must not hydrate the whole issue set. + +API-key requests exercise the authenticator. The existing session_client +fixture force-authenticates; it exercises routing/permissions, not login. +Only the asynchronous recent-visit dispatch is mocked, not the ORM. +""" + +from unittest import mock +from uuid import uuid4 + +import pytest +import requests +from django.utils import timezone + +from plane.db.models import Issue, Project, ProjectMember, State, User, Workspace, WorkspaceMember + +pytestmark = [pytest.mark.contract, pytest.mark.django_db(databases="__all__")] + + +@pytest.fixture +def pagination_project(workspace, create_user): + project = Project.objects.create( + workspace=workspace, name="Pagination", identifier="PG", network=0, guest_view_all_features=False + ) + ProjectMember.objects.create(project=project, member=create_user, role=20, is_active=True) + state = State.objects.create(workspace=workspace, project=project, name="Todo", group="backlog", default=True) + return project, state + + +@pytest.fixture +def pagination_issues(workspace, create_user, pagination_project): + project, state = pagination_project + rows = [] + for index in range(9): + issue = Issue.objects.create( + workspace=workspace, + project=project, + state=state, + name=f"Visible {index:02}", + priority="high" if index % 2 == 0 else "low", + description_html="
" + "large description " * 128 + "
", + ) + Issue.objects.filter(pk=issue.pk).update(created_by_id=create_user.pk) + rows.append(issue) + return rows + + +def list_url(surface, workspace, project): + prefix = "/api/v1" if surface == "public" else "/api" + return f"{prefix}/workspaces/{workspace.slug}/projects/{project.pk}/issues/" + + +def fetch_page(client, url, params): + # Dispatch remains outside the database work being measured. The full + # repository test stack still supplies its normal Redis/broker services. + with mock.patch("plane.app.views.issue.base.recent_visited_task.delay"): + response = client.get(url, params) + body = response.json() + assert response.status_code == 200, body + return body + + +@pytest.mark.parametrize("surface", ["public", "app"]) +@pytest.mark.parametrize("page_number", [0, 1, 2, 3]) +def test_http_page_hydration_is_bounded_by_returned_items( + surface, page_number, workspace, pagination_project, pagination_issues, api_key_client, session_client +): + project, _ = pagination_project + client = api_key_client if surface == "public" else session_client + params = {"per_page": 4, "cursor": f"4:{page_number}:0", "order_by": "sequence_id"} + if surface == "public": + params["fields"] = "id,name" + url = list_url(surface, workspace, project) + fetch_page(client, url, params) + with mock.patch.object(Issue, "from_db", wraps=Issue.from_db) as hydrate: + body = fetch_page(client, url, params) + expected = pagination_issues[page_number * 4 : page_number * 4 + 4] + assert [str(row["id"]) for row in body["results"]] == [str(issue.pk) for issue in expected] + assert body["total_count"] == body["total_results"] == 9 + assert body["count"] == len(expected) + assert body["next_page_results"] is (9 > page_number * 4 + 4) + assert body["prev_page_results"] is (page_number > 0) + # Public serializer consumes only the page; the app callback uses values(). + assert hydrate.call_count == (len(expected) if surface == "public" else 0) + + +@pytest.mark.parametrize("surface", ["public", "app"]) +def test_http_cursor_round_trip_preserves_order_and_count( + surface, workspace, pagination_project, pagination_issues, api_key_client, session_client +): + project, _ = pagination_project + client = api_key_client if surface == "public" else session_client + url = list_url(surface, workspace, project) + params = {"per_page": 4, "order_by": "sequence_id"} + first = fetch_page(client, url, params) + second = fetch_page(client, url, {**params, "cursor": first["next_cursor"]}) + previous = fetch_page(client, url, {**params, "cursor": second["prev_cursor"]}) + assert [row["id"] for row in previous["results"]] == [row["id"] for row in first["results"]] + assert {row["id"] for row in first["results"]}.isdisjoint(row["id"] for row in second["results"]) + assert first["count"] == second["count"] == previous["count"] == 4 + + +@pytest.mark.parametrize("surface", ["public", "app"]) +def test_hidden_model_states_and_other_workspace_do_not_enter_totals( + surface, workspace, pagination_project, pagination_issues, create_user, api_key_client, session_client +): + project, state = pagination_project + for attrs in ({"is_draft": True}, {"archived_at": timezone.now().date()}, {"deleted_at": timezone.now()}): + issue = Issue.objects.create(workspace=workspace, project=project, state=state, name="Hidden") + Issue.objects.filter(pk=issue.pk).update(**attrs) + triage = State.objects.create(workspace=workspace, project=project, name="Triage", group="triage", is_triage=True) + Issue.objects.create(workspace=workspace, project=project, state=triage, name="Hidden triage") + sibling = Project.objects.create(workspace=workspace, name="Sibling private", identifier="SB", network=0) + Issue.objects.create(workspace=workspace, project=sibling, name="Sibling private issue") + other = Workspace.objects.create(name="Foreign", slug=f"foreign-{uuid4().hex}", owner=create_user) + foreign_project = Project.objects.create(workspace=other, name="Foreign", identifier="FG") + Issue.objects.create(workspace=other, project=foreign_project, name="Foreign private issue") + client = api_key_client if surface == "public" else session_client + body = fetch_page(client, list_url(surface, workspace, project), {"per_page": 4, "order_by": "sequence_id"}) + assert body["total_count"] == body["total_results"] == 9 + assert body["count"] == 4 + assert all(row["name"].startswith("Visible ") for row in body["results"]) + + +def test_app_legacy_filter_applies_to_rows_and_count(workspace, pagination_project, pagination_issues, session_client): + project, _ = pagination_project + body = fetch_page( + session_client, + list_url("app", workspace, project), + {"per_page": 4, "priority": "high", "order_by": "sequence_id"}, + ) + assert body["total_count"] == 5 + assert body["count"] == 4 + assert all(row["priority"] == "high" for row in body["results"]) + + +def test_app_restricted_guest_total_contains_only_their_work_items( + workspace, pagination_project, pagination_issues, session_client +): + project, state = pagination_project + token = uuid4().hex + guest = User.objects.create(email=f"guest-{token}@example.test", username=f"guest-{token}") + WorkspaceMember.objects.create(workspace=workspace, member=guest, role=5, is_active=True) + ProjectMember.objects.create(project=project, member=guest, role=5, is_active=True) + own = Issue.objects.create(workspace=workspace, project=project, state=state, name="Guest owned") + Issue.objects.filter(pk=own.pk).update(created_by_id=guest.pk) + session_client.force_authenticate(user=guest) + body = fetch_page(session_client, list_url("app", workspace, project), {"per_page": 4}) + assert body["total_count"] == body["count"] == 1 + assert [str(row["id"]) for row in body["results"]] == [str(own.pk)] + + +def test_app_nonmember_cannot_read_project_or_total(workspace, pagination_project, pagination_issues, session_client): + project, _ = pagination_project + token = uuid4().hex + outsider = User.objects.create(email=f"outsider-{token}@example.test", username=f"outsider-{token}") + session_client.force_authenticate(user=outsider) + with mock.patch("plane.app.views.issue.base.recent_visited_task.delay"): + response = session_client.get(list_url("app", workspace, project), {"per_page": 4}) + assert response.status_code in (403, 404) + assert "total_count" not in response.json() + + +@pytest.mark.parametrize("subgroup", [False, True]) +def test_app_grouped_metadata_counts_rows_not_group_containers( + subgroup, workspace, pagination_project, pagination_issues, session_client +): + project, _ = pagination_project + params = {"per_page": 2, "group_by": "priority", "order_by": "sequence_id"} + if subgroup: + params["sub_group_by"] = "state_id" + body = fetch_page(session_client, list_url("app", workspace, project), params) + rows = [] + for group in body["results"].values(): + if subgroup: + for nested in group["results"].values(): + rows.extend(nested["results"]) + else: + rows.extend(group["results"]) + assert body["total_count"] == 9 + assert body["count"] == len(rows) == 4 + assert len({row["id"] for row in rows}) == 4 + + +@pytest.mark.parametrize("surface", ["public", "app"]) +def test_empty_project_returns_zero_metadata(surface, workspace, pagination_project, api_key_client, session_client): + project, _ = pagination_project + client = api_key_client if surface == "public" else session_client + body = fetch_page(client, list_url(surface, workspace, project), {"per_page": 4}) + assert body["results"] == [] + assert body["count"] == body["total_count"] == body["total_pages"] == 0 + assert body["next_page_results"] is False + + +def test_public_nonmember_token_does_not_expose_project_total( + workspace, pagination_project, pagination_issues, api_client +): + from plane.db.models.api import APIToken + + project, _ = pagination_project + token = uuid4().hex + outsider = User.objects.create(email=f"token-outsider-{token}@example.test", username=f"out-{token}") + api_token = APIToken.objects.create(user=outsider, label="Pagination outsider", token=uuid4().hex) + api_client.credentials(HTTP_X_API_KEY=api_token.token) + response = api_client.get(list_url("public", workspace, project), {"per_page": 4}) + assert response.status_code in (403, 404) + assert "total_count" not in response.json() + + +@pytest.mark.django_db(transaction=True, databases="__all__") +def test_public_work_item_page_over_real_tcp(plane_server, api_token, workspace, pagination_project, pagination_issues): + """Socket smoke test; not a latency benchmark or a production deployment.""" + project, _ = pagination_project + response = requests.get( + plane_server.url + list_url("public", workspace, project), + headers={"X-API-Key": api_token.token}, + params={"per_page": 4, "fields": "id,name", "order_by": "sequence_id"}, + timeout=10, + ) + assert response.status_code == 200, response.text + body = response.json() + assert body["total_count"] == body["total_results"] == 9 + assert body["count"] == 4 + assert body["next_page_results"] is True + assert [row["name"] for row in body["results"]] == [f"Visible {index:02}" for index in range(4)] diff --git a/apps/api/plane/tests/contract/test_paginator_evaluation.py b/apps/api/plane/tests/contract/test_paginator_evaluation.py new file mode 100644 index 000000000000..8e843a6ea661 --- /dev/null +++ b/apps/api/plane/tests/contract/test_paginator_evaluation.py @@ -0,0 +1,284 @@ +# Copyright (c) 2023-present Plane Software, Inc. and contributors +# SPDX-License-Identifier: AGPL-3.0-only +# See the LICENSE file for details. + +"""Real-ORM regression coverage for count-only pagination. + +Run with the repository's PostgreSQL test settings. No ORM/authentication is +mocked; from_db is wrapped solely to observe real model construction. +""" + +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack, contextmanager +from types import SimpleNamespace +from unittest import mock +from uuid import uuid4 + +import pytest +from django.db import DatabaseError, connections, transaction +from django.db.models import Count, F, Window +from django.db.models.functions import RowNumber + +from plane.db.models import User +from plane.utils.paginator import BasePaginator, Cursor, CursorResult, OffsetPaginator + +pytestmark = [pytest.mark.contract, pytest.mark.django_db(databases="__all__")] + + +@contextmanager +def capture_statements(): + statements = [] + + def record(execute, sql, params, many, context): + statements.append({"alias": context["connection"].alias, "sql": sql}) + return execute(sql, params, many, context) + + with ExitStack() as stack: + for alias in connections: + stack.enter_context(connections[alias].execute_wrapper(record)) + yield statements + + +@pytest.fixture +def user_ids(): + token = uuid4().hex + users = User.objects.bulk_create( + [ + User( + email=f"paginator-{token}-{index}@example.test", + username=f"paginator-{token}-{index}", + first_name=f"User {index:02}", + is_active=True, + ) + for index in range(9) + ] + ) + return [user.pk for user in users] + + +def fresh_users(user_ids): + return User._base_manager.using("default").filter(pk__in=user_ids).order_by("id") + + +def request(page=0): + return SimpleNamespace(GET={"per_page": "4", "cursor": f"4:{page}:0"}) + + +class StaticPaginator: + """Only replaces page selection when testing a callback's count contract.""" + + def __init__(self, queryset): + self.queryset = queryset + + def get_result(self, limit, cursor): + return CursorResult( + self.queryset, + Cursor(limit, 1, False, False), + Cursor(limit, -1, True, False), + hits=9, + max_hits=1, + ) + + +def test_total_count_does_not_fetch_models_or_populate_queryset_cache(user_ids): + count_queryset = fresh_users(user_ids) + page_queryset = fresh_users(user_ids) + with capture_statements() as statements, mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + result = OffsetPaginator(page_queryset, total_count_queryset=count_queryset).get_result(limit=4) + assert result.hits == 9 + assert count_queryset._result_cache is None + assert result.results._result_cache is None + hydrate.assert_not_called() + assert statements + assert all("COUNT(" in statement["sql"].upper() for statement in statements) + assert {statement["alias"] for statement in statements} == {"default"} + + +@pytest.mark.parametrize("empty_kind", ["none", "filtered"]) +def test_explicit_empty_count_queryset_is_authoritative(user_ids, empty_kind): + count_queryset = fresh_users(user_ids) + count_queryset = count_queryset.none() if empty_kind == "none" else count_queryset.filter(pk=uuid4()) + with mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + result = OffsetPaginator(fresh_users(user_ids), total_count_queryset=count_queryset).get_result(limit=4) + assert result.hits == 0 + assert result.max_hits == 0 + assert count_queryset._result_cache is None + hydrate.assert_not_called() + # Deliberately inconsistent inputs test only the optional-count-source + # contract. Normal callers must supply equivalently scoped querysets. + + +def test_omitted_count_queryset_still_counts_the_original_rows(user_ids): + with mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + result = OffsetPaginator(fresh_users(user_ids)).get_result(limit=4) + assert result.hits == 9 + hydrate.assert_not_called() + + +def test_existing_count_queryset_cache_is_reused(user_ids): + count_queryset = fresh_users(user_ids) + loaded = list(count_queryset) + with capture_statements() as statements, mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + result = OffsetPaginator(fresh_users(user_ids), total_count_queryset=count_queryset).get_result(limit=4) + assert result.hits == len(loaded) + # Any remaining probe is count-only. Do not depend on #9948's separate + # proposal to remove the existing next-page count query. + assert all("COUNT(" in statement["sql"].upper() for statement in statements) + hydrate.assert_not_called() + + +@pytest.mark.parametrize("page", [0, 1, 2, 3, 20]) +def test_projected_page_count_does_not_rehydrate_original_page(user_ids, page): + expected = list(fresh_users(user_ids).values("id", "first_name"))[page * 4 : page * 4 + 4] + original_pages = [] + + def project(rows): + original_pages.append(rows) + return list(rows.values("id", "first_name")) + + count_queryset = fresh_users(user_ids) + with capture_statements() as statements, mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + response = BasePaginator().paginate( + request(page), queryset=fresh_users(user_ids), total_count_queryset=count_queryset, on_results=project + ) + assert response.data["results"] == expected + assert response.data["count"] == len(expected) + assert response.data["total_count"] == 9 + assert response.data["total_results"] == 9 + assert response.data["next_page_results"] is (9 > page * 4 + 4) + assert count_queryset._result_cache is None + assert original_pages[0]._result_cache is None + hydrate.assert_not_called() + row_queries = [statement for statement in statements if "COUNT(" not in statement["sql"].upper()] + assert len(row_queries) == 1, statements + assert '"first_name"' in row_queries[0]["sql"] + + +@pytest.mark.parametrize("kind", ["filter", "mapping", "deduplicated_projection"]) +def test_transformed_cardinality_does_not_replace_raw_page_count(user_ids, kind): + raw = fresh_users(user_ids)[:9] + + def transform(rows): + if kind == "deduplicated_projection": + return [{"is_active": value} for value in {row["is_active"] for row in rows.values("is_active")}] + projected = list(rows.values("id")) + return projected[:1] if kind == "filter" else {"items": projected} + + with mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + response = BasePaginator().paginate(request(), paginator=StaticPaginator(raw), on_results=transform) + assert response.data["count"] == 9 + assert len(response.data["results"]) == 1 + assert raw._result_cache is None + hydrate.assert_not_called() + + +@pytest.mark.parametrize("callback_kind", ["none", "passthrough", "serialize"]) +def test_raw_or_serialized_page_is_fetched_only_once(user_ids, callback_kind): + page = fresh_users(user_ids)[:4] + + def passthrough(rows): + return rows + + def serialize(rows): + return [{"id": user.id} for user in rows] + + callback = {"none": None, "passthrough": passthrough, "serialize": serialize}[callback_kind] + with capture_statements() as statements, mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + response = BasePaginator().paginate(request(), paginator=StaticPaginator(page), on_results=callback) + list(response.data["results"]) + assert response.data["count"] == 4 + assert hydrate.call_count == 4 + assert len(statements) == 1, statements + assert "COUNT(" not in statements[0]["sql"].upper() + + +def test_grouped_annotation_count_keeps_raw_queryset_cardinality(user_ids): + raw = fresh_users(user_ids).order_by().values("is_active").annotate(user_count=Count("pk")) + with capture_statements() as statements, mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + response = BasePaginator().paginate( + request(), paginator=StaticPaginator(raw), on_results=lambda rows: list(rows.values("user_count")) + ) + assert response.data["count"] == 1 + assert response.data["results"] == [{"user_count": 9}] + # Unsliced custom paginator results deliberately retain legacy evaluation. + assert raw._result_cache is not None + hydrate.assert_not_called() + assert statements + + +def test_window_filtered_page_is_counted_without_hydration(user_ids): + raw = ( + fresh_users(user_ids).annotate(position=Window(RowNumber(), order_by=F("id").asc())).filter(position__lte=2)[:2] + ) + with mock.patch.object(User, "from_db", wraps=User.from_db) as hydrate: + response = BasePaginator().paginate( + request(), paginator=StaticPaginator(raw), on_results=lambda rows: list(rows.values("id", "position")) + ) + assert response.data["count"] == 2 + assert len(response.data["results"]) == 2 + assert raw._result_cache is None + hydrate.assert_not_called() + + +def test_values_queryset_input_remains_supported(user_ids): + raw = fresh_users(user_ids).values("id", "first_name")[:4] + with capture_statements() as statements: + response = BasePaginator().paginate(request(), paginator=StaticPaginator(raw), on_results=list) + assert response.data["count"] == 4 + assert len(response.data["results"]) == 4 + assert len(statements) == 1 + + +@pytest.mark.django_db(transaction=True, databases="__all__") +def test_locking_page_still_acquires_and_releases_row_locks(user_ids): + if connections["default"].vendor != "postgresql": + pytest.skip("The lock-conflict assertion requires PostgreSQL") + + def try_lock(): + try: + with transaction.atomic(using="default"): + User._base_manager.using("default").select_for_update(nowait=True).get(pk=user_ids[0]) + return "acquired" + except DatabaseError as exc: + assert getattr(exc.__cause__, "sqlstate", None) == "55P03", repr(exc) + return "locked" + finally: + connections["default"].close() + + # Ignoring the page in the callback makes the metadata evaluation the only + # lock acquisition. A blind replacement with .count() loses these locks. + with ThreadPoolExecutor(max_workers=1) as executor: + with transaction.atomic(using="default"): + raw = fresh_users(user_ids).select_for_update()[:9] + response = BasePaginator().paginate(request(), paginator=StaticPaginator(raw), on_results=lambda _: []) + assert response.data["count"] == 9 + assert executor.submit(try_lock).result(timeout=10) == "locked" + assert executor.submit(try_lock).result(timeout=10) == "acquired" + + +def test_unsliced_ordering_sensitive_distinct_preserves_original_length(user_ids): + # DISTINCT also selects the id needed by order_by(), so iteration yields + # nine rows even though all nine visible is_active values are identical. + # An unsliced COUNT can drop that ordering column and see only one row. + raw = fresh_users(user_ids).values("is_active").distinct() + expected = list(raw.all()) + assert len(expected) == 9 + response = BasePaginator().paginate( + request(), paginator=StaticPaginator(raw), on_results=lambda rows: list(rows.values("is_active")) + ) + assert response.data["results"] == expected + assert response.data["count"] == len(expected) + + +def test_sliced_ordering_sensitive_distinct_keeps_the_page_boundary(user_ids): + raw = fresh_users(user_ids).values("is_active").distinct()[:4] + expected = list(raw.all()) + assert len(expected) == 4 + with capture_statements() as statements: + response = BasePaginator().paginate( + request(), paginator=StaticPaginator(raw), on_results=lambda rows: list(rows.values("is_active")) + ) + assert response.data["results"] == expected + assert response.data["count"] == 4 + assert raw._result_cache is None + assert any("COUNT(" in item["sql"].upper() for item in statements) diff --git a/apps/api/plane/tests/unit/utils/test_paginator_evaluation.py b/apps/api/plane/tests/unit/utils/test_paginator_evaluation.py new file mode 100644 index 000000000000..662d9a8c125e --- /dev/null +++ b/apps/api/plane/tests/unit/utils/test_paginator_evaluation.py @@ -0,0 +1,259 @@ +# Copyright (c) 2023-present Plane Software, Inc. and contributors +# SPDX-License-Identifier: AGPL-3.0-only +# See the LICENSE file for details. + +"""Count-only pagination must not require model iteration. + +These tests cover the orchestration contract. The companion contract suite +executes actual Django querysets and checks SQL and model construction. +""" + +from types import SimpleNamespace +from unittest import mock + +import pytest +from django.db.models import QuerySet + +from plane.utils.paginator import ( + BadPaginationError, + BasePaginator, + Cursor, + CursorResult, + OffsetPaginator, +) + +pytestmark = pytest.mark.unit + + +class _CountSource: + """A supplied count source must never be tested for truthiness.""" + + def __init__(self, total): + self.count = mock.Mock(return_value=total) + + def __bool__(self): + raise AssertionError("Testing a count queryset for truthiness loads its rows") + + +class _PageQuery: + """Lazy slicing/count protocol; deliberately not a Django ORM emulator.""" + + def __init__(self, rows): + self.rows = rows + self.count = mock.Mock(return_value=len(rows)) + + def __getitem__(self, key): + return type(self)(self.rows[key]) + + def values(self, *fields): + return type(self)(self.rows) + + def __len__(self): + raise AssertionError("The paginator must leave page evaluation to its consumer") + + +class _StaticPaginator: + def __init__(self, result): + self.result = result + + def get_result(self, limit, cursor): + return self.result + + def process_results(self, results): + return {"group": {"results": results}} + + +def _request(**params): + return SimpleNamespace(GET={"per_page": "4", **params}) + + +def _result(rows, cls=CursorResult): + return cls( + results=rows, + next=Cursor(4, 1, False, False), + prev=Cursor(4, -1, True, False), + hits=4, + max_hits=1, + ) + + +def _queryset(total=4): + rows = mock.MagicMock(spec=QuerySet) + rows.query = SimpleNamespace(select_for_update=False, is_sliced=True) + rows.count.return_value = total + rows.__len__.return_value = total + return rows + + +@pytest.mark.parametrize("total", [0, 1, 4, 5, 100_000]) +def test_supplied_count_source_is_not_boolean_tested(total): + count_source = _CountSource(total) + page = _PageQuery(range(5)) + result = OffsetPaginator(page, total_count_queryset=count_source).get_result(limit=4) + assert result.hits == total + count_source.count.assert_called_once_with() + page.count.assert_not_called() + + +def test_missing_count_source_uses_the_page_queryset(): + page = _PageQuery(range(5)) + result = OffsetPaginator(page).get_result(limit=4) + assert result.hits == 5 + page.count.assert_called_once_with() + + +@pytest.mark.parametrize("total", [0, 1, 4, 5, 8, 9, 21]) +@pytest.mark.parametrize("page_number", [0, 1, 2, 7]) +def test_page_boundaries_and_lazy_result_are_preserved(total, page_number): + source = _PageQuery(range(total)) + result = OffsetPaginator(source).get_result(limit=4, cursor=Cursor(4, page_number, False)) + offset = page_number * 4 + assert list(result.results.rows) == list(range(total))[offset : offset + 4] + assert result.hits == total + assert result.max_hits == (total + 3) // 4 + assert result.next.has_results is (total > offset + 4) + assert result.prev.has_results is (page_number > 0) + assert str(result.next) == f"4:{page_number + 1}:0" + assert str(result.prev) == f"4:{page_number - 1}:1" + + +def test_maximum_page_size_is_preserved(): + result = OffsetPaginator(_PageQuery(range(10)), max_limit=3).get_result(limit=20) + assert list(result.results.rows) == [0, 1, 2] + assert result.next.value == 3 + + +@pytest.mark.parametrize("offset", [-1, 3]) +def test_offset_limits_still_raise(offset): + with pytest.raises(BadPaginationError): + OffsetPaginator(_PageQuery(range(20)), max_offset=12).get_result(limit=4, cursor=Cursor(4, offset)) + + +def test_previous_cursor_with_unchanged_page_size_is_preserved(): + result = OffsetPaginator(_PageQuery(range(12))).get_result(limit=4, cursor=Cursor(4, 1, True)) + assert list(result.results.rows) == [4, 5, 6, 7] + assert result.prev.has_results is True + + +@pytest.mark.parametrize("transformed", [[], [{"id": 1}], {"items": []}, (), None]) +def test_callback_output_cardinality_is_not_used_for_page_count(transformed): + rows = _queryset() + callback = mock.Mock(return_value=transformed) + response = BasePaginator().paginate(_request(), paginator=_StaticPaginator(_result(rows)), on_results=callback) + callback.assert_called_once_with(rows) + assert response.data["results"] is transformed + assert response.data["count"] == 4 + rows.count.assert_called_once_with() + rows.__len__.assert_not_called() + + +def test_controller_output_does_not_change_source_count(): + rows = _queryset() + response = BasePaginator().paginate( + _request(), + paginator=_StaticPaginator(_result(rows)), + on_results=lambda _: [{"id": 1}], + controller=lambda _: {"summary": "not a row list"}, + ) + assert response.data["count"] == 4 + assert response.data["results"] == {"summary": "not a row list"} + rows.count.assert_called_once_with() + rows.__len__.assert_not_called() + + +def test_group_container_count_is_not_confused_with_raw_row_count(): + rows = _queryset() + response = BasePaginator().paginate( + _request(), + paginator=_StaticPaginator(_result(rows)), + on_results=lambda _: [{"id": 1}], + group_by_field_name="priority", + ) + assert response.data["count"] == 4 + assert len(response.data["results"]) == 1 + rows.__len__.assert_not_called() + + +@pytest.mark.parametrize("callback", [None, lambda rows: rows]) +def test_raw_queryset_response_keeps_the_single_evaluation_path(callback): + rows = _queryset() + response = BasePaginator().paginate(_request(), paginator=_StaticPaginator(_result(rows)), on_results=callback) + assert response.data["results"] is rows + assert response.data["count"] == 4 + rows.count.assert_not_called() + rows.__len__.assert_called_once_with() + + +def test_controller_returning_original_queryset_preserves_cache_filling(): + rows = _queryset() + response = BasePaginator().paginate( + _request(), + paginator=_StaticPaginator(_result(rows)), + on_results=lambda _: [], + controller=lambda _: rows, + ) + assert response.data["results"] is rows + rows.count.assert_not_called() + rows.__len__.assert_called_once_with() + + +@pytest.mark.parametrize("rows", [[], [1, 2], (1, 2, 3), "abc"]) +def test_non_queryset_result_keeps_sequence_length(rows): + response = BasePaginator().paginate( + _request(), paginator=_StaticPaginator(_result(rows)), on_results=lambda _: {"filtered": []} + ) + assert response.data["count"] == len(rows) + + +def test_custom_cursor_result_length_contract_is_preserved(): + class CustomResult(CursorResult): + def __len__(self): + return 17 + + rows = _queryset() + response = BasePaginator().paginate( + _request(), paginator=_StaticPaginator(_result(rows, cls=CustomResult)), on_results=lambda _: [] + ) + assert response.data["count"] == 17 + rows.count.assert_not_called() + rows.__len__.assert_not_called() + + +def test_cursor_result_sequence_contract_itself_is_unchanged(): + result = _result([1, 2, 3]) + assert len(result) == 3 + assert list(result) == [1, 2, 3] + assert result[1:] == [2, 3] + assert repr(result) == "