feat(sdk): add lazy pagination pagers (#3256)

* feat(sdk): add lazy pagination pagers

Signed-off-by: Drew Newberry <anewberry@nvidia.com>

* test(sdk): harden pager edge cases

Signed-off-by: Drew Newberry <anewberry@nvidia.com>

* test(sdk): cover initial resume token

Signed-off-by: Drew Newberry <anewberry@nvidia.com>

* docs(go): fix all-workspaces pager examples

Signed-off-by: Drew Newberry <anewberry@nvidia.com>

---------

Signed-off-by: Drew Newberry <anewberry@nvidia.com>
This commit is contained in:
Drew Newberry
2026-09-11 00:02:08 +00:00
committed by GitHub
parent 33bbda3d33
commit 1860010850
76 changed files with 1532 additions and 597 deletions
+4
View File
@@ -9,6 +9,8 @@ from .sandbox import (
ClientCredentialsAuth,
ExecChunk,
ExecResult,
Page,
Pager,
Sandbox,
SandboxClient,
SandboxError,
@@ -33,6 +35,8 @@ __all__ = [
"ClientCredentialsAuth",
"ExecChunk",
"ExecResult",
"Page",
"Pager",
"Sandbox",
"SandboxClient",
"SandboxError",
+153 -48
View File
@@ -17,7 +17,7 @@ import threading
import time
from collections import namedtuple
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Never, SupportsIndex, cast
from typing import TYPE_CHECKING, Any, Generic, Never, SupportsIndex, TypeVar, cast
from urllib.parse import urlparse
import grpc
@@ -35,6 +35,37 @@ _ClientCallDetailsBase = namedtuple(
)
_OAUTH_MAX_RESPONSE_BYTES = 1 << 20
T = TypeVar("T")
@dataclass(frozen=True)
class Page(Generic[T]):
"""One response page from a list operation."""
items: builtins.list[T]
next_page_token: str
class Pager(Generic[T]):
"""Lazy, single-pass iterator that fetches one RPC page per advance."""
def __init__(self, fetch: Callable[[str], Page[T]], page_token: str = "") -> None:
self._fetch = fetch
self._page_token: str | None = page_token
def __iter__(self) -> Pager[T]:
return self
def __next__(self) -> Page[T]:
if self._page_token is None:
raise StopIteration
page = self._fetch(self._page_token)
self._page_token = page.next_page_token or None
return page
def all(self) -> builtins.list[T]:
"""Consume the pager and collect every remaining item."""
return [item for page in self for item in page.items]
def _workspace_scope(workspace: str) -> datamodel_pb2.WorkspaceSelector:
@@ -820,47 +851,77 @@ class SandboxClient:
*,
workspace: str,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[SandboxRef]:
sandboxes: builtins.list[SandboxRef] = []
page_token = ""
while True:
) -> Pager[SandboxRef]:
def fetch(token: str) -> Page[SandboxRef]:
response = self._stub.ListSandboxes(
openshell_pb2.ListSandboxesRequest(
workspace_scope=_workspace_scope(workspace),
page_size=page_size,
page_token=page_token,
page_token=token,
label_selector=label_selector or "",
),
timeout=self._timeout,
)
sandboxes.extend(_sandbox_ref(item) for item in response.sandboxes)
if not getattr(response, "next_page_token", ""):
return sandboxes
page_token = response.next_page_token
return Page(
items=[_sandbox_ref(item) for item in response.sandboxes],
next_page_token=getattr(response, "next_page_token", ""),
)
return Pager(fetch, page_token)
def list_all(
self,
*,
workspace: str,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[SandboxRef]:
return self.list(
workspace=workspace,
page_size=page_size,
page_token=page_token,
label_selector=label_selector,
).all()
def list_for_all_workspaces(
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[SandboxRef]:
sandboxes: builtins.list[SandboxRef] = []
page_token = ""
while True:
) -> Pager[SandboxRef]:
def fetch(token: str) -> Page[SandboxRef]:
response = self._stub.ListSandboxes(
openshell_pb2.ListSandboxesRequest(
workspace_scope=_all_workspaces_scope(),
page_size=page_size,
page_token=page_token,
page_token=token,
label_selector=label_selector or "",
),
timeout=self._timeout,
)
sandboxes.extend(_sandbox_ref(item) for item in response.sandboxes)
if not getattr(response, "next_page_token", ""):
return sandboxes
page_token = response.next_page_token
return Page(
items=[_sandbox_ref(item) for item in response.sandboxes],
next_page_token=getattr(response, "next_page_token", ""),
)
return Pager(fetch, page_token)
def list_all_for_all_workspaces(
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[SandboxRef]:
return self.list_for_all_workspaces(
page_size=page_size,
page_token=page_token,
label_selector=label_selector,
).all()
def list_ids(
self,
@@ -871,7 +932,7 @@ class SandboxClient:
) -> builtins.list[str]:
return [
item.id
for item in self.list(
for item in self.list_all(
workspace=workspace,
page_size=page_size,
label_selector=label_selector,
@@ -886,7 +947,7 @@ class SandboxClient:
) -> builtins.list[str]:
return [
item.id
for item in self.list_for_all_workspaces(
for item in self.list_all_for_all_workspaces(
page_size=page_size,
label_selector=label_selector,
)
@@ -1192,47 +1253,77 @@ class SandboxTemplateClient:
*,
workspace: str,
page_size: int = 100,
page_token: str = "",
label_selector: str = "",
) -> builtins.list[openshell_pb2.SandboxWorkloadTemplate]:
templates: builtins.list[openshell_pb2.SandboxWorkloadTemplate] = []
page_token = ""
while True:
) -> Pager[openshell_pb2.SandboxWorkloadTemplate]:
def fetch(token: str) -> Page[openshell_pb2.SandboxWorkloadTemplate]:
response = self._stub.ListSandboxTemplates(
openshell_pb2.ListSandboxTemplatesRequest(
workspace_scope=_workspace_scope(workspace),
page_size=page_size,
page_token=page_token,
page_token=token,
label_selector=label_selector,
),
timeout=self._timeout,
)
templates.extend(response.templates)
if not getattr(response, "next_page_token", ""):
return templates
page_token = response.next_page_token
return Page(
items=list(response.templates),
next_page_token=getattr(response, "next_page_token", ""),
)
return Pager(fetch, page_token)
def list_all(
self,
*,
workspace: str,
page_size: int = 100,
page_token: str = "",
label_selector: str = "",
) -> builtins.list[openshell_pb2.SandboxWorkloadTemplate]:
return self.list(
workspace=workspace,
page_size=page_size,
page_token=page_token,
label_selector=label_selector,
).all()
def list_for_all_workspaces(
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str = "",
) -> builtins.list[openshell_pb2.SandboxWorkloadTemplate]:
templates: builtins.list[openshell_pb2.SandboxWorkloadTemplate] = []
page_token = ""
while True:
) -> Pager[openshell_pb2.SandboxWorkloadTemplate]:
def fetch(token: str) -> Page[openshell_pb2.SandboxWorkloadTemplate]:
response = self._stub.ListSandboxTemplates(
openshell_pb2.ListSandboxTemplatesRequest(
workspace_scope=_all_workspaces_scope(),
page_size=page_size,
page_token=page_token,
page_token=token,
label_selector=label_selector,
),
timeout=self._timeout,
)
templates.extend(response.templates)
if not getattr(response, "next_page_token", ""):
return templates
page_token = response.next_page_token
return Page(
items=list(response.templates),
next_page_token=getattr(response, "next_page_token", ""),
)
return Pager(fetch, page_token)
def list_all_for_all_workspaces(
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str = "",
) -> builtins.list[openshell_pb2.SandboxWorkloadTemplate]:
return self.list_for_all_workspaces(
page_size=page_size,
page_token=page_token,
label_selector=label_selector,
).all()
def delete(self, name: str, *, workspace: str) -> bool:
response = self._stub.DeleteSandboxTemplate(
@@ -1297,23 +1388,37 @@ class WorkspaceClient:
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[WorkspaceRef]:
workspaces: builtins.list[WorkspaceRef] = []
page_token = ""
while True:
) -> Pager[WorkspaceRef]:
def fetch(token: str) -> Page[WorkspaceRef]:
response = self._stub.ListWorkspaces(
openshell_pb2.ListWorkspacesRequest(
page_size=page_size,
page_token=page_token,
page_token=token,
label_selector=label_selector or "",
),
timeout=self._timeout,
)
workspaces.extend(_workspace_ref(ws) for ws in response.workspaces)
if not getattr(response, "next_page_token", ""):
return workspaces
page_token = response.next_page_token
return Page(
items=[_workspace_ref(ws) for ws in response.workspaces],
next_page_token=getattr(response, "next_page_token", ""),
)
return Pager(fetch, page_token)
def list_all(
self,
*,
page_size: int = 100,
page_token: str = "",
label_selector: str | None = None,
) -> builtins.list[WorkspaceRef]:
return self.list(
page_size=page_size,
page_token=page_token,
label_selector=label_selector,
).all()
def delete(self, name: str) -> bool:
response = self._stub.DeleteWorkspace(
+50 -10
View File
@@ -23,6 +23,8 @@ from openshell.sandbox import (
_PYTHON_CLOUDPICKLE_BOOTSTRAP,
_SANDBOX_PYTHON_BIN,
ClientCredentialsAuth,
Page,
Pager,
Sandbox,
SandboxClient,
SandboxError,
@@ -2408,7 +2410,7 @@ def test_sandbox_template_client_crud_forwards_requests() -> None:
assert stub.get_template_request.name == "gpu-kata"
assert _request_workspace(stub.get_template_request) == "default"
listed = client.list(
listed = client.list_all(
workspace="default", page_size=50, label_selector="team=runtime"
)
assert len(listed) == 1
@@ -2429,7 +2431,7 @@ def test_sandbox_template_list_for_all_workspaces_selects_all() -> None:
stub = _FakeSandboxStub()
client = _template_client_with_fake_stub(stub)
client.list_for_all_workspaces(page_size=100, label_selector="team=runtime")
client.list_all_for_all_workspaces(page_size=100, label_selector="team=runtime")
assert stub.list_template_request is not None
assert _request_selects_all_workspaces(stub.list_template_request)
@@ -2533,7 +2535,7 @@ def test_list_forwards_label_selector() -> None:
stub = _FakeSandboxStub()
client = _client_with_fake_stub(stub)
client.list(workspace="default", label_selector="aiq=deep-research")
client.list_all(workspace="default", label_selector="aiq=deep-research")
assert stub.list_request is not None
assert stub.list_request.label_selector == "aiq=deep-research"
@@ -2544,7 +2546,7 @@ def test_list_without_selector_sends_empty_string() -> None:
stub = _FakeSandboxStub()
client = _client_with_fake_stub(stub)
client.list(workspace="default")
client.list_all(workspace="default")
assert stub.list_request is not None
assert stub.list_request.label_selector == ""
@@ -2559,17 +2561,55 @@ def test_list_follows_continuation_tokens() -> None:
)
client = _client_with_fake_stub(stub)
sandboxes = client.list(
workspace="default", page_size=1, label_selector="team=core"
)
pager = client.list(workspace="default", page_size=1, label_selector="team=core")
assert [sandbox.name for sandbox in sandboxes] == ["job-1", "job-2"]
assert stub.list_requests == []
first = next(pager)
assert [sandbox.name for sandbox in first.items] == ["job-1"]
assert first.next_page_token == "1"
second = next(pager)
assert [sandbox.name for sandbox in second.items] == ["job-2"]
assert second.next_page_token == ""
with pytest.raises(StopIteration):
next(pager)
assert len(stub.list_requests) == 2
assert stub.list_requests[0].page_token == ""
assert stub.list_requests[1].page_token == "1"
assert stub.list_requests[1].label_selector == "team=core"
def test_list_passes_initial_page_token() -> None:
stub = _FakeSandboxStub(
listed_pages=[
[_make_sandbox_proto("sandbox-1", "skipped")],
[_make_sandbox_proto("sandbox-2", "resumed")],
]
)
client = _client_with_fake_stub(stub)
page = next(client.list(workspace="default", page_token="1"))
assert [sandbox.name for sandbox in page.items] == ["resumed"]
assert stub.list_requests[0].page_token == "1"
def test_pager_retries_same_token_after_fetch_error() -> None:
tokens: list[str] = []
def fetch(token: str) -> Page[int]:
tokens.append(token)
if len(tokens) == 1:
raise RuntimeError("temporary failure")
return Page(items=[1], next_page_token="")
pager = Pager(fetch, page_token="resume")
with pytest.raises(RuntimeError, match="temporary failure"):
next(pager)
assert next(pager).items == [1]
assert tokens == ["resume", "resume"]
def test_list_ids_forwards_label_selector() -> None:
stub = _FakeSandboxStub(listed=[_make_sandbox_proto("sandbox-1", "job-1")])
client = _client_with_fake_stub(stub)
@@ -2804,7 +2844,7 @@ def test_list_for_all_workspaces_sets_flag() -> None:
stub = _FakeSandboxStub()
client = _client_with_fake_stub(stub)
client.list_for_all_workspaces()
client.list_all_for_all_workspaces()
assert stub.list_request is not None
assert _request_selects_all_workspaces(stub.list_request)
@@ -2815,7 +2855,7 @@ def test_list_with_workspace_passes_workspace() -> None:
stub = _FakeSandboxStub()
client = _client_with_fake_stub(stub)
client.list(workspace="staging")
client.list_all(workspace="staging")
assert stub.list_request is not None
assert _request_workspace(stub.list_request) == "staging"