Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions stripe/_object_classes.py
Original file line number Diff line number Diff line change
Expand Up @@ -344,6 +344,10 @@
}

V2_OBJECT_CLASSES: Dict[str, Tuple[str, str]] = {
"v2.search_result": (
"stripe.v2._search_result_object",
"SearchResultObject",
),
# V2 Object classes: The beginning of the section generated from our OpenAPI spec
"v2.billing.meter_event": ("stripe.v2.billing._meter_event", "MeterEvent"),
"v2.billing.meter_event_adjustment": (
Expand Down
1 change: 1 addition & 0 deletions stripe/_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,7 @@ def _convert_to_stripe_object(
and (
(getattr(obj, "object") == "list")
or (getattr(obj, "object") == "search_result")
or (getattr(obj, "object") == "v2.search_result")
)
):
obj._retrieve_params = params
Expand Down
1 change: 1 addition & 0 deletions stripe/v2/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from typing_extensions import TYPE_CHECKING
from stripe.v2._list_object import ListObject as ListObject
from stripe.v2._search_result_object import SearchResultObject as SearchResultObject
from stripe.v2._amount import Amount as Amount, AmountParam as AmountParam


Expand Down
67 changes: 67 additions & 0 deletions stripe/v2/_search_result_object.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
from copy import deepcopy
from typing import (
AsyncIterator,
Generic,
Iterator,
List,
Mapping,
Any,
Optional,
TypeVar,
)

from stripe._any_iterator import AnyIterator
from stripe._stripe_object import StripeObject


T = TypeVar("T", bound=StripeObject)


class SearchResultObject(StripeObject, Generic[T]):
"""A page of API v2 search results with POST-based auto-pagination."""

OBJECT_NAME = "v2.search_result"
data: List[T]
next_page_url: Optional[str]
previous_page_url: Optional[str]
total_count: int

def __iter__(self) -> Iterator[T]:
return getattr(self, "data", []).__iter__()

def __len__(self) -> int:
return getattr(self, "data", []).__len__()

def auto_paging_iter(self) -> AnyIterator[T]:
return AnyIterator(self._auto_paging_iter(), self._auto_paging_iter_async())

def _original_params(self) -> Mapping[str, Any]:
return deepcopy(self._retrieve_params)

def _auto_paging_iter(self) -> Iterator[T]:
page: SearchResultObject[T] = self
params = self._original_params()
while True:
for item in page.data:
yield item
if page.next_page_url is None:
break
result = self._request(
"post", page.next_page_url, params=params, base_address="api"
)
assert isinstance(result, SearchResultObject)
page = result

async def _auto_paging_iter_async(self) -> AsyncIterator[T]:
page: SearchResultObject[T] = self
params = self._original_params()
while True:
for item in page.data:
yield item
if page.next_page_url is None:
break
result = await self._request_async(
"post", page.next_page_url, params=params, base_address="api"
)
assert isinstance(result, SearchResultObject)
page = result
71 changes: 71 additions & 0 deletions tests/test_v2_search_result_object.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
import pytest

from stripe.v2 import SearchResultObject


class TestSearchResult(SearchResultObject):
requests = []
pages = []

def _request(self, method, url, params=None, **kwargs):
self.requests.append((method, url, params, kwargs))
return self.pages.pop(0)

async def _request_async(self, method, url, params=None, **kwargs):
return self._request(method, url, params=params, **kwargs)


def make_result(data, next_page_url):
result = TestSearchResult._construct_from(
values={
"object": "v2.search_result",
"data": data,
"next_page_url": next_page_url,
"previous_page_url": None,
"total_count": 2,
},
last_response=None,
requestor=None,
api_mode="V2",
)
return result


def test_auto_paging_replays_original_post_body():
params = {
"query": 'status:"active"',
"sort": ["name", "-created"],
"limit": 2,
"future_field": {"enabled": True},
}
first = make_result(["one"], "/v2/widgets/search?page=2")
first._retrieve_params = params
TestSearchResult.requests = []
TestSearchResult.pages = [
make_result([], "/v2/widgets/search?page=3"),
make_result(["two"], None),
]

assert list(first.auto_paging_iter()) == ["one", "two"]
assert [
(method, url, body) for method, url, body, _ in TestSearchResult.requests
] == [
("post", "/v2/widgets/search?page=2", params),
("post", "/v2/widgets/search?page=3", params),
]


@pytest.mark.asyncio
async def test_async_auto_paging_replays_original_post_body():
first = make_result(["one"], "/v2/widgets/search?page=2")
first._retrieve_params = {"query": "widgets"}
TestSearchResult.requests = []
TestSearchResult.pages = [make_result(["two"], None)]

results = [item async for item in first.auto_paging_iter()]
assert results == ["one", "two"]
assert TestSearchResult.requests[0][:3] == (
"post",
"/v2/widgets/search?page=2",
{"query": "widgets"},
)
Loading