mirror of
https://github.com/NVIDIA/OpenShell.git
synced 2026-10-02 07:34:45 +08:00
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:
@@ -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
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user