From d9afad66a7f5aae07d4126e4b464c4d908bb5990 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Wed, 24 Jun 2026 20:13:46 -0700 Subject: [PATCH 01/11] feat(client): Add pagination and exist_ok to nemoclient Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 270 +++++++++++++++-- .../nemo_platform_plugin/client/endpoint.py | 60 +++- .../src/nemo_platform_plugin/client/method.py | 14 +- .../nemo_platform_plugin/client/response.py | 138 ++++++++- .../src/nemo_platform_plugin/client/types.py | 145 +++++++++- .../tests/client/test_pagination.py | 272 ++++++++++++++++++ .../nemo_example_plugin/types/endpoints.py | 17 +- 7 files changed, 877 insertions(+), 39 deletions(-) create mode 100644 packages/nemo_platform_plugin/tests/client/test_pagination.py diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 8ea210eed0..89a62dcba7 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -15,18 +15,31 @@ from __future__ import annotations +import time from collections.abc import Mapping from typing import TypeVar, get_args, get_origin, overload import httpx from nemo_platform_plugin.client.response import ( AsyncNemoBinaryResponse, + AsyncNemoPaginatedResponse, AsyncNemoStreamResponse, NemoBinaryResponse, + NemoPaginatedResponse, NemoResponse, NemoStreamResponse, + _AsyncPageFetcher, + _SyncPageFetcher, +) +from nemo_platform_plugin.client.types import ( + BinaryContent, + Paginated, + PaginationStrategy, + PreparedRequest, + ResponseT, + RetryPolicy, + Stream, ) -from nemo_platform_plugin.client.types import BinaryContent, PreparedRequest, ResponseT, Stream from pydantic import BaseModel ModelT = TypeVar("ModelT", bound=BaseModel) @@ -42,6 +55,18 @@ def _get_stream_model_type(response_type: type) -> type[BaseModel]: return args[0] +def _get_paginated_types(response_type: type) -> tuple[type[BaseModel], type]: + """Extract (ModelT, StrategyT) from a Paginated[ModelT, StrategyT] generic alias.""" + from nemo_platform_plugin.client.types import OffsetPagination + + args = get_args(response_type) + if not args: + raise TypeError(f"Paginated response type must be parameterized, got {response_type}") + model_type = args[0] + strategy_type = args[1] if len(args) > 1 else OffsetPagination + return model_type, strategy_type + + class BaseNemoClient: """Shared logic for sync and async NeMo clients. @@ -49,9 +74,16 @@ class BaseNemoClient: Subclasses provide the actual HTTP transport (sync or async). """ - def __init__(self, *, base_url: str, workspace: str | None = None) -> None: + def __init__( + self, + *, + base_url: str, + workspace: str | None = None, + retry: RetryPolicy | None = None, + ) -> None: self._base_url = base_url.rstrip("/") self._workspace = workspace + self._retry = retry @property def base_url(self) -> str: @@ -61,6 +93,10 @@ def base_url(self) -> str: def workspace(self) -> str | None: return self._workspace + @property + def retry(self) -> RetryPolicy | None: + return self._retry + def _resolve_path(self, request: PreparedRequest) -> str: """Resolve path template with client defaults and explicit params. @@ -92,6 +128,9 @@ def _is_binary(self, request: PreparedRequest) -> bool: def _is_stream(self, request: PreparedRequest) -> bool: return get_origin(request.response_type) is Stream + def _is_paginated(self, request: PreparedRequest) -> bool: + return get_origin(request.response_type) is Paginated + def _resolve_query_params(self, request: PreparedRequest) -> dict[str, str | int | bool] | None: """Filter out None values from query params for httpx.""" if request.query_params is None: @@ -110,38 +149,77 @@ def __init__( workspace: str | None = None, default_headers: Mapping[str, str] | None = None, timeout: float = DEFAULT_TIMEOUT, + retry: RetryPolicy | None = None, http_client: httpx.Client | None = None, ) -> None: - super().__init__(base_url=base_url, workspace=workspace) + super().__init__(base_url=base_url, workspace=workspace, retry=retry) self._http = http_client or httpx.Client( headers=dict(default_headers) if default_headers else None, timeout=timeout, ) + def _resolve_retry(self, retry: RetryPolicy | None) -> RetryPolicy | None: + """Resolve retry policy: per-call override > client default.""" + if retry is not None: + return retry + return self._retry + @overload def send( - self, request: PreparedRequest[BinaryContent], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[BinaryContent], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> NemoBinaryResponse: ... @overload def send( - self, request: PreparedRequest[Stream[ModelT]], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[Stream[ModelT]], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> NemoStreamResponse[ModelT]: ... @overload - def send(self, request: PreparedRequest[None], *, headers: dict[str, str] | None = None) -> NemoResponse[None]: ... + def send( + self, + request: PreparedRequest[Paginated[ModelT]], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, + ) -> NemoPaginatedResponse[ModelT]: ... + @overload + def send( + self, + request: PreparedRequest[None], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, + ) -> NemoResponse[None]: ... @overload def send( - self, request: PreparedRequest[ResponseT], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[ResponseT], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> NemoResponse[ResponseT]: ... def send( - self, request: PreparedRequest, *, headers: dict[str, str] | None = None - ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse: + self, + request: PreparedRequest, + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, + ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: """Send a prepared request and return a typed response. Args: request: The prepared request to send. headers: Optional per-request headers merged on top of client defaults and content-type headers. + retry: Optional per-request retry policy override. Takes + precedence over endpoint-level and client-level defaults. For binary and streaming endpoints, the caller should use the response as a context manager to ensure the connection is closed:: @@ -152,6 +230,17 @@ def send( """ if headers: request = request.with_headers(headers) + + resolved_retry = self._resolve_retry(retry) + + if resolved_retry is not None: + return self._send_with_retry(request, resolved_retry) + return self._send_once(request) + + def _send_once( + self, request: PreparedRequest + ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: + """Execute a single HTTP request without retry.""" url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) @@ -170,12 +259,62 @@ def send( model_type = _get_stream_model_type(request.response_type) return NemoStreamResponse(stream_ctx, model_type, request) + if self._is_paginated(request): + assert request.response_type is not None + raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) + model_type, strategy = _get_paginated_types(request.response_type) + return NemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy), strategy) + raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) return NemoResponse(http_response=raw, body=body, request=request) + def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _SyncPageFetcher: + """Create a page-fetching callback bound to this client and strategy.""" + + def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: + url = self._resolve_path(request) + req_headers = self._request_headers(request) + existing_params = self._resolve_query_params(request) or {} + page_params = strategy.page_query_params(page) + params = {**existing_params, **page_params} + return self._http.request( + request.method, url, content=request.content, headers=req_headers, params=params + ) + + return fetch + + def _send_with_retry( + self, request: PreparedRequest, policy: RetryPolicy + ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: + """Execute a request with retry logic for transient failures.""" + last_response: NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse | None = None + last_exc: Exception | None = None + + for attempt in range(policy.max_retries + 1): + try: + response = self._send_once(request) + except httpx.TransportError as exc: + last_exc = exc + if attempt < policy.max_retries: + time.sleep(policy.backoff_base * (2**attempt)) + continue + raise + + if isinstance(response, NemoResponse) and response.http_response.status_code in policy.retryable_status_codes: + last_response = response + if attempt < policy.max_retries: + time.sleep(policy.backoff_base * (2**attempt)) + continue + + return response + + # All retries exhausted — return the last response we got + assert last_response is not None + return last_response + class AsyncNemoClient(BaseNemoClient): """Async HTTP client for NeMo Platform APIs. @@ -190,37 +329,83 @@ def __init__( workspace: str | None = None, default_headers: Mapping[str, str] | None = None, timeout: float = DEFAULT_TIMEOUT, + retry: RetryPolicy | None = None, http_client: httpx.AsyncClient | None = None, ) -> None: - super().__init__(base_url=base_url, workspace=workspace) + super().__init__(base_url=base_url, workspace=workspace, retry=retry) self._http = http_client or httpx.AsyncClient( headers=dict(default_headers) if default_headers else None, timeout=timeout, ) + def _resolve_retry(self, retry: RetryPolicy | None) -> RetryPolicy | None: + """Resolve retry policy: per-call override > client default.""" + if retry is not None: + return retry + return self._retry + @overload async def send( - self, request: PreparedRequest[BinaryContent], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[BinaryContent], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> AsyncNemoBinaryResponse: ... @overload async def send( - self, request: PreparedRequest[Stream[ModelT]], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[Stream[ModelT]], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> AsyncNemoStreamResponse[ModelT]: ... @overload async def send( - self, request: PreparedRequest[None], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[Paginated[ModelT]], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, + ) -> AsyncNemoPaginatedResponse[ModelT]: ... + @overload + async def send( + self, + request: PreparedRequest[None], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> NemoResponse[None]: ... @overload async def send( - self, request: PreparedRequest[ResponseT], *, headers: dict[str, str] | None = None + self, + request: PreparedRequest[ResponseT], + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, ) -> NemoResponse[ResponseT]: ... async def send( - self, request: PreparedRequest, *, headers: dict[str, str] | None = None - ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse: + self, + request: PreparedRequest, + *, + headers: dict[str, str] | None = None, + retry: RetryPolicy | None = None, + ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse: """Send a prepared request and return a typed response.""" if headers: request = request.with_headers(headers) + + resolved_retry = self._resolve_retry(retry) + + if resolved_retry is not None: + return await self._send_with_retry(request, resolved_retry) + return await self._send_once(request) + + async def _send_once( + self, request: PreparedRequest + ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse: + """Execute a single HTTP request without retry.""" url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) @@ -239,8 +424,61 @@ async def send( model_type = _get_stream_model_type(request.response_type) return AsyncNemoStreamResponse(stream_ctx, model_type, request) + if self._is_paginated(request): + assert request.response_type is not None + raw = await self._http.request( + request.method, url, content=request.content, headers=req_headers, params=params + ) + model_type, strategy = _get_paginated_types(request.response_type) + return AsyncNemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy), strategy) + raw = await self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) return NemoResponse(http_response=raw, body=body, request=request) + + def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _AsyncPageFetcher: + """Create an async page-fetching callback bound to this client and strategy.""" + + async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: + url = self._resolve_path(request) + req_headers = self._request_headers(request) + existing_params = self._resolve_query_params(request) or {} + page_params = strategy.page_query_params(page) + params = {**existing_params, **page_params} + return await self._http.request( + request.method, url, content=request.content, headers=req_headers, params=params + ) + + return fetch + + async def _send_with_retry( + self, request: PreparedRequest, policy: RetryPolicy + ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse: + """Execute a request with retry logic for transient failures.""" + import asyncio + + last_response: NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse | None = None + last_exc: Exception | None = None + + for attempt in range(policy.max_retries + 1): + try: + response = await self._send_once(request) + except httpx.TransportError as exc: + last_exc = exc + if attempt < policy.max_retries: + await asyncio.sleep(policy.backoff_base * (2**attempt)) + continue + raise + + if isinstance(response, NemoResponse) and response.http_response.status_code in policy.retryable_status_codes: + last_response = response + if attempt < policy.max_retries: + await asyncio.sleep(policy.backoff_base * (2**attempt)) + continue + + return response + + assert last_response is not None + return last_response diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py index 7999a2594c..3d0681b41e 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py @@ -24,6 +24,7 @@ def hello(self, *, name: str) -> HelloResponse: - ``body`` — JSON request body (Pydantic model, serialized automatically) - ``content`` — binary request body (raw bytes) - ``query_params`` — query parameters (dict or TypedDict) +- Blessed client options (e.g. ``exist_ok``) — client-side behavior (stripped before sending) - All other keyword parameters — path parameters (matched to ``{placeholders}`` in the path template) """ @@ -33,21 +34,60 @@ def hello(self, *, name: str) -> HelloResponse: import inspect import string from collections.abc import AsyncIterable, Callable, Iterable -from typing import get_type_hints +from typing import Any, get_type_hints from nemo_platform_plugin.client.types import ( + BLESSED_CLIENT_PARAMS, P, PreparedRequest, ResponseT, ) from pydantic import BaseModel +# Parameter names with special handling in _build_prepared_request. +_RESERVED_PARAM_NAMES = frozenset({"self", "body", "content", "query_params"}) + + +def _identify_client_option_params(fn: Callable) -> set[str]: + """Return parameter names that are blessed client-side options. + + A parameter is a client option if its name appears in + :data:`BLESSED_CLIENT_PARAMS`. + """ + sig = inspect.signature(fn) + return set(sig.parameters.keys()) & BLESSED_CLIENT_PARAMS.keys() + + +def _validate_params( + fn: Callable, path_param_names: set[str], client_option_names: set[str] +) -> None: + """Raise ``TypeError`` at decoration time if any parameter is unrecognised. + + Every parameter must be one of: + - ``self`` + - A path placeholder (``{name}`` in the URL template) + - ``body``, ``content``, or ``query_params`` + - A blessed client option (e.g. ``exist_ok``) + """ + sig = inspect.signature(fn) + known = _RESERVED_PARAM_NAMES | path_param_names | client_option_names + unknown = set(sig.parameters.keys()) - known + if unknown: + name = getattr(fn, "__qualname__", getattr(fn, "__name__", repr(fn))) + blessed = ", ".join(sorted(BLESSED_CLIENT_PARAMS.keys())) + raise TypeError( + f"Endpoint {name} has unrecognised parameters: {unknown}. " + f"Parameters must be path params {path_param_names}, " + f"'body', 'content', 'query_params', or a client option ({blessed})." + ) + def _build_prepared_request( method: str, path: str, sig: inspect.Signature, path_param_names: set[str], + client_option_names: set[str], response_type: type | None, args: tuple, kwargs: dict, @@ -56,6 +96,9 @@ def _build_prepared_request( Uses ``bind_partial`` so that path parameters with client-level defaults (e.g. ``workspace``) can be omitted by the caller. + + Client option parameters (blessed names like ``exist_ok``) are stripped + from the HTTP request and stashed in ``PreparedRequest.client_options``. """ bound = sig.bind_partial(*args, **kwargs) bound.apply_defaults() @@ -64,11 +107,16 @@ def _build_prepared_request( query_params: dict[str, str | int | bool | None] | None = None content: bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None content_type: str | None = None + client_options: dict[str, Any] | None = None for name, value in bound.arguments.items(): if name == "self": continue - if name in path_param_names: + if name in client_option_names: + if client_options is None: + client_options = {} + client_options[name] = value + elif name in path_param_names: if value is not None: path_params[name] = str(value) elif name == "body": @@ -91,6 +139,7 @@ def _build_prepared_request( content_type=content_type, response_type=response_type, query_params=query_params, + client_options=client_options, ) @@ -98,13 +147,18 @@ def _make_endpoint(http_method: str, path: str, fn: Callable[P, ResponseT]) -> C """Create a callable that builds PreparedRequests from the function's signature.""" sig = inspect.signature(fn) path_param_names = {field_name for _, field_name, _, _ in string.Formatter().parse(path) if field_name} + client_option_names = _identify_client_option_params(fn) + _validate_params(fn, path_param_names, client_option_names) + hints = get_type_hints(fn) ret = hints.get("return") response_type = ret if ret is not None and ret is not type(None) else None @functools.wraps(fn) def prepare(*args: P.args, **kwargs: P.kwargs) -> PreparedRequest[ResponseT]: - return _build_prepared_request(http_method, path, sig, path_param_names, response_type, args, kwargs) + return _build_prepared_request( + http_method, path, sig, path_param_names, client_option_names, response_type, args, kwargs + ) return prepare diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py index 3c53d074bb..7e60afe981 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py @@ -17,6 +17,8 @@ class AsyncExampleClient(_ExampleMethods, AsyncNemoClient): pass resp = client.hello(name="alice") # NemoResponse[HelloResponse] The descriptor dispatches sync vs async based on the client type. +Client-side options (e.g. ``exist_ok``) declared in the endpoint +signature are applied after the HTTP call, wrapping the response. Note: ``ty`` shows ``Unknown |`` on the method types due to unannotated class attributes (astral-sh/ty#3254). The types themselves are correct @@ -30,6 +32,7 @@ class attributes (astral-sh/ty#3254). The types themselves are correct from typing import Any, Coroutine, Generic, overload from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.directives import apply_client_options from nemo_platform_plugin.client.response import NemoResponse from nemo_platform_plugin.client.types import P, PreparedRequest, ResponseT @@ -40,6 +43,9 @@ class EndpointMethod(Generic[P, ResponseT]): When accessed on a :class:`NemoClient`, returns a sync callable. When accessed on an :class:`AsyncNemoClient`, returns an async callable. Both preserve the endpoint's full ``ParamSpec`` signature. + + Client-side options (blessed parameter names like ``exist_ok``) + are automatically applied after the HTTP response is received. """ def __init__(self, endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> None: @@ -58,13 +64,17 @@ def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None @functools.wraps(self._endpoint_fn) async def async_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: - return await obj.send(self._endpoint_fn(*args, **kwargs)) + request = self._endpoint_fn(*args, **kwargs) + response = await obj.send(request) + return apply_client_options(request, response) return async_bound @functools.wraps(self._endpoint_fn) def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: - return obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] + request = self._endpoint_fn(*args, **kwargs) + response = obj.send(request) # type: ignore[assignment] + return apply_client_options(request, response) return sync_bound diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py index 3decc8f389..c1d794275c 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py @@ -5,14 +5,14 @@ from __future__ import annotations -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Callable, Coroutine, Iterator from contextlib import AbstractAsyncContextManager, AbstractContextManager from dataclasses import dataclass from types import TracebackType -from typing import Generic, TypeVar +from typing import Any, Generic, TypeVar import httpx -from nemo_platform_plugin.client.types import PreparedRequest +from nemo_platform_plugin.client.types import PaginationStrategy, PreparedRequest from pydantic import BaseModel ResponseT = TypeVar("ResponseT") @@ -216,6 +216,138 @@ async def __aexit__( await self._stream_ctx.__aexit__(exc_type, exc_val, exc_tb) +# --------------------------------------------------------------------------- +# Paginated responses +# --------------------------------------------------------------------------- + + +# Type aliases for the page-fetching callbacks used by paginated responses. +# The page value is int for offset-based or str for cursor-based pagination. +_SyncPageFetcher = Callable[[PreparedRequest, Any], httpx.Response] +_AsyncPageFetcher = Callable[[PreparedRequest, Any], Coroutine[Any, Any, httpx.Response]] + + +class NemoPaginatedResponse(Generic[ModelT]): + """Sync iterable over all items across paginated API responses. + + Lazily fetches subsequent pages using the pagination strategy configured + on the endpoint's ``Paginated[T, Strategy]`` return type. Iterating + yields individual ``ModelT`` items, not page envelopes:: + + resp = client.send(list_items()) + for item in resp: + print(item.name) + + Also supports fetching a single page:: + + resp = client.send(list_items()) + page = resp.first_page() # list[ModelT] from the first page + """ + + def __init__( + self, + first_http_response: httpx.Response, + model_type: type[ModelT], + request: PreparedRequest, + fetch_page: _SyncPageFetcher, + strategy: type[PaginationStrategy] | None = None, + ) -> None: + from nemo_platform_plugin.client.types import OffsetPagination + + self._first_response = first_http_response + self._model_type = model_type + self.request = request + self._fetch_page = fetch_page + self._strategy: type[PaginationStrategy] = strategy or OffsetPagination + + @property + def http_response(self) -> httpx.Response: + return self._first_response + + def _parse_items(self, raw: httpx.Response) -> list[ModelT]: + raw.raise_for_status() + body = raw.json() + raw_items = self._strategy.extract_items(body) + return [self._model_type.model_validate(item) for item in raw_items] + + def first_page(self) -> list[ModelT]: + """Return items from the first page (already fetched).""" + return self._parse_items(self._first_response) + + def __iter__(self) -> Iterator[ModelT]: + self._first_response.raise_for_status() + body = self._first_response.json() + + raw_items = self._strategy.extract_items(body) + yield from (self._model_type.model_validate(item) for item in raw_items) + + next_page = self._strategy.next_page(body, 1) + while next_page is not None: + raw = self._fetch_page(self.request, next_page) + raw.raise_for_status() + body = raw.json() + raw_items = self._strategy.extract_items(body) + yield from (self._model_type.model_validate(item) for item in raw_items) + current = next_page + next_page = self._strategy.next_page(body, current) + + +class AsyncNemoPaginatedResponse(Generic[ModelT]): + """Async iterable over all items across paginated API responses. + + Async twin of :class:`NemoPaginatedResponse`:: + + resp = await client.send(list_items()) + async for item in resp: + print(item.name) + """ + + def __init__( + self, + first_http_response: httpx.Response, + model_type: type[ModelT], + request: PreparedRequest, + fetch_page: _AsyncPageFetcher, + strategy: type[PaginationStrategy] | None = None, + ) -> None: + from nemo_platform_plugin.client.types import OffsetPagination + + self._first_response = first_http_response + self._model_type = model_type + self.request = request + self._fetch_page = fetch_page + self._strategy: type[PaginationStrategy] = strategy or OffsetPagination + + @property + def http_response(self) -> httpx.Response: + return self._first_response + + def first_page(self) -> list[ModelT]: + self._first_response.raise_for_status() + body = self._first_response.json() + raw_items = self._strategy.extract_items(body) + return [self._model_type.model_validate(item) for item in raw_items] + + async def __aiter__(self) -> AsyncIterator[ModelT]: + self._first_response.raise_for_status() + body = self._first_response.json() + + raw_items = self._strategy.extract_items(body) + for item in raw_items: + yield self._model_type.model_validate(item) + + next_page = self._strategy.next_page(body, 1) + while next_page is not None: + raw = await self._fetch_page(self.request, next_page) + raw.raise_for_status() + body = raw.json() + raw_items = self._strategy.extract_items(body) + for item in raw_items: + yield self._model_type.model_validate(item) + current = next_page + next_page = self._strategy.next_page(body, current) + + # --------------------------------------------------------------------------- # Errors # --------------------------------------------------------------------------- diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index 2a0ac85b33..f5c32f7fb6 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -3,17 +3,18 @@ """Shared types for the NeMo client infrastructure. -This module contains marker types, TypeVars, and data classes that are -used across the client package. +This module contains marker types, TypeVars, data classes, pagination +strategies, and client-side option definitions used across the client package. """ from __future__ import annotations from collections.abc import AsyncIterable, Iterable from dataclasses import dataclass, replace -from typing import Generic, ParamSpec, TypeVar +from typing import Any, ClassVar, Generic, ParamSpec, Protocol, TypeVar from pydantic import BaseModel +from typing_extensions import TypeVar as TypeVarExt P = ParamSpec("P") ModelT = TypeVar("ModelT", bound=BaseModel) @@ -40,6 +41,143 @@ def ChatEndpoint(body: ChatRequest, *, workspace: str) -> Stream[ChatChunk]: ... """ +# --------------------------------------------------------------------------- +# Pagination strategies +# --------------------------------------------------------------------------- + + +class PaginationStrategy(Protocol): + """Protocol for pagination strategies. + + Pagination strategies control how the client extracts items from a page + response, determines the next page identifier, and builds query params + to fetch the next page. + """ + + @classmethod + def extract_items(cls, response_body: dict) -> list[dict]: ... + + @classmethod + def next_page(cls, response_body: dict, current_page: Any) -> Any | None: ... + + @classmethod + def page_query_params(cls, page: Any) -> dict[str, Any]: ... + + +class OffsetPagination: + """Offset-based pagination using ``page`` query parameter. + + This is the default strategy, matching the standard ``NemoListResponse`` + envelope used by NeMo Platform services:: + + {"data": [...], "pagination": {"page": 1, "total_pages": 5, ...}} + + Subclass to customise field names for non-standard envelopes:: + + class MyPagination(OffsetPagination): + items_field = "results" + page_param = "offset" + """ + + items_field: ClassVar[str] = "data" + page_param: ClassVar[str] = "page" + pagination_field: ClassVar[str] = "pagination" + total_pages_field: ClassVar[str] = "total_pages" + + @classmethod + def extract_items(cls, response_body: dict) -> list[dict]: + return response_body.get(cls.items_field, []) + + @classmethod + def next_page(cls, response_body: dict, current_page: int) -> int | None: + pagination = response_body.get(cls.pagination_field) + if pagination is None: + return None + total = pagination.get(cls.total_pages_field, 1) + if current_page < total: + return current_page + 1 + return None + + @classmethod + def page_query_params(cls, page: int) -> dict[str, int]: + return {cls.page_param: page} + + +StrategyT = TypeVarExt("StrategyT", default=OffsetPagination) + + +class Paginated(Generic[ModelT, StrategyT]): + """Marker type: endpoint returns paginated results of ``ModelT``. + + The second type parameter selects the pagination strategy. It defaults + to :class:`OffsetPagination`, which matches the standard + ``NemoListResponse`` envelope. + + Usage:: + + # Default offset-based pagination + @get("/apis/example/v2/workspaces/{workspace}/items") + def list_items(...) -> Paginated[Item]: ... + + # Cursor-based pagination + @get("/apis/example/v2/workspaces/{workspace}/logs") + def list_logs(...) -> Paginated[LogEntry, CursorPagination]: ... + + # Custom strategy + class MyPagination(OffsetPagination): + items_field = "results" + page_param = "offset" + + @get("/apis/example/v2/workspaces/{workspace}/widgets") + def list_widgets(...) -> Paginated[Widget, MyPagination]: ... + + Caller experience:: + + for item in client.list_items(): + print(item.name) + """ + + +# --------------------------------------------------------------------------- +# Client-side options (blessed parameter names) +# --------------------------------------------------------------------------- + +# Parameters with these names are recognised in endpoint signatures as +# client-side options. They are stripped from the HTTP request and stashed +# in ``PreparedRequest.client_options`` for the client to act on. +# +# Each entry maps a parameter name to its expected Python type. +# Unknown parameters in an endpoint signature trigger a ``TypeError`` +# at decoration time. +BLESSED_CLIENT_PARAMS: dict[str, type] = { + "exist_ok": bool, +} + + +@dataclass(frozen=True, slots=True) +class RetryPolicy: + """Retry transient failures with exponential backoff. + + Set as a client-level default via the ``retry`` constructor parameter, + or override per-request via ``send()``'s ``retry`` keyword argument. + + This is an operational concern, not a per-endpoint directive — it does + not belong in endpoint signatures. + + Usage:: + + # Client-level default + client = MyClient(base_url="...", retry=RetryPolicy(max_retries=3)) + + # Per-request override via send() + client.send(endpoint_fn(...), retry=RetryPolicy(max_retries=10)) + """ + + max_retries: int = 3 + backoff_base: float = 0.5 + retryable_status_codes: tuple[int, ...] = (502, 503, 504, 429) + + @dataclass(frozen=True, slots=True) class PreparedRequest(Generic[ResponseT]): """A request ready to be sent — carries the endpoint metadata and payload. @@ -57,6 +195,7 @@ class PreparedRequest(Generic[ResponseT]): response_type: type[ResponseT] | None query_params: dict[str, str | int | bool | None] | None = None extra_headers: dict[str, str] | None = None + client_options: dict[str, Any] | None = None def with_headers(self, headers: dict[str, str]) -> PreparedRequest[ResponseT]: """Return a new PreparedRequest with additional headers merged in.""" diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py new file mode 100644 index 0000000000..dc6fad9b0d --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -0,0 +1,272 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for Paginated[T] — automatic pagination via return type marker.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock, call + +import httpx +import pytest +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.endpoint import get +from nemo_platform_plugin.client.method import method +from nemo_platform_plugin.client.response import AsyncNemoPaginatedResponse, NemoPaginatedResponse +from nemo_platform_plugin.client.types import OffsetPagination, Paginated +from pydantic import BaseModel + +BASE = "http://test:8000" + + +class Item(BaseModel): + id: int + name: str + + +@get("/apis/test/v2/workspaces/{workspace}/items") +def LIST_ITEMS(*, workspace: str | None = None) -> Paginated[Item]: + raise NotImplementedError + + +def _page_response(items: list[dict], page: int, total_pages: int, page_size: int = 2) -> httpx.Response: + """Helper to build a paginated response matching NemoListResponse format.""" + return httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/items"), + json={ + "data": items, + "pagination": { + "page": page, + "page_size": page_size, + "current_page_size": len(items), + "total_pages": total_pages, + "total_results": total_pages * page_size, + }, + }, + ) + + +# --------------------------------------------------------------------------- +# Sync +# --------------------------------------------------------------------------- + + +class TestPaginatedSync: + def test_single_page_iteration(self) -> None: + """When total_pages=1, iterating should yield items from only one page.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = _page_response( + [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}], page=1, total_pages=1 + ) + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + + assert isinstance(resp, NemoPaginatedResponse) + items = list(resp) + assert len(items) == 2 + assert items[0].name == "a" + assert items[1].name == "b" + # Only one request made (no additional page fetches) + assert mock_http.request.call_count == 1 + + def test_multi_page_iteration(self) -> None: + """Iterating should automatically fetch all pages.""" + mock_http = MagicMock(spec=httpx.Client) + # First call (page 1) via send() + mock_http.request.side_effect = [ + _page_response([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}], page=1, total_pages=3), + # Pages 2 and 3 fetched by _fetch_page + _page_response([{"id": 3, "name": "c"}, {"id": 4, "name": "d"}], page=2, total_pages=3), + _page_response([{"id": 5, "name": "e"}], page=3, total_pages=3), + ] + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + + items = list(resp) + assert len(items) == 5 + assert [i.name for i in items] == ["a", "b", "c", "d", "e"] + assert mock_http.request.call_count == 3 + + def test_first_page_method(self) -> None: + """first_page() returns items from the already-fetched first page.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = _page_response( + [{"id": 1, "name": "a"}], page=1, total_pages=5 + ) + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + + first = resp.first_page() + assert len(first) == 1 + assert first[0].name == "a" + # No additional requests for first_page() + assert mock_http.request.call_count == 1 + + def test_empty_page(self) -> None: + """Empty data list should yield nothing.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = _page_response([], page=1, total_pages=1) + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + + items = list(resp) + assert items == [] + + def test_no_pagination_metadata(self) -> None: + """When pagination is None, treat as single page.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/items"), + json={"data": [{"id": 1, "name": "a"}]}, + ) + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + + items = list(resp) + assert len(items) == 1 + assert mock_http.request.call_count == 1 + + def test_page_query_param_passed_on_subsequent_pages(self) -> None: + """Subsequent page fetches should include page=N in query params.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + _page_response([{"id": 1, "name": "a"}], page=1, total_pages=2), + _page_response([{"id": 2, "name": "b"}], page=2, total_pages=2), + ] + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_ITEMS()) + list(resp) # consume all pages + + # Second call should have page=2 in params + second_call_params = mock_http.request.call_args_list[1][1].get("params") or mock_http.request.call_args_list[1][0] + assert second_call_params.get("page") == 2 if isinstance(second_call_params, dict) else True + + +# --------------------------------------------------------------------------- +# Via method() descriptor +# --------------------------------------------------------------------------- + + +class TestPaginatedViaMethod: + def test_method_descriptor_returns_paginated_response(self) -> None: + """method() wrapping a Paginated endpoint should return NemoPaginatedResponse.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = _page_response( + [{"id": 1, "name": "a"}], page=1, total_pages=1 + ) + + class _Methods: + list_items = method(LIST_ITEMS) + + class TestClient(_Methods, NemoClient): + pass + + client = TestClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.list_items() + + # Directives are applied but shouldn't break pagination + items = list(resp) + assert len(items) == 1 + assert items[0].name == "a" + + +# --------------------------------------------------------------------------- +# Async +# --------------------------------------------------------------------------- + + +class TestPaginatedAsync: + @pytest.mark.asyncio + async def test_async_multi_page_iteration(self) -> None: + """Async iteration should automatically fetch all pages.""" + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.side_effect = [ + _page_response([{"id": 1, "name": "a"}], page=1, total_pages=2), + _page_response([{"id": 2, "name": "b"}], page=2, total_pages=2), + ] + + client = AsyncNemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = await client.send(LIST_ITEMS()) + + assert isinstance(resp, AsyncNemoPaginatedResponse) + items = [item async for item in resp] + assert len(items) == 2 + assert items[0].name == "a" + assert items[1].name == "b" + + +# --------------------------------------------------------------------------- +# Custom pagination strategy +# --------------------------------------------------------------------------- + + +class ResultsPagination(OffsetPagination): + """Custom strategy: items in 'results', page param is 'offset'.""" + + items_field = "results" + page_param = "offset" + + +@get("/apis/test/v2/workspaces/{workspace}/things") +def LIST_THINGS(*, workspace: str | None = None) -> Paginated[Item, ResultsPagination]: + raise NotImplementedError + + +class TestCustomStrategy: + def test_custom_items_field(self) -> None: + """Custom strategy should extract items from the configured field.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), + json={ + "results": [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}], + "pagination": {"page": 1, "page_size": 10, "current_page_size": 2, "total_pages": 1, "total_results": 2}, + }, + ) + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = client.send(LIST_THINGS()) + + items = list(resp) + assert len(items) == 2 + assert items[0].name == "a" + + def test_custom_page_param(self) -> None: + """Custom strategy should use the configured page query param.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), + json={ + "results": [{"id": 1, "name": "a"}], + "pagination": {"page": 1, "page_size": 1, "current_page_size": 1, "total_pages": 2, "total_results": 2}, + }, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), + json={ + "results": [{"id": 2, "name": "b"}], + "pagination": {"page": 2, "page_size": 1, "current_page_size": 1, "total_pages": 2, "total_results": 2}, + }, + ), + ] + + client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) + items = list(client.send(LIST_THINGS())) + + assert len(items) == 2 + # Verify the second call used "offset" not "page" + second_call_params = mock_http.request.call_args_list[1][1]["params"] + assert "offset" in second_call_params + assert second_call_params["offset"] == 2 diff --git a/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py b/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py index d9a16e9d3a..db1fa33660 100644 --- a/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py +++ b/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py @@ -10,25 +10,18 @@ from __future__ import annotations from abc import abstractmethod -from typing import NotRequired, TypedDict from nemo_example_plugin.entities import ExampleItem from nemo_example_plugin.types.payloads import ( BlobUploadResponse, CountRequest, CreateExampleItemRequest, - ExampleItemPage, HelloResponse, Tick, UpdateExampleItemRequest, ) from nemo_platform_plugin.client.endpoint import delete, get, patch, post, put -from nemo_platform_plugin.client.types import BinaryContent, Stream - - -class ListItemsQueryParams(TypedDict, total=False): - page: NotRequired[int] - page_size: NotRequired[int] +from nemo_platform_plugin.client.types import BinaryContent, Paginated, Stream @get("/apis/example/hello/{name}") @@ -38,14 +31,14 @@ def hello(*, name: str) -> HelloResponse: ... @post("/apis/example/v2/workspaces/{workspace}/items") @abstractmethod -def create_item(*, workspace: str | None = None, body: CreateExampleItemRequest) -> ExampleItem: ... +def create_item( + *, workspace: str | None = None, body: CreateExampleItemRequest, exist_ok: bool = False +) -> ExampleItem: ... @get("/apis/example/v2/workspaces/{workspace}/items") @abstractmethod -def list_items( - *, workspace: str | None = None, query_params: ListItemsQueryParams | None = None -) -> ExampleItemPage: ... +def list_items(*, workspace: str | None = None) -> Paginated[ExampleItem]: ... @get("/apis/example/v2/workspaces/{workspace}/items/{name}") From 6ee1880fc9931d10938f46e975e78a8eae771a71 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Wed, 24 Jun 2026 20:22:31 -0700 Subject: [PATCH 02/11] fix client Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 27 +- .../src/nemo_platform_plugin/client/method.py | 14 +- .../src/nemo_platform_plugin/client/types.py | 8 +- .../tests/client/test_client_options.py | 377 ++++++++++++++++++ .../tests/client/test_pagination.py | 2 +- 5 files changed, 407 insertions(+), 21 deletions(-) create mode 100644 packages/nemo_platform_plugin/tests/client/test_client_options.py diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 89a62dcba7..33e8470f33 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -138,6 +138,27 @@ def _resolve_query_params(self, request: PreparedRequest) -> dict[str, str | int filtered = {k: v for k, v in request.query_params.items() if v is not None} return filtered or None + def _apply_client_options(self, request: PreparedRequest, response: NemoResponse) -> NemoResponse: + """Apply blessed client options (e.g. ``exist_ok``) to the response. + + Options are stashed on ``PreparedRequest.client_options`` by the + endpoint decorator and applied here after the HTTP call completes. + """ + if not request.client_options: + return response + + if request.client_options.get("exist_ok"): + if response.http_response.status_code == 409: + body = response.body + if body is None and request.response_type is not None: + try: + body = request.response_type.model_validate(response.http_response.json()) + except Exception: + pass + return NemoResponse(http_response=response.http_response, body=body, request=request) + + return response + class NemoClient(BaseNemoClient): """Sync HTTP client for NeMo Platform APIs.""" @@ -269,7 +290,8 @@ def _send_once( body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) - return NemoResponse(http_response=raw, body=body, request=request) + response = NemoResponse(http_response=raw, body=body, request=request) + return self._apply_client_options(request, response) def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _SyncPageFetcher: """Create a page-fetching callback bound to this client and strategy.""" @@ -436,7 +458,8 @@ async def _send_once( body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) - return NemoResponse(http_response=raw, body=body, request=request) + response = NemoResponse(http_response=raw, body=body, request=request) + return self._apply_client_options(request, response) def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _AsyncPageFetcher: """Create an async page-fetching callback bound to this client and strategy.""" diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py index 7e60afe981..3c53d074bb 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py @@ -17,8 +17,6 @@ class AsyncExampleClient(_ExampleMethods, AsyncNemoClient): pass resp = client.hello(name="alice") # NemoResponse[HelloResponse] The descriptor dispatches sync vs async based on the client type. -Client-side options (e.g. ``exist_ok``) declared in the endpoint -signature are applied after the HTTP call, wrapping the response. Note: ``ty`` shows ``Unknown |`` on the method types due to unannotated class attributes (astral-sh/ty#3254). The types themselves are correct @@ -32,7 +30,6 @@ class attributes (astral-sh/ty#3254). The types themselves are correct from typing import Any, Coroutine, Generic, overload from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient -from nemo_platform_plugin.client.directives import apply_client_options from nemo_platform_plugin.client.response import NemoResponse from nemo_platform_plugin.client.types import P, PreparedRequest, ResponseT @@ -43,9 +40,6 @@ class EndpointMethod(Generic[P, ResponseT]): When accessed on a :class:`NemoClient`, returns a sync callable. When accessed on an :class:`AsyncNemoClient`, returns an async callable. Both preserve the endpoint's full ``ParamSpec`` signature. - - Client-side options (blessed parameter names like ``exist_ok``) - are automatically applied after the HTTP response is received. """ def __init__(self, endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> None: @@ -64,17 +58,13 @@ def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None @functools.wraps(self._endpoint_fn) async def async_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: - request = self._endpoint_fn(*args, **kwargs) - response = await obj.send(request) - return apply_client_options(request, response) + return await obj.send(self._endpoint_fn(*args, **kwargs)) return async_bound @functools.wraps(self._endpoint_fn) def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: - request = self._endpoint_fn(*args, **kwargs) - response = obj.send(request) # type: ignore[assignment] - return apply_client_options(request, response) + return obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] return sync_bound diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index f5c32f7fb6..67f93b9081 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -119,10 +119,6 @@ class Paginated(Generic[ModelT, StrategyT]): @get("/apis/example/v2/workspaces/{workspace}/items") def list_items(...) -> Paginated[Item]: ... - # Cursor-based pagination - @get("/apis/example/v2/workspaces/{workspace}/logs") - def list_logs(...) -> Paginated[LogEntry, CursorPagination]: ... - # Custom strategy class MyPagination(OffsetPagination): items_field = "results" @@ -161,8 +157,8 @@ class RetryPolicy: Set as a client-level default via the ``retry`` constructor parameter, or override per-request via ``send()``'s ``retry`` keyword argument. - This is an operational concern, not a per-endpoint directive — it does - not belong in endpoint signatures. + This is an operational concern — it does not belong in endpoint + signatures. Usage:: diff --git a/packages/nemo_platform_plugin/tests/client/test_client_options.py b/packages/nemo_platform_plugin/tests/client/test_client_options.py new file mode 100644 index 0000000000..f34a5083de --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_client_options.py @@ -0,0 +1,377 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for client-side options (exist_ok), RetryPolicy, and param validation.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.endpoint import delete, get, post +from nemo_platform_plugin.client.method import method +from nemo_platform_plugin.client.types import PreparedRequest, RetryPolicy +from pydantic import BaseModel + +BASE = "http://test:8000" + + +class ItemRequest(BaseModel): + name: str + + +class ItemResponse(BaseModel): + id: int + name: str + + +# --------------------------------------------------------------------------- +# Endpoint definitions with client options +# --------------------------------------------------------------------------- + + +@post("/apis/test/v2/items") +def CREATE_ITEM(body: ItemRequest, *, exist_ok: bool = False) -> ItemResponse: + raise NotImplementedError + + +@get("/apis/test/v2/items/{name}") +def GET_ITEM(*, name: str) -> ItemResponse: + raise NotImplementedError + + +@delete("/apis/test/v2/items/{name}") +def DELETE_ITEM(*, name: str) -> None: + raise NotImplementedError + + +# --------------------------------------------------------------------------- +# exist_ok: stripped from request, stashed in client_options +# --------------------------------------------------------------------------- + + +class TestExistOkOption: + def test_exist_ok_stripped_from_request(self) -> None: + prepared = CREATE_ITEM(ItemRequest(name="alice"), exist_ok=True) + + assert isinstance(prepared, PreparedRequest) + assert prepared.content is not None + assert prepared.client_options is not None + assert prepared.client_options["exist_ok"] is True + + def test_exist_ok_default_false(self) -> None: + prepared = CREATE_ITEM(ItemRequest(name="alice")) + + assert prepared.client_options is not None + assert prepared.client_options["exist_ok"] is False + + def test_endpoint_without_options_has_none(self) -> None: + prepared = GET_ITEM(name="alice") + assert prepared.client_options is None + + +# --------------------------------------------------------------------------- +# exist_ok: applied via send() +# --------------------------------------------------------------------------- + + +class TestExistOkViaSend: + def test_exist_ok_via_send_swallows_409(self) -> None: + """exist_ok should work when calling client.send() directly.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 409, + request=httpx.Request("POST", f"{BASE}/apis/test/v2/items"), + json={"id": 1, "name": "alice"}, + ) + + client = NemoClient(base_url=BASE, http_client=mock_http) + resp = client.send(CREATE_ITEM(ItemRequest(name="alice"), exist_ok=True)) + + assert resp.http_response.status_code == 409 + assert resp.body is not None + assert resp.body.name == "alice" + + +# --------------------------------------------------------------------------- +# exist_ok: applied via EndpointMethod +# --------------------------------------------------------------------------- + + +class TestExistOkViaMethod: + def test_exist_ok_true_swallows_409(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 409, + request=httpx.Request("POST", f"{BASE}/apis/test/v2/items"), + json={"id": 1, "name": "alice"}, + ) + + class _Methods: + create_item = method(CREATE_ITEM) + + class TestClient(_Methods, NemoClient): + pass + + client = TestClient(base_url=BASE, http_client=mock_http) + resp = client.create_item(body=ItemRequest(name="alice"), exist_ok=True) + + assert resp.http_response.status_code == 409 + assert resp.body is not None + assert resp.body.name == "alice" + + def test_exist_ok_false_returns_409_as_is(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 409, + request=httpx.Request("POST", f"{BASE}/apis/test/v2/items"), + json={"detail": "Already exists"}, + ) + + class _Methods: + create_item = method(CREATE_ITEM) + + class TestClient(_Methods, NemoClient): + pass + + client = TestClient(base_url=BASE, http_client=mock_http) + resp = client.create_item(body=ItemRequest(name="alice")) + + assert resp.http_response.status_code == 409 + assert resp.body is None + + def test_exist_ok_true_non_409_passes_through(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 201, + request=httpx.Request("POST", f"{BASE}/apis/test/v2/items"), + json={"id": 1, "name": "alice"}, + ) + + class _Methods: + create_item = method(CREATE_ITEM) + + class TestClient(_Methods, NemoClient): + pass + + client = TestClient(base_url=BASE, http_client=mock_http) + resp = client.create_item(body=ItemRequest(name="alice"), exist_ok=True) + + assert resp.http_response.status_code == 201 + assert resp.body.name == "alice" + + +# --------------------------------------------------------------------------- +# Param validation at decoration time +# --------------------------------------------------------------------------- + + +class TestParamValidation: + def test_unknown_param_raises_at_decoration_time(self) -> None: + with pytest.raises(TypeError, match="unrecognised parameters"): + + @post("/apis/test/v2/items") + def bad_endpoint(body: ItemRequest, *, bogus: str = "oops") -> ItemResponse: + raise NotImplementedError + + def test_blessed_param_is_allowed(self) -> None: + @post("/apis/test/v2/items") + def ok_endpoint(body: ItemRequest, *, exist_ok: bool = False) -> ItemResponse: + raise NotImplementedError + + prepared = ok_endpoint(ItemRequest(name="x")) + assert isinstance(prepared, PreparedRequest) + + def test_path_params_are_allowed(self) -> None: + @get("/items/{workspace}/{name}") + def ok_endpoint(*, workspace: str, name: str) -> ItemResponse: + raise NotImplementedError + + prepared = ok_endpoint(workspace="default", name="x") + assert prepared.path_params == {"workspace": "default", "name": "x"} + + +# --------------------------------------------------------------------------- +# RetryPolicy: client-level default +# --------------------------------------------------------------------------- + + +class TestRetryPolicy: + def test_retry_on_503(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "Service Unavailable"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + resp = client.send(GET_ITEM(name="alice")) + + assert resp.http_response.status_code == 200 + assert resp.body.name == "alice" + assert mock_http.request.call_count == 2 + + def test_retry_exhausted_returns_last_response(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "Service Unavailable"}, + ) + + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + resp = client.send(GET_ITEM(name="alice")) + + assert resp.http_response.status_code == 503 + assert mock_http.request.call_count == 3 + + def test_no_retry_on_non_retryable_status(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 404, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "Not found"}, + ) + + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + resp = client.send(GET_ITEM(name="alice")) + + assert resp.http_response.status_code == 404 + assert mock_http.request.call_count == 1 + + def test_retry_on_transport_error(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.ConnectError("Connection refused"), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + resp = client.send(GET_ITEM(name="alice")) + + assert resp.body.name == "alice" + assert mock_http.request.call_count == 2 + + def test_per_request_retry_overrides_client_default(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "unavailable"}, + ) + + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=5, backoff_base=0.0), + ) + resp = client.send(GET_ITEM(name="alice"), retry=RetryPolicy(max_retries=1, backoff_base=0.0)) + + assert mock_http.request.call_count == 2 + + def test_no_retry_without_policy(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "unavailable"}, + ) + + client = NemoClient(base_url=BASE, http_client=mock_http) + resp = client.send(GET_ITEM(name="alice")) + + assert resp.http_response.status_code == 503 + assert mock_http.request.call_count == 1 + + +# --------------------------------------------------------------------------- +# Async: exist_ok +# --------------------------------------------------------------------------- + + +class TestAsyncExistOk: + @pytest.mark.asyncio + async def test_exist_ok_true_swallows_409_async(self) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.return_value = httpx.Response( + 409, + request=httpx.Request("POST", f"{BASE}/apis/test/v2/items"), + json={"id": 1, "name": "alice"}, + ) + + class _Methods: + create_item = method(CREATE_ITEM) + + class TestAsyncClient(_Methods, AsyncNemoClient): + pass + + client = TestAsyncClient(base_url=BASE, http_client=mock_http) + resp = await client.create_item(body=ItemRequest(name="alice"), exist_ok=True) + + assert resp.http_response.status_code == 409 + assert resp.body is not None + assert resp.body.name == "alice" + + +# --------------------------------------------------------------------------- +# Async: RetryPolicy +# --------------------------------------------------------------------------- + + +class TestAsyncRetryPolicy: + @pytest.mark.asyncio + async def test_retry_on_503_async(self) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.side_effect = [ + httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "unavailable"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + + client = AsyncNemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + resp = await client.send(GET_ITEM(name="alice")) + + assert resp.http_response.status_code == 200 + assert resp.body.name == "alice" + assert mock_http.request.call_count == 2 diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py index dc6fad9b0d..fd8cf4ff93 100644 --- a/packages/nemo_platform_plugin/tests/client/test_pagination.py +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -172,7 +172,7 @@ class TestClient(_Methods, NemoClient): client = TestClient(base_url=BASE, workspace="default", http_client=mock_http) resp = client.list_items() - # Directives are applied but shouldn't break pagination + # Client options are applied but shouldn't break pagination items = list(resp) assert len(items) == 1 assert items[0].name == "a" From b631148d7dffd4d640cf49d70d1d2327ab141c25 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Wed, 24 Jun 2026 20:26:39 -0700 Subject: [PATCH 03/11] lint Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 26 +++++++------ .../nemo_platform_plugin/client/endpoint.py | 4 +- .../tests/client/test_client_options.py | 2 +- .../tests/client/test_pagination.py | 38 +++++++++++++------ 4 files changed, 43 insertions(+), 27 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 33e8470f33..2f5a0e8ab4 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -302,9 +302,7 @@ def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: existing_params = self._resolve_query_params(request) or {} page_params = strategy.page_query_params(page) params = {**existing_params, **page_params} - return self._http.request( - request.method, url, content=request.content, headers=req_headers, params=params - ) + return self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) return fetch @@ -313,19 +311,20 @@ def _send_with_retry( ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: """Execute a request with retry logic for transient failures.""" last_response: NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse | None = None - last_exc: Exception | None = None for attempt in range(policy.max_retries + 1): try: response = self._send_once(request) - except httpx.TransportError as exc: - last_exc = exc + except httpx.TransportError: if attempt < policy.max_retries: time.sleep(policy.backoff_base * (2**attempt)) continue raise - if isinstance(response, NemoResponse) and response.http_response.status_code in policy.retryable_status_codes: + if ( + isinstance(response, NemoResponse) + and response.http_response.status_code in policy.retryable_status_codes + ): last_response = response if attempt < policy.max_retries: time.sleep(policy.backoff_base * (2**attempt)) @@ -482,20 +481,23 @@ async def _send_with_retry( """Execute a request with retry logic for transient failures.""" import asyncio - last_response: NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse | None = None - last_exc: Exception | None = None + last_response: ( + NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse | None + ) = None for attempt in range(policy.max_retries + 1): try: response = await self._send_once(request) - except httpx.TransportError as exc: - last_exc = exc + except httpx.TransportError: if attempt < policy.max_retries: await asyncio.sleep(policy.backoff_base * (2**attempt)) continue raise - if isinstance(response, NemoResponse) and response.http_response.status_code in policy.retryable_status_codes: + if ( + isinstance(response, NemoResponse) + and response.http_response.status_code in policy.retryable_status_codes + ): last_response = response if attempt < policy.max_retries: await asyncio.sleep(policy.backoff_base * (2**attempt)) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py index 3d0681b41e..35ad1dabf7 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py @@ -58,9 +58,7 @@ def _identify_client_option_params(fn: Callable) -> set[str]: return set(sig.parameters.keys()) & BLESSED_CLIENT_PARAMS.keys() -def _validate_params( - fn: Callable, path_param_names: set[str], client_option_names: set[str] -) -> None: +def _validate_params(fn: Callable, path_param_names: set[str], client_option_names: set[str]) -> None: """Raise ``TypeError`` at decoration time if any parameter is unrecognised. Every parameter must be one of: diff --git a/packages/nemo_platform_plugin/tests/client/test_client_options.py b/packages/nemo_platform_plugin/tests/client/test_client_options.py index f34a5083de..4f85c67391 100644 --- a/packages/nemo_platform_plugin/tests/client/test_client_options.py +++ b/packages/nemo_platform_plugin/tests/client/test_client_options.py @@ -295,7 +295,7 @@ def test_per_request_retry_overrides_client_default(self) -> None: http_client=mock_http, retry=RetryPolicy(max_retries=5, backoff_base=0.0), ) - resp = client.send(GET_ITEM(name="alice"), retry=RetryPolicy(max_retries=1, backoff_base=0.0)) + client.send(GET_ITEM(name="alice"), retry=RetryPolicy(max_retries=1, backoff_base=0.0)) assert mock_http.request.call_count == 2 diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py index fd8cf4ff93..9c005bc3bc 100644 --- a/packages/nemo_platform_plugin/tests/client/test_pagination.py +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -5,7 +5,7 @@ from __future__ import annotations -from unittest.mock import AsyncMock, MagicMock, call +from unittest.mock import AsyncMock, MagicMock import httpx import pytest @@ -93,9 +93,7 @@ def test_multi_page_iteration(self) -> None: def test_first_page_method(self) -> None: """first_page() returns items from the already-fetched first page.""" mock_http = MagicMock(spec=httpx.Client) - mock_http.request.return_value = _page_response( - [{"id": 1, "name": "a"}], page=1, total_pages=5 - ) + mock_http.request.return_value = _page_response([{"id": 1, "name": "a"}], page=1, total_pages=5) client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) resp = client.send(LIST_ITEMS()) @@ -146,7 +144,9 @@ def test_page_query_param_passed_on_subsequent_pages(self) -> None: list(resp) # consume all pages # Second call should have page=2 in params - second_call_params = mock_http.request.call_args_list[1][1].get("params") or mock_http.request.call_args_list[1][0] + second_call_params = ( + mock_http.request.call_args_list[1][1].get("params") or mock_http.request.call_args_list[1][0] + ) assert second_call_params.get("page") == 2 if isinstance(second_call_params, dict) else True @@ -159,9 +159,7 @@ class TestPaginatedViaMethod: def test_method_descriptor_returns_paginated_response(self) -> None: """method() wrapping a Paginated endpoint should return NemoPaginatedResponse.""" mock_http = MagicMock(spec=httpx.Client) - mock_http.request.return_value = _page_response( - [{"id": 1, "name": "a"}], page=1, total_pages=1 - ) + mock_http.request.return_value = _page_response([{"id": 1, "name": "a"}], page=1, total_pages=1) class _Methods: list_items = method(LIST_ITEMS) @@ -229,7 +227,13 @@ def test_custom_items_field(self) -> None: request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), json={ "results": [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}], - "pagination": {"page": 1, "page_size": 10, "current_page_size": 2, "total_pages": 1, "total_results": 2}, + "pagination": { + "page": 1, + "page_size": 10, + "current_page_size": 2, + "total_pages": 1, + "total_results": 2, + }, }, ) @@ -249,7 +253,13 @@ def test_custom_page_param(self) -> None: request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), json={ "results": [{"id": 1, "name": "a"}], - "pagination": {"page": 1, "page_size": 1, "current_page_size": 1, "total_pages": 2, "total_results": 2}, + "pagination": { + "page": 1, + "page_size": 1, + "current_page_size": 1, + "total_pages": 2, + "total_results": 2, + }, }, ), httpx.Response( @@ -257,7 +267,13 @@ def test_custom_page_param(self) -> None: request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/things"), json={ "results": [{"id": 2, "name": "b"}], - "pagination": {"page": 2, "page_size": 1, "current_page_size": 1, "total_pages": 2, "total_results": 2}, + "pagination": { + "page": 2, + "page_size": 1, + "current_page_size": 1, + "total_pages": 2, + "total_results": 2, + }, }, ), ] From b111e0f9943e35bbccab8af271fc7668ce329dda Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Wed, 24 Jun 2026 20:41:39 -0700 Subject: [PATCH 04/11] fix Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 92 +++++++++++++++++-- 1 file changed, 82 insertions(+), 10 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 2f5a0e8ab4..1d0cef7b65 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -17,7 +17,7 @@ import time from collections.abc import Mapping -from typing import TypeVar, get_args, get_origin, overload +from typing import Any, TypeVar, get_args, get_origin, overload import httpx from nemo_platform_plugin.client.response import ( @@ -204,7 +204,7 @@ def send( @overload def send( self, - request: PreparedRequest[Paginated[ModelT]], + request: PreparedRequest[Paginated[ModelT, Any]], *, headers: dict[str, str] | None = None, retry: RetryPolicy | None = None, @@ -284,7 +284,8 @@ def _send_once( assert request.response_type is not None raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) model_type, strategy = _get_paginated_types(request.response_type) - return NemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy), strategy) + retry = self._resolve_retry(None) + return NemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy, retry), strategy) raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) body = None @@ -293,7 +294,9 @@ def _send_once( response = NemoResponse(http_response=raw, body=body, request=request) return self._apply_client_options(request, response) - def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _SyncPageFetcher: + def _make_page_fetcher( + self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None + ) -> _SyncPageFetcher: """Create a page-fetching callback bound to this client and strategy.""" def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: @@ -302,10 +305,40 @@ def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: existing_params = self._resolve_query_params(request) or {} page_params = strategy.page_query_params(page) params = {**existing_params, **page_params} - return self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) + return self._fetch_with_retry(request, url, req_headers, params, retry) return fetch + def _fetch_with_retry( + self, + request: PreparedRequest, + url: str, + headers: dict[str, str] | None, + params: dict, + retry: RetryPolicy | None, + ) -> httpx.Response: + """Execute a single HTTP request with optional retry.""" + if retry is None: + return self._http.request(request.method, url, content=request.content, headers=headers, params=params) + + last_response: httpx.Response | None = None + for attempt in range(retry.max_retries + 1): + try: + raw = self._http.request(request.method, url, content=request.content, headers=headers, params=params) + except httpx.TransportError: + if attempt < retry.max_retries: + time.sleep(retry.backoff_base * (2**attempt)) + continue + raise + if raw.status_code in retry.retryable_status_codes and attempt < retry.max_retries: + last_response = raw + time.sleep(retry.backoff_base * (2**attempt)) + continue + return raw + + assert last_response is not None + return last_response + def _send_with_retry( self, request: PreparedRequest, policy: RetryPolicy ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: @@ -384,7 +417,7 @@ async def send( @overload async def send( self, - request: PreparedRequest[Paginated[ModelT]], + request: PreparedRequest[Paginated[ModelT, Any]], *, headers: dict[str, str] | None = None, retry: RetryPolicy | None = None, @@ -451,7 +484,10 @@ async def _send_once( request.method, url, content=request.content, headers=req_headers, params=params ) model_type, strategy = _get_paginated_types(request.response_type) - return AsyncNemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy), strategy) + retry = self._resolve_retry(None) + return AsyncNemoPaginatedResponse( + raw, model_type, request, self._make_page_fetcher(strategy, retry), strategy + ) raw = await self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) body = None @@ -460,7 +496,9 @@ async def _send_once( response = NemoResponse(http_response=raw, body=body, request=request) return self._apply_client_options(request, response) - def _make_page_fetcher(self, strategy: type[PaginationStrategy]) -> _AsyncPageFetcher: + def _make_page_fetcher( + self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None + ) -> _AsyncPageFetcher: """Create an async page-fetching callback bound to this client and strategy.""" async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: @@ -469,11 +507,45 @@ async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: existing_params = self._resolve_query_params(request) or {} page_params = strategy.page_query_params(page) params = {**existing_params, **page_params} + return await self._fetch_with_retry(request, url, req_headers, params, retry) + + return fetch + + async def _fetch_with_retry( + self, + request: PreparedRequest, + url: str, + headers: dict[str, str] | None, + params: dict, + retry: RetryPolicy | None, + ) -> httpx.Response: + """Execute a single async HTTP request with optional retry.""" + import asyncio + + if retry is None: return await self._http.request( - request.method, url, content=request.content, headers=req_headers, params=params + request.method, url, content=request.content, headers=headers, params=params ) - return fetch + last_response: httpx.Response | None = None + for attempt in range(retry.max_retries + 1): + try: + raw = await self._http.request( + request.method, url, content=request.content, headers=headers, params=params + ) + except httpx.TransportError: + if attempt < retry.max_retries: + await asyncio.sleep(retry.backoff_base * (2**attempt)) + continue + raise + if raw.status_code in retry.retryable_status_codes and attempt < retry.max_retries: + last_response = raw + await asyncio.sleep(retry.backoff_base * (2**attempt)) + continue + return raw + + assert last_response is not None + return last_response async def _send_with_retry( self, request: PreparedRequest, policy: RetryPolicy From c25ed74210660563b1fc59681f25221fc0e02fa8 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 09:15:09 -0700 Subject: [PATCH 05/11] fix retries and add better pagination Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 259 +++++++----------- .../nemo_platform_plugin/client/endpoint.py | 17 +- .../nemo_platform_plugin/client/response.py | 116 +++++--- .../tests/client/test_pagination.py | 52 +++- 4 files changed, 227 insertions(+), 217 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 1d0cef7b65..be8ac7ea81 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -11,10 +11,12 @@ - ``None`` → :class:`~.response.NemoResponse[None]` - ``BinaryContent`` → :class:`~.response.NemoBinaryResponse` - ``Stream[T]`` → :class:`~.response.NemoStreamResponse[T]` +- ``Paginated[T]`` → :class:`~.response.NemoPaginatedResponse[T]` """ from __future__ import annotations +import asyncio import time from collections.abc import Mapping from typing import Any, TypeVar, get_args, get_origin, overload @@ -33,6 +35,7 @@ ) from nemo_platform_plugin.client.types import ( BinaryContent, + OffsetPagination, Paginated, PaginationStrategy, PreparedRequest, @@ -57,8 +60,6 @@ def _get_stream_model_type(response_type: type) -> type[BaseModel]: def _get_paginated_types(response_type: type) -> tuple[type[BaseModel], type]: """Extract (ModelT, StrategyT) from a Paginated[ModelT, StrategyT] generic alias.""" - from nemo_platform_plugin.client.types import OffsetPagination - args = get_args(response_type) if not args: raise TypeError(f"Paginated response type must be parameterized, got {response_type}") @@ -67,6 +68,33 @@ def _get_paginated_types(response_type: type) -> tuple[type[BaseModel], type]: return model_type, strategy_type +# --------------------------------------------------------------------------- +# Retry helper +# --------------------------------------------------------------------------- + + +def _should_retry( + response: httpx.Response | None, + exc: httpx.TransportError | None, + attempt: int, + policy: RetryPolicy, +) -> float | None: + """Decide whether to retry and return the backoff duration, or None to stop. + + Shared decision logic used by both sync and async retry paths. + Returns the sleep duration if a retry should happen, or ``None`` if + the response should be returned / the exception re-raised. + """ + is_last = attempt >= policy.max_retries + if is_last: + return None + if exc is not None: + return policy.backoff_base * (2**attempt) + if response is not None and response.status_code in policy.retryable_status_codes: + return policy.backoff_base * (2**attempt) + return None + + class BaseNemoClient: """Shared logic for sync and async NeMo clients. @@ -97,6 +125,12 @@ def workspace(self) -> str | None: def retry(self) -> RetryPolicy | None: return self._retry + def _resolve_retry(self, retry: RetryPolicy | None) -> RetryPolicy | None: + """Resolve retry policy: per-call override > client default.""" + if retry is not None: + return retry + return self._retry + def _resolve_path(self, request: PreparedRequest) -> str: """Resolve path template with client defaults and explicit params. @@ -179,12 +213,6 @@ def __init__( timeout=timeout, ) - def _resolve_retry(self, retry: RetryPolicy | None) -> RetryPolicy | None: - """Resolve retry policy: per-call override > client default.""" - if retry is not None: - return retry - return self._retry - @overload def send( self, @@ -240,7 +268,7 @@ def send( headers: Optional per-request headers merged on top of client defaults and content-type headers. retry: Optional per-request retry policy override. Takes - precedence over endpoint-level and client-level defaults. + precedence over client-level defaults. For binary and streaming endpoints, the caller should use the response as a context manager to ensure the connection is closed:: @@ -252,19 +280,10 @@ def send( if headers: request = request.with_headers(headers) - resolved_retry = self._resolve_retry(retry) - - if resolved_retry is not None: - return self._send_with_retry(request, resolved_retry) - return self._send_once(request) - - def _send_once( - self, request: PreparedRequest - ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: - """Execute a single HTTP request without retry.""" url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) + resolved_retry = self._resolve_retry(retry) if self._is_binary(request): stream_ctx = self._http.stream( @@ -282,92 +301,63 @@ def _send_once( if self._is_paginated(request): assert request.response_type is not None - raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) + raw = self._request_with_retry(request, url, req_headers, params, resolved_retry) model_type, strategy = _get_paginated_types(request.response_type) - retry = self._resolve_retry(None) - return NemoPaginatedResponse(raw, model_type, request, self._make_page_fetcher(strategy, retry), strategy) + return NemoPaginatedResponse( + raw, model_type, request, self._make_page_fetcher(strategy, resolved_retry), strategy + ) - raw = self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) + raw = self._request_with_retry(request, url, req_headers, params, resolved_retry) body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) response = NemoResponse(http_response=raw, body=body, request=request) return self._apply_client_options(request, response) - def _make_page_fetcher( - self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None - ) -> _SyncPageFetcher: - """Create a page-fetching callback bound to this client and strategy.""" - - def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: - url = self._resolve_path(request) - req_headers = self._request_headers(request) - existing_params = self._resolve_query_params(request) or {} - page_params = strategy.page_query_params(page) - params = {**existing_params, **page_params} - return self._fetch_with_retry(request, url, req_headers, params, retry) - - return fetch - - def _fetch_with_retry( + def _request_with_retry( self, request: PreparedRequest, url: str, headers: dict[str, str] | None, - params: dict, + params: dict | None, retry: RetryPolicy | None, ) -> httpx.Response: """Execute a single HTTP request with optional retry.""" - if retry is None: - return self._http.request(request.method, url, content=request.content, headers=headers, params=params) - last_response: httpx.Response | None = None - for attempt in range(retry.max_retries + 1): + for attempt in range(retry.max_retries + 1 if retry else 1): try: raw = self._http.request(request.method, url, content=request.content, headers=headers, params=params) - except httpx.TransportError: - if attempt < retry.max_retries: - time.sleep(retry.backoff_base * (2**attempt)) + except httpx.TransportError as exc: + backoff = _should_retry(None, exc, attempt, retry) if retry else None + if backoff is not None: + time.sleep(backoff) continue raise - if raw.status_code in retry.retryable_status_codes and attempt < retry.max_retries: - last_response = raw - time.sleep(retry.backoff_base * (2**attempt)) - continue + if retry: + backoff = _should_retry(raw, None, attempt, retry) + if backoff is not None: + last_response = raw + time.sleep(backoff) + continue return raw assert last_response is not None return last_response - def _send_with_retry( - self, request: PreparedRequest, policy: RetryPolicy - ) -> NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse: - """Execute a request with retry logic for transient failures.""" - last_response: NemoResponse | NemoBinaryResponse | NemoStreamResponse | NemoPaginatedResponse | None = None - - for attempt in range(policy.max_retries + 1): - try: - response = self._send_once(request) - except httpx.TransportError: - if attempt < policy.max_retries: - time.sleep(policy.backoff_base * (2**attempt)) - continue - raise - - if ( - isinstance(response, NemoResponse) - and response.http_response.status_code in policy.retryable_status_codes - ): - last_response = response - if attempt < policy.max_retries: - time.sleep(policy.backoff_base * (2**attempt)) - continue + def _make_page_fetcher( + self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None + ) -> _SyncPageFetcher: + """Create a page-fetching callback bound to this client and strategy.""" - return response + def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: + url = self._resolve_path(request) + req_headers = self._request_headers(request) + existing_params = self._resolve_query_params(request) or {} + page_params = strategy.page_query_params(page) + params = {**existing_params, **page_params} + return self._request_with_retry(request, url, req_headers, params, retry) - # All retries exhausted — return the last response we got - assert last_response is not None - return last_response + return fetch class AsyncNemoClient(BaseNemoClient): @@ -392,12 +382,6 @@ def __init__( timeout=timeout, ) - def _resolve_retry(self, retry: RetryPolicy | None) -> RetryPolicy | None: - """Resolve retry policy: per-call override > client default.""" - if retry is not None: - return retry - return self._retry - @overload async def send( self, @@ -450,19 +434,10 @@ async def send( if headers: request = request.with_headers(headers) - resolved_retry = self._resolve_retry(retry) - - if resolved_retry is not None: - return await self._send_with_retry(request, resolved_retry) - return await self._send_once(request) - - async def _send_once( - self, request: PreparedRequest - ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse: - """Execute a single HTTP request without retry.""" url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) + resolved_retry = self._resolve_retry(retry) if self._is_binary(request): stream_ctx = self._http.stream( @@ -480,102 +455,62 @@ async def _send_once( if self._is_paginated(request): assert request.response_type is not None - raw = await self._http.request( - request.method, url, content=request.content, headers=req_headers, params=params - ) + raw = await self._request_with_retry(request, url, req_headers, params, resolved_retry) model_type, strategy = _get_paginated_types(request.response_type) - retry = self._resolve_retry(None) return AsyncNemoPaginatedResponse( - raw, model_type, request, self._make_page_fetcher(strategy, retry), strategy + raw, model_type, request, self._make_page_fetcher(strategy, resolved_retry), strategy ) - raw = await self._http.request(request.method, url, content=request.content, headers=req_headers, params=params) + raw = await self._request_with_retry(request, url, req_headers, params, resolved_retry) body = None if raw.is_success and request.response_type is not None: body = request.response_type.model_validate(raw.json()) response = NemoResponse(http_response=raw, body=body, request=request) return self._apply_client_options(request, response) - def _make_page_fetcher( - self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None - ) -> _AsyncPageFetcher: - """Create an async page-fetching callback bound to this client and strategy.""" - - async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: - url = self._resolve_path(request) - req_headers = self._request_headers(request) - existing_params = self._resolve_query_params(request) or {} - page_params = strategy.page_query_params(page) - params = {**existing_params, **page_params} - return await self._fetch_with_retry(request, url, req_headers, params, retry) - - return fetch - - async def _fetch_with_retry( + async def _request_with_retry( self, request: PreparedRequest, url: str, headers: dict[str, str] | None, - params: dict, + params: dict | None, retry: RetryPolicy | None, ) -> httpx.Response: """Execute a single async HTTP request with optional retry.""" - import asyncio - - if retry is None: - return await self._http.request( - request.method, url, content=request.content, headers=headers, params=params - ) - last_response: httpx.Response | None = None - for attempt in range(retry.max_retries + 1): + for attempt in range(retry.max_retries + 1 if retry else 1): try: raw = await self._http.request( request.method, url, content=request.content, headers=headers, params=params ) - except httpx.TransportError: - if attempt < retry.max_retries: - await asyncio.sleep(retry.backoff_base * (2**attempt)) + except httpx.TransportError as exc: + backoff = _should_retry(None, exc, attempt, retry) if retry else None + if backoff is not None: + await asyncio.sleep(backoff) continue raise - if raw.status_code in retry.retryable_status_codes and attempt < retry.max_retries: - last_response = raw - await asyncio.sleep(retry.backoff_base * (2**attempt)) - continue + if retry: + backoff = _should_retry(raw, None, attempt, retry) + if backoff is not None: + last_response = raw + await asyncio.sleep(backoff) + continue return raw assert last_response is not None return last_response - async def _send_with_retry( - self, request: PreparedRequest, policy: RetryPolicy - ) -> NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse: - """Execute a request with retry logic for transient failures.""" - import asyncio - - last_response: ( - NemoResponse | AsyncNemoBinaryResponse | AsyncNemoStreamResponse | AsyncNemoPaginatedResponse | None - ) = None - - for attempt in range(policy.max_retries + 1): - try: - response = await self._send_once(request) - except httpx.TransportError: - if attempt < policy.max_retries: - await asyncio.sleep(policy.backoff_base * (2**attempt)) - continue - raise - - if ( - isinstance(response, NemoResponse) - and response.http_response.status_code in policy.retryable_status_codes - ): - last_response = response - if attempt < policy.max_retries: - await asyncio.sleep(policy.backoff_base * (2**attempt)) - continue + def _make_page_fetcher( + self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None + ) -> _AsyncPageFetcher: + """Create an async page-fetching callback bound to this client and strategy.""" - return response + async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: + url = self._resolve_path(request) + req_headers = self._request_headers(request) + existing_params = self._resolve_query_params(request) or {} + page_params = strategy.page_query_params(page) + params = {**existing_params, **page_params} + return await self._request_with_retry(request, url, req_headers, params, retry) - assert last_response is not None - return last_response + return fetch diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py index 35ad1dabf7..9c1357b1bd 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py @@ -65,20 +65,31 @@ def _validate_params(fn: Callable, path_param_names: set[str], client_option_nam - ``self`` - A path placeholder (``{name}`` in the URL template) - ``body``, ``content``, or ``query_params`` - - A blessed client option (e.g. ``exist_ok``) + - A blessed client option (e.g. ``exist_ok``) with the correct type """ sig = inspect.signature(fn) known = _RESERVED_PARAM_NAMES | path_param_names | client_option_names unknown = set(sig.parameters.keys()) - known + fn_name = getattr(fn, "__qualname__", getattr(fn, "__name__", repr(fn))) if unknown: - name = getattr(fn, "__qualname__", getattr(fn, "__name__", repr(fn))) blessed = ", ".join(sorted(BLESSED_CLIENT_PARAMS.keys())) raise TypeError( - f"Endpoint {name} has unrecognised parameters: {unknown}. " + f"Endpoint {fn_name} has unrecognised parameters: {unknown}. " f"Parameters must be path params {path_param_names}, " f"'body', 'content', 'query_params', or a client option ({blessed})." ) + # Validate that blessed client option params have the expected type annotation. + hints = get_type_hints(fn) + for param_name in client_option_names: + expected_type = BLESSED_CLIENT_PARAMS[param_name] + actual_type = hints.get(param_name) + if actual_type is not None and actual_type is not expected_type: + raise TypeError( + f"Endpoint {fn_name}: client option '{param_name}' must be " + f"annotated as '{expected_type.__name__}', got '{actual_type}'." + ) + def _build_prepared_request( method: str, diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py index c1d794275c..01a8d76108 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py @@ -12,7 +12,7 @@ from typing import Any, Generic, TypeVar import httpx -from nemo_platform_plugin.client.types import PaginationStrategy, PreparedRequest +from nemo_platform_plugin.client.types import OffsetPagination, PaginationStrategy, PreparedRequest from pydantic import BaseModel ResponseT = TypeVar("ResponseT") @@ -227,6 +227,27 @@ async def __aexit__( _AsyncPageFetcher = Callable[[PreparedRequest, Any], Coroutine[Any, Any, httpx.Response]] +@dataclass(frozen=True, slots=True) +class PageResult(Generic[ModelT]): + """A single page of results with pagination metadata. + + Returned by :meth:`NemoPaginatedResponse.data` for callers who want + one page at a time rather than auto-iterating all pages:: + + resp = client.send(list_items()) + page = resp.data() + print(f"Page {page.page} of {page.total_pages} ({page.total_results} total)") + for item in page.items: + print(item.name) + """ + + items: list[ModelT] + page: int | None = None + page_size: int | None = None + total_pages: int | None = None + total_results: int | None = None + + class NemoPaginatedResponse(Generic[ModelT]): """Sync iterable over all items across paginated API responses. @@ -234,14 +255,13 @@ class NemoPaginatedResponse(Generic[ModelT]): on the endpoint's ``Paginated[T, Strategy]`` return type. Iterating yields individual ``ModelT`` items, not page envelopes:: - resp = client.send(list_items()) - for item in resp: + for item in client.send(list_items()): print(item.name) - Also supports fetching a single page:: + For single-page access with metadata, use :meth:`data`:: - resp = client.send(list_items()) - page = resp.first_page() # list[ModelT] from the first page + page = client.send(list_items()).data() + print(f"{page.total_results} total across {page.total_pages} pages") """ def __init__( @@ -252,8 +272,6 @@ def __init__( fetch_page: _SyncPageFetcher, strategy: type[PaginationStrategy] | None = None, ) -> None: - from nemo_platform_plugin.client.types import OffsetPagination - self._first_response = first_http_response self._model_type = model_type self.request = request @@ -264,30 +282,33 @@ def __init__( def http_response(self) -> httpx.Response: return self._first_response - def _parse_items(self, raw: httpx.Response) -> list[ModelT]: + def _parse_page(self, raw: httpx.Response) -> tuple[list[ModelT], dict]: + """Parse a page response into (items, raw_body).""" raw.raise_for_status() body = raw.json() - raw_items = self._strategy.extract_items(body) - return [self._model_type.model_validate(item) for item in raw_items] - - def first_page(self) -> list[ModelT]: - """Return items from the first page (already fetched).""" - return self._parse_items(self._first_response) + items = [self._model_type.model_validate(item) for item in self._strategy.extract_items(body)] + return items, body + + def data(self) -> PageResult[ModelT]: + """Return the first page as a :class:`PageResult` with metadata.""" + items, body = self._parse_page(self._first_response) + pagination = body.get("pagination") or {} + return PageResult( + items=items, + page=pagination.get("page"), + page_size=pagination.get("page_size"), + total_pages=pagination.get("total_pages"), + total_results=pagination.get("total_results"), + ) def __iter__(self) -> Iterator[ModelT]: - self._first_response.raise_for_status() - body = self._first_response.json() - - raw_items = self._strategy.extract_items(body) - yield from (self._model_type.model_validate(item) for item in raw_items) + items, body = self._parse_page(self._first_response) + yield from items next_page = self._strategy.next_page(body, 1) while next_page is not None: - raw = self._fetch_page(self.request, next_page) - raw.raise_for_status() - body = raw.json() - raw_items = self._strategy.extract_items(body) - yield from (self._model_type.model_validate(item) for item in raw_items) + items, body = self._parse_page(self._fetch_page(self.request, next_page)) + yield from items current = next_page next_page = self._strategy.next_page(body, current) @@ -297,8 +318,7 @@ class AsyncNemoPaginatedResponse(Generic[ModelT]): Async twin of :class:`NemoPaginatedResponse`:: - resp = await client.send(list_items()) - async for item in resp: + async for item in await client.send(list_items()): print(item.name) """ @@ -310,8 +330,6 @@ def __init__( fetch_page: _AsyncPageFetcher, strategy: type[PaginationStrategy] | None = None, ) -> None: - from nemo_platform_plugin.client.types import OffsetPagination - self._first_response = first_http_response self._model_type = model_type self.request = request @@ -322,28 +340,36 @@ def __init__( def http_response(self) -> httpx.Response: return self._first_response - def first_page(self) -> list[ModelT]: - self._first_response.raise_for_status() - body = self._first_response.json() - raw_items = self._strategy.extract_items(body) - return [self._model_type.model_validate(item) for item in raw_items] + def _parse_page(self, raw: httpx.Response) -> tuple[list[ModelT], dict]: + """Parse a page response into (items, raw_body).""" + raw.raise_for_status() + body = raw.json() + items = [self._model_type.model_validate(item) for item in self._strategy.extract_items(body)] + return items, body + + def data(self) -> PageResult[ModelT]: + """Return the first page as a :class:`PageResult` with metadata.""" + items, body = self._parse_page(self._first_response) + pagination = body.get("pagination") or {} + return PageResult( + items=items, + page=pagination.get("page"), + page_size=pagination.get("page_size"), + total_pages=pagination.get("total_pages"), + total_results=pagination.get("total_results"), + ) async def __aiter__(self) -> AsyncIterator[ModelT]: - self._first_response.raise_for_status() - body = self._first_response.json() - - raw_items = self._strategy.extract_items(body) - for item in raw_items: - yield self._model_type.model_validate(item) + items, body = self._parse_page(self._first_response) + for item in items: + yield item next_page = self._strategy.next_page(body, 1) while next_page is not None: raw = await self._fetch_page(self.request, next_page) - raw.raise_for_status() - body = raw.json() - raw_items = self._strategy.extract_items(body) - for item in raw_items: - yield self._model_type.model_validate(item) + items, body = self._parse_page(raw) + for item in items: + yield item current = next_page next_page = self._strategy.next_page(body, current) diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py index 9c005bc3bc..d01e0097d1 100644 --- a/packages/nemo_platform_plugin/tests/client/test_pagination.py +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -13,7 +13,7 @@ from nemo_platform_plugin.client.endpoint import get from nemo_platform_plugin.client.method import method from nemo_platform_plugin.client.response import AsyncNemoPaginatedResponse, NemoPaginatedResponse -from nemo_platform_plugin.client.types import OffsetPagination, Paginated +from nemo_platform_plugin.client.types import OffsetPagination, Paginated, RetryPolicy from pydantic import BaseModel BASE = "http://test:8000" @@ -90,18 +90,22 @@ def test_multi_page_iteration(self) -> None: assert [i.name for i in items] == ["a", "b", "c", "d", "e"] assert mock_http.request.call_count == 3 - def test_first_page_method(self) -> None: - """first_page() returns items from the already-fetched first page.""" + def test_data_returns_page_result_with_metadata(self) -> None: + """data() returns a PageResult with items and pagination metadata.""" mock_http = MagicMock(spec=httpx.Client) mock_http.request.return_value = _page_response([{"id": 1, "name": "a"}], page=1, total_pages=5) client = NemoClient(base_url=BASE, workspace="default", http_client=mock_http) resp = client.send(LIST_ITEMS()) - first = resp.first_page() - assert len(first) == 1 - assert first[0].name == "a" - # No additional requests for first_page() + page = resp.data() + assert len(page.items) == 1 + assert page.items[0].name == "a" + assert page.page == 1 + assert page.total_pages == 5 + assert page.total_results == 10 + assert page.page_size == 2 + # No additional requests for data() assert mock_http.request.call_count == 1 def test_empty_page(self) -> None: @@ -201,6 +205,40 @@ async def test_async_multi_page_iteration(self) -> None: assert items[1].name == "b" +# --------------------------------------------------------------------------- +# Retry on subsequent pages +# --------------------------------------------------------------------------- + + +class TestPaginatedRetry: + def test_retry_on_subsequent_page_503(self) -> None: + """A 503 on page 2 should be retried when a retry policy is configured.""" + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + # Page 1: success + _page_response([{"id": 1, "name": "a"}], page=1, total_pages=2), + # Page 2: first attempt 503, second attempt success + httpx.Response( + 503, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/workspaces/default/items"), + json={"detail": "unavailable"}, + ), + _page_response([{"id": 2, "name": "b"}], page=2, total_pages=2), + ] + + client = NemoClient( + base_url=BASE, + workspace="default", + http_client=mock_http, + retry=RetryPolicy(max_retries=2, backoff_base=0.0), + ) + items = list(client.send(LIST_ITEMS())) + + assert len(items) == 2 + assert [i.name for i in items] == ["a", "b"] + assert mock_http.request.call_count == 3 # page 1 + page 2 fail + page 2 retry + + # --------------------------------------------------------------------------- # Custom pagination strategy # --------------------------------------------------------------------------- From bd67bb728669a17e1cd3d73dd50c7bff04085252 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 09:22:56 -0700 Subject: [PATCH 06/11] make pagination more generic Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/client.py | 10 +++---- .../nemo_platform_plugin/client/response.py | 28 ++++++------------- .../src/nemo_platform_plugin/client/types.py | 20 +++++++++++-- .../tests/client/test_pagination.py | 15 ++++++++++ 4 files changed, 46 insertions(+), 27 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index be8ac7ea81..4ee332ed62 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -26,12 +26,12 @@ AsyncNemoBinaryResponse, AsyncNemoPaginatedResponse, AsyncNemoStreamResponse, + AsyncPageFetcher, NemoBinaryResponse, NemoPaginatedResponse, NemoResponse, NemoStreamResponse, - _AsyncPageFetcher, - _SyncPageFetcher, + SyncPageFetcher, ) from nemo_platform_plugin.client.types import ( BinaryContent, @@ -187,7 +187,7 @@ def _apply_client_options(self, request: PreparedRequest, response: NemoResponse if body is None and request.response_type is not None: try: body = request.response_type.model_validate(response.http_response.json()) - except Exception: + except (ValueError, TypeError): pass return NemoResponse(http_response=response.http_response, body=body, request=request) @@ -346,7 +346,7 @@ def _request_with_retry( def _make_page_fetcher( self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None - ) -> _SyncPageFetcher: + ) -> SyncPageFetcher: """Create a page-fetching callback bound to this client and strategy.""" def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: @@ -502,7 +502,7 @@ async def _request_with_retry( def _make_page_fetcher( self, strategy: type[PaginationStrategy], retry: RetryPolicy | None = None - ) -> _AsyncPageFetcher: + ) -> AsyncPageFetcher: """Create an async page-fetching callback bound to this client and strategy.""" async def fetch(request: PreparedRequest, page: int | str) -> httpx.Response: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py index 01a8d76108..92f7971bc1 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py @@ -223,8 +223,8 @@ async def __aexit__( # Type aliases for the page-fetching callbacks used by paginated responses. # The page value is int for offset-based or str for cursor-based pagination. -_SyncPageFetcher = Callable[[PreparedRequest, Any], httpx.Response] -_AsyncPageFetcher = Callable[[PreparedRequest, Any], Coroutine[Any, Any, httpx.Response]] +SyncPageFetcher = Callable[[PreparedRequest, Any], httpx.Response] +AsyncPageFetcher = Callable[[PreparedRequest, Any], Coroutine[Any, Any, httpx.Response]] @dataclass(frozen=True, slots=True) @@ -269,7 +269,7 @@ def __init__( first_http_response: httpx.Response, model_type: type[ModelT], request: PreparedRequest, - fetch_page: _SyncPageFetcher, + fetch_page: SyncPageFetcher, strategy: type[PaginationStrategy] | None = None, ) -> None: self._first_response = first_http_response @@ -292,14 +292,8 @@ def _parse_page(self, raw: httpx.Response) -> tuple[list[ModelT], dict]: def data(self) -> PageResult[ModelT]: """Return the first page as a :class:`PageResult` with metadata.""" items, body = self._parse_page(self._first_response) - pagination = body.get("pagination") or {} - return PageResult( - items=items, - page=pagination.get("page"), - page_size=pagination.get("page_size"), - total_pages=pagination.get("total_pages"), - total_results=pagination.get("total_results"), - ) + metadata = self._strategy.extract_metadata(body) + return PageResult(items=items, **metadata) def __iter__(self) -> Iterator[ModelT]: items, body = self._parse_page(self._first_response) @@ -327,7 +321,7 @@ def __init__( first_http_response: httpx.Response, model_type: type[ModelT], request: PreparedRequest, - fetch_page: _AsyncPageFetcher, + fetch_page: AsyncPageFetcher, strategy: type[PaginationStrategy] | None = None, ) -> None: self._first_response = first_http_response @@ -350,14 +344,8 @@ def _parse_page(self, raw: httpx.Response) -> tuple[list[ModelT], dict]: def data(self) -> PageResult[ModelT]: """Return the first page as a :class:`PageResult` with metadata.""" items, body = self._parse_page(self._first_response) - pagination = body.get("pagination") or {} - return PageResult( - items=items, - page=pagination.get("page"), - page_size=pagination.get("page_size"), - total_pages=pagination.get("total_pages"), - total_results=pagination.get("total_results"), - ) + metadata = self._strategy.extract_metadata(body) + return PageResult(items=items, **metadata) async def __aiter__(self) -> AsyncIterator[ModelT]: items, body = self._parse_page(self._first_response) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index 67f93b9081..63af5d73cb 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -50,8 +50,8 @@ class PaginationStrategy(Protocol): """Protocol for pagination strategies. Pagination strategies control how the client extracts items from a page - response, determines the next page identifier, and builds query params - to fetch the next page. + response, determines the next page identifier, builds query params + to fetch the next page, and extracts pagination metadata. """ @classmethod @@ -63,6 +63,9 @@ def next_page(cls, response_body: dict, current_page: Any) -> Any | None: ... @classmethod def page_query_params(cls, page: Any) -> dict[str, Any]: ... + @classmethod + def extract_metadata(cls, response_body: dict) -> dict[str, Any]: ... + class OffsetPagination: """Offset-based pagination using ``page`` query parameter. @@ -82,7 +85,10 @@ class MyPagination(OffsetPagination): items_field: ClassVar[str] = "data" page_param: ClassVar[str] = "page" pagination_field: ClassVar[str] = "pagination" + page_field: ClassVar[str] = "page" + page_size_field: ClassVar[str] = "page_size" total_pages_field: ClassVar[str] = "total_pages" + total_results_field: ClassVar[str] = "total_results" @classmethod def extract_items(cls, response_body: dict) -> list[dict]: @@ -102,6 +108,16 @@ def next_page(cls, response_body: dict, current_page: int) -> int | None: def page_query_params(cls, page: int) -> dict[str, int]: return {cls.page_param: page} + @classmethod + def extract_metadata(cls, response_body: dict) -> dict[str, Any]: + pagination = response_body.get(cls.pagination_field) or {} + return { + "page": pagination.get(cls.page_field), + "page_size": pagination.get(cls.page_size_field), + "total_pages": pagination.get(cls.total_pages_field), + "total_results": pagination.get(cls.total_results_field), + } + StrategyT = TypeVarExt("StrategyT", default=OffsetPagination) diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py index d01e0097d1..09ae55f6b7 100644 --- a/packages/nemo_platform_plugin/tests/client/test_pagination.py +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -204,6 +204,21 @@ async def test_async_multi_page_iteration(self) -> None: assert items[0].name == "a" assert items[1].name == "b" + @pytest.mark.asyncio + async def test_async_data_returns_page_result(self) -> None: + """Async data() returns a PageResult with metadata.""" + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.return_value = _page_response([{"id": 1, "name": "a"}], page=1, total_pages=3) + + client = AsyncNemoClient(base_url=BASE, workspace="default", http_client=mock_http) + resp = await client.send(LIST_ITEMS()) + + page = resp.data() + assert len(page.items) == 1 + assert page.page == 1 + assert page.total_pages == 3 + assert mock_http.request.call_count == 1 + # --------------------------------------------------------------------------- # Retry on subsequent pages From 6d7922c53a4695493a5822f907c55a6cc565c835 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 09:42:21 -0700 Subject: [PATCH 07/11] fix types Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/endpoint.py | 6 ++---- .../src/nemo_platform_plugin/client/types.py | 7 ++++++- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py index 9c1357b1bd..aeb4483dfb 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/endpoint.py @@ -38,15 +38,13 @@ def hello(self, *, name: str) -> HelloResponse: from nemo_platform_plugin.client.types import ( BLESSED_CLIENT_PARAMS, + RESERVED_PARAM_NAMES, P, PreparedRequest, ResponseT, ) from pydantic import BaseModel -# Parameter names with special handling in _build_prepared_request. -_RESERVED_PARAM_NAMES = frozenset({"self", "body", "content", "query_params"}) - def _identify_client_option_params(fn: Callable) -> set[str]: """Return parameter names that are blessed client-side options. @@ -68,7 +66,7 @@ def _validate_params(fn: Callable, path_param_names: set[str], client_option_nam - A blessed client option (e.g. ``exist_ok``) with the correct type """ sig = inspect.signature(fn) - known = _RESERVED_PARAM_NAMES | path_param_names | client_option_names + known = RESERVED_PARAM_NAMES | path_param_names | client_option_names unknown = set(sig.parameters.keys()) - known fn_name = getattr(fn, "__qualname__", getattr(fn, "__name__", repr(fn))) if unknown: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index 63af5d73cb..c0707f34d5 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -151,9 +151,14 @@ def list_widgets(...) -> Paginated[Widget, MyPagination]: ... # --------------------------------------------------------------------------- -# Client-side options (blessed parameter names) +# Endpoint parameter registries # --------------------------------------------------------------------------- +# Parameter names with special handling in the endpoint decorator. +# These are routed to specific fields on ``PreparedRequest`` (body, content, +# query_params) and are not treated as path parameters. +RESERVED_PARAM_NAMES: frozenset[str] = frozenset({"self", "body", "content", "query_params"}) + # Parameters with these names are recognised in endpoint signatures as # client-side options. They are stripped from the HTTP request and stashed # in ``PreparedRequest.client_options`` for the client to act on. From f160fe98850fdad491c5b203e9b54b82673530b1 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 10:04:02 -0700 Subject: [PATCH 08/11] fix(example-plugin): update test_sdk.py for Paginated return type list_items now returns Paginated[ExampleItem] instead of ExampleItemPage. Update tests to use PageResult.items instead of .data, and remove the query_params test since pagination is now handled internally by the client. Co-Authored-By: Claude Opus 4.6 (1M context) Signed-off-by: Matthew Grossman --- plugins/example-plugin/tests/test_sdk.py | 22 ++++------------------ 1 file changed, 4 insertions(+), 18 deletions(-) diff --git a/plugins/example-plugin/tests/test_sdk.py b/plugins/example-plugin/tests/test_sdk.py index 7c326b4a80..0d070e7b48 100644 --- a/plugins/example-plugin/tests/test_sdk.py +++ b/plugins/example-plugin/tests/test_sdk.py @@ -119,22 +119,8 @@ def test_sync_list_items() -> None: resp = client.list_items() page = resp.data() - assert len(page.data) == 1 - assert page.data[0].name == "my-item" - - -def test_sync_list_items_with_query_params() -> None: - client, mock_http = _sync_client() - mock_http.request.return_value = _resp( - 200, {"data": [ITEM_PAYLOAD], "pagination": None, "sort": None, "filter": None} - ) - - resp = client.list_items(query_params={"page": 2, "page_size": 5}) - page = resp.data() - - assert len(page.data) == 1 - _, kwargs = mock_http.request.call_args - assert kwargs["params"] == {"page": 2, "page_size": 5} + assert len(page.items) == 1 + assert page.items[0].name == "my-item" def test_sync_update_item() -> None: @@ -206,8 +192,8 @@ async def test_async_list_items() -> None: resp = await client.list_items() page = resp.data() - assert len(page.data) == 1 - assert page.data[0].name == "my-item" + assert len(page.items) == 1 + assert page.items[0].name == "my-item" @pytest.mark.asyncio From 115c37d0905215aba9aea2fc9fee4e940fad8b39 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 10:09:05 -0700 Subject: [PATCH 09/11] fix(client): add PaginatedEndpointMethod for correct return types EndpointMethod.__get__ always returned Callable[P, NemoResponse[ResponseT]], which meant Paginated endpoints returned NemoResponse[Paginated[...]] to the type checker instead of NemoPaginatedResponse[T]. This caused ty to reject .data().items on paginated responses. Add PaginatedEndpointMethod descriptor that returns NemoPaginatedResponse[T], and overload method() to select it automatically for Paginated endpoints. Co-Authored-By: Claude Opus 4.6 (1M context) Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/method.py | 74 ++++++++++++++++++- 1 file changed, 70 insertions(+), 4 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py index 3c53d074bb..fd9324de81 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py @@ -30,8 +30,12 @@ class attributes (astral-sh/ty#3254). The types themselves are correct from typing import Any, Coroutine, Generic, overload from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient -from nemo_platform_plugin.client.response import NemoResponse -from nemo_platform_plugin.client.types import P, PreparedRequest, ResponseT +from nemo_platform_plugin.client.response import ( + AsyncNemoPaginatedResponse, + NemoPaginatedResponse, + NemoResponse, +) +from nemo_platform_plugin.client.types import ModelT, P, Paginated, PreparedRequest, ResponseT class EndpointMethod(Generic[P, ResponseT]): @@ -69,12 +73,74 @@ def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: return sync_bound -def method(endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> EndpointMethod[P, ResponseT]: - """Create an :class:`EndpointMethod` descriptor from an endpoint method. +class PaginatedEndpointMethod(Generic[P, ModelT]): + """Descriptor for endpoints that return ``Paginated[T]``. + + Like :class:`EndpointMethod` but returns + :class:`~nemo_platform_plugin.client.response.NemoPaginatedResponse` + instead of :class:`~nemo_platform_plugin.client.response.NemoResponse`. + """ + + def __init__(self, endpoint_fn: Callable[P, PreparedRequest]) -> None: + self._endpoint_fn = endpoint_fn + + @overload + def __get__(self, obj: NemoClient, objtype: type | None = None) -> Callable[P, NemoPaginatedResponse[ModelT]]: ... + @overload + def __get__( + self, obj: AsyncNemoClient, objtype: type | None = None + ) -> Callable[P, Coroutine[Any, Any, AsyncNemoPaginatedResponse[ModelT]]]: ... + + def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None = None) -> object: + assert obj is not None + if isinstance(obj, AsyncNemoClient): + + @functools.wraps(self._endpoint_fn) + async def async_bound(*args: P.args, **kwargs: P.kwargs) -> AsyncNemoPaginatedResponse[ModelT]: + return await obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] + + return async_bound + + @functools.wraps(self._endpoint_fn) + def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoPaginatedResponse[ModelT]: + return obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] + + return sync_bound + + +@overload +def method(endpoint_fn: Callable[P, PreparedRequest[Paginated[ModelT, Any]]]) -> PaginatedEndpointMethod[P, ModelT]: ... +@overload +def method(endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> EndpointMethod[P, ResponseT]: ... + + +def method( + endpoint_fn: Callable[P, PreparedRequest[ResponseT]], +) -> EndpointMethod[P, ResponseT] | PaginatedEndpointMethod: + """Create a descriptor from an endpoint function. + + Returns :class:`PaginatedEndpointMethod` for endpoints with a + ``Paginated[T]`` return type, :class:`EndpointMethod` for everything + else. Usage:: class _MyMethods: create_item = method(MyEndpoints.create_item) + list_items = method(MyEndpoints.list_items) # Paginated """ + # Inspect the endpoint's return type at wiring time to pick the right + # descriptor class. + import inspect + from typing import get_args, get_origin + + hints = inspect.signature(endpoint_fn).return_annotation + # The endpoint_fn is already wrapped by @get/@post — its return annotation + # is PreparedRequest[ResponseT]. We need to check the inner ResponseT. + origin = get_origin(hints) + if origin is PreparedRequest: + inner_args = get_args(hints) + if inner_args and get_origin(inner_args[0]) is Paginated: + return PaginatedEndpointMethod(endpoint_fn) # type: ignore[return-value] + return EndpointMethod(endpoint_fn) From bb7e2c47a403a8e1a2cfe7ec1eae3f4e90f7e2c3 Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 10:28:29 -0700 Subject: [PATCH 10/11] fix method.py Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/method.py | 128 +++++++++--------- 1 file changed, 65 insertions(+), 63 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py index fd9324de81..a167f7a8ff 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py @@ -17,6 +17,9 @@ class AsyncExampleClient(_ExampleMethods, AsyncNemoClient): pass resp = client.hello(name="alice") # NemoResponse[HelloResponse] The descriptor dispatches sync vs async based on the client type. +The ``method()`` function is overloaded so that the return type of the +bound callable matches what ``send()`` returns for each response-type +marker (``BinaryContent``, ``Stream[T]``, ``Paginated[T]``, plain model). Note: ``ty`` shows ``Unknown |`` on the method types due to unannotated class attributes (astral-sh/ty#3254). The types themselves are correct @@ -27,120 +30,119 @@ class attributes (astral-sh/ty#3254). The types themselves are correct import functools from collections.abc import Callable -from typing import Any, Coroutine, Generic, overload +from typing import Any, Coroutine, Generic, TypeVar, overload from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_plugin.client.response import ( + AsyncNemoBinaryResponse, AsyncNemoPaginatedResponse, + AsyncNemoStreamResponse, + NemoBinaryResponse, NemoPaginatedResponse, NemoResponse, + NemoStreamResponse, ) -from nemo_platform_plugin.client.types import ModelT, P, Paginated, PreparedRequest, ResponseT +from nemo_platform_plugin.client.types import ( + BinaryContent, + ModelT, + P, + Paginated, + PreparedRequest, + ResponseT, + Stream, +) + +# TypeVar for the sync return type of the bound callable. +SyncReturnT = TypeVar("SyncReturnT") +# TypeVar for the async return type of the bound callable. +AsyncReturnT = TypeVar("AsyncReturnT") -class EndpointMethod(Generic[P, ResponseT]): +class EndpointMethod(Generic[P, SyncReturnT, AsyncReturnT]): """Descriptor that binds an endpoint to a client instance. - When accessed on a :class:`NemoClient`, returns a sync callable. - When accessed on an :class:`AsyncNemoClient`, returns an async callable. - Both preserve the endpoint's full ``ParamSpec`` signature. + When accessed on a :class:`NemoClient`, returns a sync callable + with return type ``SyncReturnT``. + When accessed on an :class:`AsyncNemoClient`, returns an async callable + with return type ``AsyncReturnT``. + + The type parameters are set by the ``method()`` overloads to match + the response type that ``send()`` returns for each endpoint marker. """ - def __init__(self, endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> None: + def __init__(self, endpoint_fn: Callable[P, PreparedRequest]) -> None: self._endpoint_fn = endpoint_fn @overload - def __get__(self, obj: NemoClient, objtype: type | None = None) -> Callable[P, NemoResponse[ResponseT]]: ... + def __get__(self, obj: NemoClient, objtype: type | None = None) -> Callable[P, SyncReturnT]: ... @overload def __get__( self, obj: AsyncNemoClient, objtype: type | None = None - ) -> Callable[P, Coroutine[Any, Any, NemoResponse[ResponseT]]]: ... + ) -> Callable[P, Coroutine[Any, Any, AsyncReturnT]]: ... def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None = None) -> object: assert obj is not None if isinstance(obj, AsyncNemoClient): @functools.wraps(self._endpoint_fn) - async def async_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: - return await obj.send(self._endpoint_fn(*args, **kwargs)) + async def async_bound(*args: P.args, **kwargs: P.kwargs) -> AsyncReturnT: + return await obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] return async_bound @functools.wraps(self._endpoint_fn) - def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoResponse[ResponseT]: + def sync_bound(*args: P.args, **kwargs: P.kwargs) -> SyncReturnT: return obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] return sync_bound -class PaginatedEndpointMethod(Generic[P, ModelT]): - """Descriptor for endpoints that return ``Paginated[T]``. - - Like :class:`EndpointMethod` but returns - :class:`~nemo_platform_plugin.client.response.NemoPaginatedResponse` - instead of :class:`~nemo_platform_plugin.client.response.NemoResponse`. - """ - - def __init__(self, endpoint_fn: Callable[P, PreparedRequest]) -> None: - self._endpoint_fn = endpoint_fn +# --------------------------------------------------------------------------- +# method() overloads — one per response-type marker +# --------------------------------------------------------------------------- - @overload - def __get__(self, obj: NemoClient, objtype: type | None = None) -> Callable[P, NemoPaginatedResponse[ModelT]]: ... - @overload - def __get__( - self, obj: AsyncNemoClient, objtype: type | None = None - ) -> Callable[P, Coroutine[Any, Any, AsyncNemoPaginatedResponse[ModelT]]]: ... - def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None = None) -> object: - assert obj is not None - if isinstance(obj, AsyncNemoClient): +@overload +def method( + endpoint_fn: Callable[P, PreparedRequest[BinaryContent]], +) -> EndpointMethod[P, NemoBinaryResponse, AsyncNemoBinaryResponse]: ... - @functools.wraps(self._endpoint_fn) - async def async_bound(*args: P.args, **kwargs: P.kwargs) -> AsyncNemoPaginatedResponse[ModelT]: - return await obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] - return async_bound +@overload +def method( + endpoint_fn: Callable[P, PreparedRequest[Stream[ModelT]]], +) -> EndpointMethod[P, NemoStreamResponse[ModelT], AsyncNemoStreamResponse[ModelT]]: ... - @functools.wraps(self._endpoint_fn) - def sync_bound(*args: P.args, **kwargs: P.kwargs) -> NemoPaginatedResponse[ModelT]: - return obj.send(self._endpoint_fn(*args, **kwargs)) # type: ignore[return-value] - return sync_bound +@overload +def method( + endpoint_fn: Callable[P, PreparedRequest[Paginated[ModelT, Any]]], +) -> EndpointMethod[P, NemoPaginatedResponse[ModelT], AsyncNemoPaginatedResponse[ModelT]]: ... @overload -def method(endpoint_fn: Callable[P, PreparedRequest[Paginated[ModelT, Any]]]) -> PaginatedEndpointMethod[P, ModelT]: ... -@overload -def method(endpoint_fn: Callable[P, PreparedRequest[ResponseT]]) -> EndpointMethod[P, ResponseT]: ... +def method( + endpoint_fn: Callable[P, PreparedRequest[None]], +) -> EndpointMethod[P, NemoResponse[None], NemoResponse[None]]: ... +@overload def method( endpoint_fn: Callable[P, PreparedRequest[ResponseT]], -) -> EndpointMethod[P, ResponseT] | PaginatedEndpointMethod: - """Create a descriptor from an endpoint function. +) -> EndpointMethod[P, NemoResponse[ResponseT], NemoResponse[ResponseT]]: ... + - Returns :class:`PaginatedEndpointMethod` for endpoints with a - ``Paginated[T]`` return type, :class:`EndpointMethod` for everything - else. +def method(endpoint_fn: Callable[P, PreparedRequest]) -> EndpointMethod: + """Create an :class:`EndpointMethod` descriptor from an endpoint function. + + The return type of the bound callable is determined by the endpoint's + response-type marker via overloads, matching the dispatch in ``send()``. Usage:: class _MyMethods: - create_item = method(MyEndpoints.create_item) - list_items = method(MyEndpoints.list_items) # Paginated + create_item = method(MyEndpoints.create_item) # NemoResponse[Item] + list_items = method(MyEndpoints.list_items) # NemoPaginatedResponse[Item] + download = method(MyEndpoints.download) # NemoBinaryResponse """ - # Inspect the endpoint's return type at wiring time to pick the right - # descriptor class. - import inspect - from typing import get_args, get_origin - - hints = inspect.signature(endpoint_fn).return_annotation - # The endpoint_fn is already wrapped by @get/@post — its return annotation - # is PreparedRequest[ResponseT]. We need to check the inner ResponseT. - origin = get_origin(hints) - if origin is PreparedRequest: - inner_args = get_args(hints) - if inner_args and get_origin(inner_args[0]) is Paginated: - return PaginatedEndpointMethod(endpoint_fn) # type: ignore[return-value] - return EndpointMethod(endpoint_fn) From 9bf9ebd9df40488356deb623d0607df186aa1f8a Mon Sep 17 00:00:00 2001 From: Matthew Grossman Date: Thu, 25 Jun 2026 11:34:56 -0700 Subject: [PATCH 11/11] code review Signed-off-by: Matthew Grossman --- .../src/nemo_platform_plugin/client/types.py | 8 ++++++++ .../tests/client/test_pagination.py | 6 ++---- .../src/nemo_example_plugin/types/endpoints.py | 10 +++++++++- 3 files changed, 19 insertions(+), 5 deletions(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index c0707f34d5..1e2a6c72b0 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -181,6 +181,14 @@ class RetryPolicy: This is an operational concern — it does not belong in endpoint signatures. + .. note:: + + The retry loop replays the full request including the body. + Callers are responsible for ensuring idempotency when retrying + non-safe methods (POST, PATCH). Consider using an + ``Idempotency-Key`` header for create operations that must be + safe to retry. + Usage:: # Client-level default diff --git a/packages/nemo_platform_plugin/tests/client/test_pagination.py b/packages/nemo_platform_plugin/tests/client/test_pagination.py index 09ae55f6b7..f19bc28705 100644 --- a/packages/nemo_platform_plugin/tests/client/test_pagination.py +++ b/packages/nemo_platform_plugin/tests/client/test_pagination.py @@ -148,10 +148,8 @@ def test_page_query_param_passed_on_subsequent_pages(self) -> None: list(resp) # consume all pages # Second call should have page=2 in params - second_call_params = ( - mock_http.request.call_args_list[1][1].get("params") or mock_http.request.call_args_list[1][0] - ) - assert second_call_params.get("page") == 2 if isinstance(second_call_params, dict) else True + second_call_params = mock_http.request.call_args_list[1][1]["params"] + assert second_call_params["page"] == 2 # --------------------------------------------------------------------------- diff --git a/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py b/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py index db1fa33660..437afbe795 100644 --- a/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py +++ b/plugins/example-plugin/src/nemo_example_plugin/types/endpoints.py @@ -10,6 +10,7 @@ from __future__ import annotations from abc import abstractmethod +from typing import NotRequired, TypedDict from nemo_example_plugin.entities import ExampleItem from nemo_example_plugin.types.payloads import ( @@ -36,9 +37,16 @@ def create_item( ) -> ExampleItem: ... +class ListItemsQueryParams(TypedDict, total=False): + page_size: NotRequired[int] + sort: NotRequired[str] + + @get("/apis/example/v2/workspaces/{workspace}/items") @abstractmethod -def list_items(*, workspace: str | None = None) -> Paginated[ExampleItem]: ... +def list_items( + *, workspace: str | None = None, query_params: ListItemsQueryParams | None = None +) -> Paginated[ExampleItem]: ... @get("/apis/example/v2/workspaces/{workspace}/items/{name}")