From df3710cb2c549fbc8c5164bb83ee687d7588e96a Mon Sep 17 00:00:00 2001 From: pantera Date: Fri, 2 Oct 2026 23:15:04 -0700 Subject: [PATCH] fix: page backward in auto_paging_iter after a before cursor --- src/workos/_base_client.py | 56 ++++++++++-- src/workos/_pagination.py | 50 ++++++++--- tests/test_pagination.py | 169 ++++++++++++++++++++++++++++++++++++- 3 files changed, 256 insertions(+), 19 deletions(-) diff --git a/src/workos/_base_client.py b/src/workos/_base_client.py index 7c6f7e64..f79504bc 100644 --- a/src/workos/_base_client.py +++ b/src/workos/_base_client.py @@ -34,7 +34,7 @@ resolve_async_backend, resolve_sync_backend, ) -from ._pagination import AsyncPage, ListMetadata, SyncPage +from ._pagination import AsyncPage, ListMetadata, PaginationDirection, SyncPage from ._types import D, Deserializable, RequestOptions try: @@ -677,8 +677,24 @@ def request_page( cast(Dict[str, Any], data.get("list_metadata", {})) ) - def _fetch(*, after: Optional[str] = None) -> SyncPage[D]: - next_params = {**(params or {}), "after": after} + direction: PaginationDirection = ( + "backward" + if params and params.get("before") and not params.get("after") + else "forward" + ) + + def _fetch( + *, after: Optional[str] = None, before: Optional[str] = None + ) -> SyncPage[D]: + # Follow-up requests send only the cursor for the page's direction; + # the API rejects requests that carry both "after" and "before". + next_params = dict(params or {}) + if direction == "backward": + next_params.pop("after", None) + next_params["before"] = before + else: + next_params.pop("before", None) + next_params["after"] = after return self.request_page( method=method, path=path, @@ -688,7 +704,12 @@ def _fetch(*, after: Optional[str] = None) -> SyncPage[D]: request_options=request_options, ) - return SyncPage(data=items, list_metadata=list_metadata, _fetch_page=_fetch) + return SyncPage( + data=items, + list_metadata=list_metadata, + _fetch_page=_fetch, + _direction=direction, + ) class AsyncWorkOSClient(_BaseWorkOSClient): @@ -920,8 +941,24 @@ async def request_page( cast(Dict[str, Any], data.get("list_metadata", {})) ) - async def _fetch(*, after: Optional[str] = None) -> AsyncPage[D]: - next_params = {**(params or {}), "after": after} + direction: PaginationDirection = ( + "backward" + if params and params.get("before") and not params.get("after") + else "forward" + ) + + async def _fetch( + *, after: Optional[str] = None, before: Optional[str] = None + ) -> AsyncPage[D]: + # Follow-up requests send only the cursor for the page's direction; + # the API rejects requests that carry both "after" and "before". + next_params = dict(params or {}) + if direction == "backward": + next_params.pop("after", None) + next_params["before"] = before + else: + next_params.pop("before", None) + next_params["after"] = after return await self.request_page( method=method, path=path, @@ -931,4 +968,9 @@ async def _fetch(*, after: Optional[str] = None) -> AsyncPage[D]: request_options=request_options, ) - return AsyncPage(data=items, list_metadata=list_metadata, _fetch_page=_fetch) + return AsyncPage( + data=items, + list_metadata=list_metadata, + _fetch_page=_fetch, + _direction=direction, + ) diff --git a/src/workos/_pagination.py b/src/workos/_pagination.py index 71c15ac9..96f4194d 100644 --- a/src/workos/_pagination.py +++ b/src/workos/_pagination.py @@ -12,6 +12,7 @@ Generic, Iterator, List, + Literal, Optional, TypeVar, ) @@ -20,6 +21,9 @@ T = TypeVar("T", bound=Deserializable) +PaginationDirection = Literal["forward", "backward"] +"""Direction a page paginates in: ``forward`` follows ``after``, ``backward`` follows ``before``.""" + @dataclass(slots=True) class ListMetadata: @@ -42,6 +46,7 @@ class SyncPage(Generic[T]): _fetch_page: Optional[Callable[..., "SyncPage[T]"]] = field( default=None, repr=False ) + _direction: PaginationDirection = field(default="forward", repr=False) @property def before(self) -> Optional[str]: @@ -58,15 +63,27 @@ def has_more(self) -> bool: return self.after is not None def auto_paging_iter(self) -> Iterator[T]: - """Iterate through all items across all pages.""" + """Iterate through all items across all pages. + + Follows the page's direction: forward pages keep fetching with + ``after``; a page first requested with a ``before`` cursor keeps + fetching with ``before`` and yields each page's items reversed. + """ page = self + backward = page._direction == "backward" while True: - yield from page.data + items = reversed(page.data) if backward else page.data + yield from items if not page.data: break - if not page.has_more() or page._fetch_page is None: - break - page = page._fetch_page(after=page.after) + if backward: + if page.before is None or page._fetch_page is None: + break + page = page._fetch_page(before=page.before) + else: + if not page.has_more() or page._fetch_page is None: + break + page = page._fetch_page(after=page.after) def __iter__(self) -> Iterator[T]: """Iterate through all items across all pages.""" @@ -82,6 +99,7 @@ class AsyncPage(Generic[T]): _fetch_page: Optional[Callable[..., Awaitable["AsyncPage[T]"]]] = field( default=None, repr=False ) + _direction: PaginationDirection = field(default="forward", repr=False) @property def before(self) -> Optional[str]: @@ -98,16 +116,28 @@ def has_more(self) -> bool: return self.after is not None async def auto_paging_iter(self) -> AsyncIterator[T]: - """Iterate through all items across all pages.""" + """Iterate through all items across all pages. + + Follows the page's direction: forward pages keep fetching with + ``after``; a page first requested with a ``before`` cursor keeps + fetching with ``before`` and yields each page's items reversed. + """ page = self + backward = page._direction == "backward" while True: - for item in page.data: + items = reversed(page.data) if backward else page.data + for item in items: yield item if not page.data: break - if not page.has_more() or page._fetch_page is None: - break - page = await page._fetch_page(after=page.after) + if backward: + if page.before is None or page._fetch_page is None: + break + page = await page._fetch_page(before=page.before) + else: + if not page.has_more() or page._fetch_page is None: + break + page = await page._fetch_page(after=page.after) def __aiter__(self) -> AsyncIterator[T]: """Iterate through all items across all pages.""" diff --git a/tests/test_pagination.py b/tests/test_pagination.py index 6d7fcdb1..9a5efa84 100644 --- a/tests/test_pagination.py +++ b/tests/test_pagination.py @@ -2,11 +2,13 @@ """Pagination tests: auto_paging_iter, before cursor stripping, and HTTP integration.""" +from dataclasses import dataclass +from typing import Any, Dict, List, Optional +from urllib.parse import parse_qs, urlparse + import pytest from workos._pagination import SyncPage, AsyncPage, ListMetadata -from dataclasses import dataclass -from typing import Any, Dict @dataclass @@ -140,3 +142,166 @@ def test_auto_paging_iter_fetches_two_pages(self, workos, httpx_mock): requests = httpx_mock.get_requests() assert len(requests) == 2 assert "after=cursor_page2" in str(requests[1].url) + + +def _user_json(user_id: str) -> Dict[str, Any]: + return { + "object": "user", + "id": user_id, + "first_name": None, + "last_name": None, + "profile_picture_url": None, + "email": f"{user_id}@example.com", + "email_verified": True, + "external_id": None, + "last_sign_in_at": None, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z", + } + + +def _user_page_json( + ids: List[str], + *, + before: Optional[str] = None, + after: Optional[str] = None, +) -> Dict[str, Any]: + return { + "data": [_user_json(i) for i in ids], + "list_metadata": {"before": before, "after": after}, + } + + +def _query(request: Any) -> Dict[str, List[str]]: + return parse_qs(urlparse(str(request.url)).query) + + +class TestAutoPagingDirection: + """auto_paging_iter follows the direction of the initial request.""" + + def test_forward_full_traversal(self, workos, httpx_mock): + for i in range(10): + ids = [f"u{n}" for n in range(i * 10 + 1, i * 10 + 11)] + httpx_mock.add_response( + json=_user_page_json(ids, after=f"u{(i + 1) * 10}" if i < 9 else None) + ) + + page = workos.user_management.list_users(limit=10) + assert [u.id for u in page.auto_paging_iter()] == [ + f"u{n}" for n in range(1, 101) + ] + + requests = httpx_mock.get_requests() + assert len(requests) == 10 + assert _query(requests[0]) == {"limit": ["10"], "order": ["desc"]} + for k in range(1, 10): + assert _query(requests[k]) == { + "limit": ["10"], + "order": ["desc"], + "after": [f"u{k * 10}"], + } + + def test_backward_full_traversal(self, workos, httpx_mock): + httpx_mock.add_response( + json=_user_page_json( + [f"u{n}" for n in range(11, 21)], before="u11", after="u21" + ) + ) + httpx_mock.add_response(json=_user_page_json([f"u{n}" for n in range(1, 11)])) + + page = workos.user_management.list_users(before="u21", limit=10) + assert [u.id for u in page.auto_paging_iter()] == [ + f"u{n}" for n in range(20, 0, -1) + ] + + requests = httpx_mock.get_requests() + assert len(requests) == 2 + assert _query(requests[0]) == { + "limit": ["10"], + "order": ["desc"], + "before": ["u21"], + } + assert _query(requests[1]) == { + "limit": ["10"], + "order": ["desc"], + "before": ["u11"], + } + + def test_request_page_does_not_mutate_params(self, workos, httpx_mock): + httpx_mock.add_response(json=_user_page_json(["u2"], before="u1", after="u3")) + httpx_mock.add_response(json=_user_page_json(["u1"])) + params: Dict[str, Any] = {"before": "u3", "limit": 1} + + page = workos.request_page( + "get", ("user_management", "users"), model=FakeItem, params=params + ) + assert [i.id for i in page.auto_paging_iter()] == ["u2", "u1"] + + assert params == {"before": "u3", "limit": 1} + assert len(httpx_mock.get_requests()) == 2 + + +@pytest.mark.asyncio +class TestAsyncAutoPagingDirection: + """Async auto_paging_iter follows the direction of the initial request.""" + + async def test_forward_full_traversal(self, async_workos, httpx_mock): + for i in range(10): + ids = [f"u{n}" for n in range(i * 10 + 1, i * 10 + 11)] + httpx_mock.add_response( + json=_user_page_json(ids, after=f"u{(i + 1) * 10}" if i < 9 else None) + ) + + page = await async_workos.user_management.list_users(limit=10) + assert [u.id async for u in page.auto_paging_iter()] == [ + f"u{n}" for n in range(1, 101) + ] + + requests = httpx_mock.get_requests() + assert len(requests) == 10 + assert _query(requests[0]) == {"limit": ["10"], "order": ["desc"]} + for k in range(1, 10): + assert _query(requests[k]) == { + "limit": ["10"], + "order": ["desc"], + "after": [f"u{k * 10}"], + } + + async def test_backward_full_traversal(self, async_workos, httpx_mock): + httpx_mock.add_response( + json=_user_page_json( + [f"u{n}" for n in range(11, 21)], before="u11", after="u21" + ) + ) + httpx_mock.add_response(json=_user_page_json([f"u{n}" for n in range(1, 11)])) + + page = await async_workos.user_management.list_users(before="u21", limit=10) + assert [u.id async for u in page.auto_paging_iter()] == [ + f"u{n}" for n in range(20, 0, -1) + ] + + requests = httpx_mock.get_requests() + assert len(requests) == 2 + assert _query(requests[0]) == { + "limit": ["10"], + "order": ["desc"], + "before": ["u21"], + } + assert _query(requests[1]) == { + "limit": ["10"], + "order": ["desc"], + "before": ["u11"], + } + + async def test_request_page_does_not_mutate_params(self, async_workos, httpx_mock): + httpx_mock.add_response(json=_user_page_json(["u2"], before="u1", after="u3")) + httpx_mock.add_response(json=_user_page_json(["u1"])) + params: Dict[str, Any] = {"before": "u3", "limit": 1} + + page = await async_workos.request_page( + "get", ("user_management", "users"), model=FakeItem, params=params + ) + assert [i.id async for i in page.auto_paging_iter()] == ["u2", "u1"] + + assert params == {"before": "u3", "limit": 1} + assert len(httpx_mock.get_requests()) == 2