Skip to content
Merged
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
10 changes: 7 additions & 3 deletions datamint/api/base_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -593,6 +593,7 @@ def _make_request_with_pagination(self,
endpoint: str,
return_field: str | None = None,
limit: int | None = None,
max_page_size: int | None = None,
**kwargs
) -> Generator[tuple[httpx.Response, list | dict | str], None, None]:
"""Make paginated HTTP requests, yielding each page of results.
Expand All @@ -602,13 +603,16 @@ def _make_request_with_pagination(self,
endpoint: API endpoint path
return_field: Optional field name to extract from each item in the response
limit: Optional maximum number of items to retrieve
max_page_size: Optional cap on the per-request page size, for endpoints
whose server route enforces a lower limit than the client default.
**kwargs: Additional arguments for the request (e.g., params, json)

Yields:
Tuples of (HTTP response, items from the current page `response.json()`, for convenience)
"""
offset = 0
total_fetched = 0
page_size = min(_PAGE_LIMIT, max_page_size) if max_page_size is not None else _PAGE_LIMIT

use_json_pagination = method.upper() == 'POST' and 'json' in kwargs and isinstance(kwargs['json'], dict)

Expand All @@ -621,10 +625,10 @@ def _make_request_with_pagination(self,
if limit is not None and total_fetched >= limit:
break

page_limit = _PAGE_LIMIT
page_limit = page_size
if limit is not None:
remaining = limit - total_fetched
page_limit = min(_PAGE_LIMIT, remaining)
page_limit = min(page_size, remaining)

if use_json_pagination:
kwargs['json']['offset'] = str(offset)
Expand All @@ -649,7 +653,7 @@ def _make_request_with_pagination(self,
yield response, items_to_yield
total_fetched += len(items_to_yield)

if len(items) < _PAGE_LIMIT:
if len(items) < page_size:
break

offset += len(items)
Expand Down
8 changes: 8 additions & 0 deletions datamint/api/endpoints/inference_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ class InferenceApi(EntityBaseApi[InferenceJob]):
(image, frame, slice, volume).
"""

_max_page_size = 500 # server enforces `le=500` on the list route's `limit` query param

def __init__(self,
config: ApiConfig,
client: httpx.Client | None = None,
Expand All @@ -54,6 +56,12 @@ def _parse_job_response(self, data: dict) -> InferenceJob:
data['id'] = data.pop('job_id')
return self._init_entity_obj(**data)

def _init_entity_obj(self, **kwargs) -> InferenceJob:
"""Rename the server's 'job_id' to 'id' before constructing the entity."""
if 'job_id' in kwargs and 'id' not in kwargs:
kwargs['id'] = kwargs.pop('job_id')
return super()._init_entity_obj(**kwargs)

def _build_common_payload(
self,
model_name: str,
Expand Down
5 changes: 5 additions & 0 deletions datamint/api/entity_base_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,10 @@ def _stream_entity_request(self,
raise ItemNotFoundError(self.endpoint_base, {'id': entity_id}) from e
raise

#: Optional cap on the per-request page size, for endpoints whose server
#: route enforces a lower limit than the client default.
_max_page_size: int | None = None

def get_list(self, limit: int | None = None,
**kwargs) -> Sequence[T]:
"""Get entities with optional filtering.
Expand All @@ -159,6 +163,7 @@ def get_list(self, limit: int | None = None,
items_gen = self._make_request_with_pagination('GET', f'/{self.endpoint_base}',
return_field=self.endpoint_base,
limit=limit,
max_page_size=self._max_page_size,
**new_kwargs)

all_items = []
Expand Down
Loading