From 21f10944f43b23c2347ef9907c6022cffc8dbeda Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Thu, 10 Sep 2026 17:13:03 -0400 Subject: [PATCH 1/8] feat(client): typed-client plumbing for the CLI migration Adds what hand-written CLI commands need from nemo_platform_plugin: - InferenceGatewayClient (provider/model/openai proxy routes, provider_ready, raw SSE streams, OpenAI model listing) with endpoint and wire tests. - client_from_platform accepts a NemoClient or AsyncNemoClient and derives the typed client with from_client, so callers no longer need to know which platform handle they hold; PlatformClient is the runtime-checkable structural type for that parameter and platform_default_headers reads identity headers off either shape. - models.refs holds the pure model-reference helpers; packages/models re-exports them so the CLI does not import the Stainless-bound package. - filesets.transfer holds SDK-free upload/download/list/delete; FilesResource delegates to it and filesets no longer imports .resources eagerly. - Bearer tokens are resolved on every HTTP attempt rather than baked into the PreparedRequest, so retries and later pages never replay a stale token. - A 409 on a create sent with exist_ok is no longer retried before send() resolves it by fetching the existing entity. - Query-param TypedDicts and request models for files, iam, virtual models and workspaces gain the fields the CLI exposes; GuardrailConfig.data is optional to match the server entity. Signed-off-by: Max Dubrinsky --- packages/filesets/src/filesets/__init__.py | 5 +- packages/filesets/src/filesets/resources.py | 365 +++--------- packages/filesets/src/filesets/transfer.py | 502 ++++++++++++++++ packages/filesets/tests/test_transfer.py | 544 ++++++++++++++++++ packages/models/src/models/resources.py | 72 +-- .../nemo_platform_plugin/client/adapter.py | 109 +++- .../src/nemo_platform_plugin/client/client.py | 91 ++- .../nemo_platform_plugin/files/endpoints.py | 7 +- .../src/nemo_platform_plugin/files/types.py | 4 + .../nemo_platform_plugin/guardrail/types.py | 2 +- .../src/nemo_platform_plugin/iam/endpoints.py | 11 +- .../src/nemo_platform_plugin/iam/types.py | 1 + .../inference_gateway/client.py | 50 +- .../inference_gateway/endpoints.py | 128 ++++- .../inference_gateway/types.py | 49 ++ .../src/nemo_platform_plugin/sdk.py | 15 +- .../virtual_models/endpoints.py | 25 +- .../workspaces/endpoints.py | 20 +- .../nemo_platform_plugin/workspaces/types.py | 6 + .../tests/client/test_adapter.py | 61 +- .../tests/client/test_auth_per_attempt.py | 241 ++++++++ .../tests/client/test_client_options.py | 33 ++ .../tests/files/test_endpoints.py | 33 ++ .../tests/guardrail/test_endpoints.py | 183 ++++++ .../tests/iam/test_client.py | 6 +- .../tests/iam/test_endpoints.py | 22 +- .../tests/inference_gateway/test_endpoints.py | 222 +++++++ .../tests/virtual_models/test_endpoints.py | 30 + .../tests/workspaces/test_client.py | 174 ++++++ .../tests/workspaces/test_endpoints.py | 169 ++++++ 30 files changed, 2729 insertions(+), 451 deletions(-) create mode 100644 packages/filesets/src/filesets/transfer.py create mode 100644 packages/filesets/tests/test_transfer.py create mode 100644 packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/types.py create mode 100644 packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py create mode 100644 packages/nemo_platform_plugin/tests/guardrail/test_endpoints.py create mode 100644 packages/nemo_platform_plugin/tests/inference_gateway/test_endpoints.py create mode 100644 packages/nemo_platform_plugin/tests/workspaces/test_client.py create mode 100644 packages/nemo_platform_plugin/tests/workspaces/test_endpoints.py diff --git a/packages/filesets/src/filesets/__init__.py b/packages/filesets/src/filesets/__init__.py index d4fc4d8a0a..288a0e9b49 100644 --- a/packages/filesets/src/filesets/__init__.py +++ b/packages/filesets/src/filesets/__init__.py @@ -7,6 +7,9 @@ fsspec integration for NeMo Platform filesets via the sdk.files.fsspec property. Located at: nemo_platform/filesets/ (after vendoring) + +``ListFilesResponse`` lives in :mod:`filesets.transfer`; it is also re-exported +here for backwards compatibility. """ from .filesystem.callbacks import RichFileProgressCallback as RichFileProgressCallback @@ -17,4 +20,4 @@ from .filesystem.filesystem import build_fileset_ref as build_fileset_ref from .filesystem.filesystem import parse_fileset_path as parse_fileset_path from .filesystem.filesystem import parse_fileset_ref as parse_fileset_ref -from .resources import ListFilesResponse as ListFilesResponse +from .transfer import ListFilesResponse as ListFilesResponse diff --git a/packages/filesets/src/filesets/resources.py b/packages/filesets/src/filesets/resources.py index 23514a1885..4b9249e696 100644 --- a/packages/filesets/src/filesets/resources.py +++ b/packages/filesets/src/filesets/resources.py @@ -7,15 +7,11 @@ backed by the NemoClient typed HTTP client and fsspec filesystem access. """ -import uuid from collections.abc import AsyncIterator, Iterator -from dataclasses import dataclass from functools import cached_property -from pathlib import PurePath from typing import Any, Protocol, runtime_checkable -from fsspec.callbacks import DEFAULT_CALLBACK, Callback -from fsspec.core import has_magic +from fsspec.callbacks import Callback from nemo_platform.resources.files.files import ( AsyncFilesResource as GeneratedAsyncFilesResource, ) @@ -25,70 +21,17 @@ from nemo_platform.resources.files.filesets import AsyncFilesetsResource, FilesetsResource from nemo_platform.resources.files.otlp.otlp import AsyncOtlpResource, OtlpResource from nemo_platform_plugin.files.client import AsyncFilesClient, FilesClient -from nemo_platform_plugin.files.types import ( - CacheStatus, - CreateFilesetRequest, - FilesetFileOutput, - FilesetOutput, - ListFilesQueryParams, -) +from nemo_platform_plugin.files.types import CreateFilesetRequest, FilesetOutput +from filesets import transfer from filesets.filesystem.filesystem import ( AsyncFilesetFileSystem, FilesetFileSystem, build_fileset_ref, parse_fileset_path, ) - - -@dataclass -class ListFilesResponse: - """Response from listing files in a fileset. - - Attributes: - data: List of files in the fileset. - - Properties: - cache_status: Aggregate cache status of all files. - - "caching" if any file is actively being cached - - "not_cached" if any file is not cached (and none are caching) - - "cached" if all files are fully cached - - "not_cacheable" if all files cannot be cached - - None if no cache information is available - """ - - data: list[FilesetFileOutput] - - @property - def cache_status(self) -> CacheStatus | None: - """Get aggregate cache status of all files. - - Returns the most relevant status based on priority: - - "caching" if any file is actively being cached - - "not_cached" if any file is not cached (and none are caching) - - "cached" if all files are fully cached - - "not_cacheable" if all files cannot be cached - - None if no cache information is available - """ - if not self.data: - return None - - statuses = [f.cache_status for f in self.data if f.cache_status is not None] - if not statuses: - return None - - # Priority: caching > not_cached > cached > not_cacheable - if "caching" in statuses: - return CacheStatus.CACHING - if "not_cached" in statuses: - return CacheStatus.NOT_CACHED - if all(s == "cached" for s in statuses): - return CacheStatus.CACHED - if all(s == "not_cacheable" for s in statuses): - return CacheStatus.NOT_CACHEABLE - - # Mixed cached/not_cacheable - return cached since some files are cached - return CacheStatus.CACHED +from filesets.transfer import ListFilesResponse as ListFilesResponse +from filesets.transfer import generate_fileset_name as _generate_fileset_name @runtime_checkable @@ -109,37 +52,6 @@ async def read(self, size: int = -1) -> bytes: ... AsyncContent = bytes | str | AsyncReadable | AsyncIterator[bytes] -def _generate_fileset_name() -> str: - """Generate a unique fileset name using UUID.""" - return f"fileset-{uuid.uuid4().hex[:8]}" - - -def _matches_glob(filepath: str, pattern: str) -> bool: - """Match filepath against a glob pattern using pathlib. - - Simple patterns (no /) only match top-level files. - Path patterns (with /) match the full relative path from the right. - - Examples: - _matches_glob("train.json", "*.json") -> True - _matches_glob("subdir/nested.json", "*.json") -> False (nested file) - _matches_glob("subdir/nested.json", "subdir/*.json") -> True - _matches_glob("subdir/nested.json", "*/*.json") -> True - - Args: - filepath: The file path to check (relative path within fileset). - pattern: Glob pattern to match against. - - Returns: - True if the filepath matches the pattern. - """ - if "/" not in pattern: - # Simple pattern - only matches top-level files - return "/" not in filepath and PurePath(filepath).match(pattern) - # Path pattern - match from the right - return PurePath(filepath).match(pattern) - - class FilesResource: """FilesResource with high-level file operations. @@ -272,47 +184,16 @@ def download( ... callback=cb ... ) """ - # Handle list of paths - if isinstance(remote_path, list): - if not remote_path: - return - ws = workspace or self._client.workspace - if fileset is None: - raise ValueError("fileset must be provided when remote_path is a list.") - if ws is None: - raise ValueError("workspace must be provided when remote_path is a list.") - # Build list of (remote, local) path pairs preserving directory structure - rpaths = [build_fileset_ref(p, workspace=ws, fileset=fileset) for p in remote_path] - lpaths = [str(PurePath(local_path) / p) for p in remote_path] - self.fsspec.get(rpath=rpaths, lpath=lpaths, callback=callback or DEFAULT_CALLBACK) - return - - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, + transfer.download( + self._client, + remote_path=remote_path, + local_path=local_path, + fileset=fileset, + workspace=workspace, + callback=callback, + max_workers=max_workers, + filesystem=self.fsspec, ) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - # Handle glob patterns by expanding to list of files first - if has_magic(path): - matching_files = self.list(remote_path=path, fileset=fileset, workspace=ws) - if not matching_files.data: - return - # Build list of (remote, local) path pairs preserving directory structure - rpaths = [build_fileset_ref(f.path, workspace=ws, fileset=fileset) for f in matching_files.data] - lpaths = [str(PurePath(local_path) / f.path) for f in matching_files.data] - self.fsspec.get(rpath=rpaths, lpath=lpaths, callback=callback or DEFAULT_CALLBACK) - else: - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - self.fsspec.get( - rpath=fileset_ref, - lpath=local_path, - recursive=True, - callback=callback or DEFAULT_CALLBACK, - ) def upload( self, @@ -383,33 +264,18 @@ def upload( ... ) >>> print(f"Uploaded to: {fileset.name}") # e.g., "fileset-a1b2c3d4" """ - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, - ) - fileset = fileset or path_fileset - - if fileset is None: - if fileset_auto_create: - fileset = _generate_fileset_name() - else: - raise ValueError( - "Fileset must be specified either as a parameter or in the remote_path when fileset_auto_create is False." - ) - - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - if fileset_auto_create: - self._ensure_fileset_exists(ws, fileset) - - self.fsspec.put( - lpath=local_path, - rpath=fileset_ref, - recursive=True, - callback=callback or DEFAULT_CALLBACK, + return transfer.upload( + self._client, + local_path=local_path, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + callback=callback, + max_workers=max_workers, + fileset_auto_create=fileset_auto_create, + filesystem=self.fsspec, ) - return self._client.get_fileset(name=fileset, workspace=ws).data() - def upload_content( self, *, @@ -617,37 +483,13 @@ def list( >>> for f in response.data: ... print(f"{f.path}: {f.cache_status}") """ - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, - ) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - # For glob patterns, list all files then filter client-side - # For path prefixes, the API handles filtering server-side - api_path = None if has_magic(path) else (path or None) - - query_params: ListFilesQueryParams = {} - if api_path is not None: - query_params["path"] = api_path - if include_cache_status: - query_params["include_cache_status"] = True - - response = self._client.list_files( - workspace=ws, - name=fileset, - query_params=query_params or None, + return transfer.list_files( + self._client, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + include_cache_status=include_cache_status, ) - response = response.data() - files = list(response.data) - - # Apply glob filtering if needed - if has_magic(path): - files = [f for f in files if _matches_glob(f.path, path)] - return ListFilesResponse(data=files) def delete( self, @@ -676,17 +518,13 @@ def delete( # Delete using full path >>> sdk.files.delete(remote_path="my-fileset#data/old-file.txt") """ - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, + transfer.delete( + self._client, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + filesystem=self.fsspec, ) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - self.fsspec.rm(fileset_ref) class AsyncFilesResource: @@ -801,48 +639,16 @@ async def download( ... local_path="./downloads/" ... ) """ - # Handle list of paths - if isinstance(remote_path, list): - if not remote_path: - return - ws = workspace or self._client.workspace - if fileset is None: - raise ValueError("fileset must be provided when remote_path is a list.") - if ws is None: - raise ValueError("workspace must be provided when remote_path is a list.") - # Build list of (remote, local) path pairs preserving directory structure - rpaths = [build_fileset_ref(p, workspace=ws, fileset=fileset) for p in remote_path] - lpaths = [str(PurePath(local_path) / p) for p in remote_path] - await self.fsspec._get(rpaths, lpaths, batch_size=max_workers, callback=callback or DEFAULT_CALLBACK) - return - - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, + await transfer.async_download( + self._client, + remote_path=remote_path, + local_path=local_path, + fileset=fileset, + workspace=workspace, + callback=callback, + max_workers=max_workers, + filesystem=self.fsspec, ) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - # Handle glob patterns by expanding to list of files first - if has_magic(path): - matching_files = await self.list(remote_path=path, fileset=fileset, workspace=ws) - if not matching_files.data: - return - # Build list of (remote, local) path pairs preserving directory structure - rpaths = [build_fileset_ref(f.path, workspace=ws, fileset=fileset) for f in matching_files.data] - lpaths = [str(PurePath(local_path) / f.path) for f in matching_files.data] - await self.fsspec._get(rpaths, lpaths, batch_size=max_workers, callback=callback or DEFAULT_CALLBACK) - else: - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - await self.fsspec._get( - fileset_ref, - local_path, - recursive=True, - batch_size=max_workers, - callback=callback or DEFAULT_CALLBACK, - ) async def upload( self, @@ -907,30 +713,17 @@ async def upload( ... ) >>> print(f"Uploaded to: {fileset.name}") # e.g., "fileset-a1b2c3d4" """ - ws, path_fileset, path = parse_fileset_path( - remote_path, - workspace_fallback=workspace or self._client.workspace, + return await transfer.async_upload( + self._client, + local_path=local_path, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + callback=callback, + max_workers=max_workers, + fileset_auto_create=fileset_auto_create, + filesystem=self.fsspec, ) - fileset = fileset or path_fileset - - if fileset is None: - if fileset_auto_create: - fileset = _generate_fileset_name() - else: - raise ValueError( - "Fileset must be specified either as a parameter or in the remote_path when fileset_auto_create is False." - ) - - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - if fileset_auto_create: - await self._ensure_fileset_exists(ws, fileset) - - kwargs: dict = {"lpath": local_path, "rpath": fileset_ref, "recursive": True, "batch_size": max_workers} - if callback is not None: - kwargs["callback"] = callback - await self.fsspec._put(**kwargs) - - return (await self._client.get_fileset(name=fileset, workspace=ws)).data() async def upload_content( self, @@ -1135,34 +928,13 @@ async def list( >>> for f in response.data: ... print(f"{f.path}: {f.cache_status}") """ - ws, path_fileset, path = parse_fileset_path(remote_path, workspace_fallback=workspace or self._client.workspace) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - # For glob patterns, list all files then filter client-side - # For path prefixes, the API handles filtering server-side - api_path = None if has_magic(path) else (path or None) - - query_params: ListFilesQueryParams = {} - if api_path is not None: - query_params["path"] = api_path - if include_cache_status: - query_params["include_cache_status"] = True - - response = await self._client.list_files( - workspace=ws, - name=fileset, - query_params=query_params or None, + return await transfer.async_list_files( + self._client, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + include_cache_status=include_cache_status, ) - response = response.data() - files = list(response.data) - - # Apply glob filtering if needed - if has_magic(path): - files = [f for f in files if _matches_glob(f.path, path)] - return ListFilesResponse(data=files) async def delete( self, @@ -1191,11 +963,10 @@ async def delete( # Delete using full path >>> await sdk.files.delete(remote_path="my-fileset#data/old-file.txt") """ - ws, path_fileset, path = parse_fileset_path(remote_path, workspace_fallback=workspace or self._client.workspace) - fileset = fileset or path_fileset - - if fileset is None: - raise ValueError("Fileset must be specified either as a parameter or in the remote_path.") - - fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) - await self.fsspec._rm(fileset_ref) + await transfer.async_delete( + self._client, + remote_path=remote_path, + fileset=fileset, + workspace=workspace, + filesystem=self.fsspec, + ) diff --git a/packages/filesets/src/filesets/transfer.py b/packages/filesets/src/filesets/transfer.py new file mode 100644 index 0000000000..9d19ee0434 --- /dev/null +++ b/packages/filesets/src/filesets/transfer.py @@ -0,0 +1,502 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""High-level fileset transfers on the typed Files client. + +``upload``, ``download``, ``list_files`` and ``delete`` (and their async twins) +drive :class:`~filesets.filesystem.filesystem.FilesetFileSystem` from a +:class:`~nemo_platform_plugin.files.client.FilesClient`. They resolve the +``[workspace/]fileset#path`` reference forms, expand glob patterns, create +filesets on demand, and report progress through fsspec callbacks. The CLI and +the SDK ``FilesResource`` both build on these functions. +""" + +from __future__ import annotations + +import uuid +from dataclasses import dataclass +from pathlib import PurePath + +from fsspec.callbacks import DEFAULT_CALLBACK, Callback +from fsspec.core import has_magic +from nemo_platform_plugin.files.client import AsyncFilesClient, FilesClient +from nemo_platform_plugin.files.types import ( + CacheStatus, + CreateFilesetRequest, + FilesetFileOutput, + FilesetOutput, + ListFilesQueryParams, +) + +from filesets.filesystem.filesystem import ( + AsyncFilesetFileSystem, + FilesetFileSystem, + build_fileset_ref, + parse_fileset_path, +) + +FILESET_REQUIRED_MESSAGE = "Fileset must be specified either as a parameter or in the remote_path." +FILESET_REQUIRED_WITHOUT_AUTO_CREATE_MESSAGE = ( + "Fileset must be specified either as a parameter or in the remote_path when fileset_auto_create is False." +) + + +@dataclass +class ListFilesResponse: + """Response from listing files in a fileset. + + Attributes: + data: List of files in the fileset. + + Properties: + cache_status: Aggregate cache status of all files. + - "caching" if any file is actively being cached + - "not_cached" if any file is not cached (and none are caching) + - "cached" if all files are fully cached + - "not_cacheable" if all files cannot be cached + - None if no cache information is available + """ + + data: list[FilesetFileOutput] + + @property + def cache_status(self) -> CacheStatus | None: + """Get aggregate cache status of all files. + + Returns the most relevant status based on priority: + - "caching" if any file is actively being cached + - "not_cached" if any file is not cached (and none are caching) + - "cached" if all files are fully cached + - "not_cacheable" if all files cannot be cached + - None if no cache information is available + """ + if not self.data: + return None + + statuses = [f.cache_status for f in self.data if f.cache_status is not None] + if not statuses: + return None + + # Priority: caching > not_cached > cached > not_cacheable + if "caching" in statuses: + return CacheStatus.CACHING + if "not_cached" in statuses: + return CacheStatus.NOT_CACHED + if all(s == "cached" for s in statuses): + return CacheStatus.CACHED + if all(s == "not_cacheable" for s in statuses): + return CacheStatus.NOT_CACHEABLE + + # Mixed cached/not_cacheable - return cached since some files are cached + return CacheStatus.CACHED + + +def generate_fileset_name() -> str: + """Generate a unique fileset name using UUID.""" + return f"fileset-{uuid.uuid4().hex[:8]}" + + +def matches_glob(filepath: str, pattern: str) -> bool: + """Match filepath against a glob pattern using pathlib. + + Simple patterns (no /) only match top-level files. + Path patterns (with /) match the full relative path from the right. + + Examples: + matches_glob("train.json", "*.json") -> True + matches_glob("subdir/nested.json", "*.json") -> False (nested file) + matches_glob("subdir/nested.json", "subdir/*.json") -> True + matches_glob("subdir/nested.json", "*/*.json") -> True + + Args: + filepath: The file path to check (relative path within fileset). + pattern: Glob pattern to match against. + + Returns: + True if the filepath matches the pattern. + """ + if "/" not in pattern: + # Simple pattern - only matches top-level files + return "/" not in filepath and PurePath(filepath).match(pattern) + # Path pattern - match from the right + return PurePath(filepath).match(pattern) + + +def _resolve_target( + remote_path: str, + *, + fileset: str | None, + workspace: str | None, + client_workspace: str | None, +) -> tuple[str, str | None, str]: + """Return ``(workspace, fileset, path)`` for a remote path that may embed a fileset ref.""" + ws, path_fileset, path = parse_fileset_path(remote_path, workspace_fallback=workspace or client_workspace) + return ws, fileset or path_fileset, path + + +def _resolve_upload_fileset(fileset: str | None, *, fileset_auto_create: bool) -> str: + if fileset is not None: + return fileset + if fileset_auto_create: + return generate_fileset_name() + raise ValueError(FILESET_REQUIRED_WITHOUT_AUTO_CREATE_MESSAGE) + + +def list_query_params(path: str, *, include_cache_status: bool = False) -> ListFilesQueryParams | None: + """Query params ``list_files`` sends for *path*: a prefix goes to the server, a glob is filtered client-side.""" + query_params: ListFilesQueryParams = {} + if not has_magic(path) and path: + query_params["path"] = path + if include_cache_status: + query_params["include_cache_status"] = True + return query_params or None + + +def _filter_listed(files: list[FilesetFileOutput], path: str) -> ListFilesResponse: + if has_magic(path): + files = [f for f in files if matches_glob(f.path, path)] + return ListFilesResponse(data=files) + + +def _pairs_for(paths: list[str], *, workspace: str, fileset: str, local_path: str) -> tuple[list[str], list[str]]: + """Build parallel remote/local path lists that preserve directory structure.""" + rpaths = [build_fileset_ref(p, workspace=workspace, fileset=fileset) for p in paths] + lpaths = [str(PurePath(local_path) / p) for p in paths] + return rpaths, lpaths + + +def _async_transfer_kwargs(callback: Callback | None, max_workers: int | None) -> dict: + kwargs: dict = {"batch_size": max_workers} + if callback is not None: + kwargs["callback"] = callback + return kwargs + + +# --------------------------------------------------------------------------- +# Sync API +# --------------------------------------------------------------------------- + + +def _fs(client: FilesClient, filesystem: FilesetFileSystem | None) -> FilesetFileSystem: + return filesystem if filesystem is not None else FilesetFileSystem(client=client) + + +def list_files( + client: FilesClient, + *, + remote_path: str = "", + fileset: str | None = None, + workspace: str | None = None, + include_cache_status: bool = False, +) -> ListFilesResponse: + """List all files in a fileset path (recursive), with optional glob pattern support. + + Args: + client: Typed Files client. + remote_path: Path within the fileset to list. Can be a full path + (e.g., "workspace/fileset#data/" or "fileset#data/") if fileset is not provided, + or a relative path (e.g., "data/") if fileset is provided. + Supports glob patterns (*, ?, []) for filtering files. + Defaults to "" (root of fileset). + fileset: Fileset name. If not provided, inferred from remote_path. + workspace: Workspace name. If not provided, inferred from remote_path + or uses the client's default workspace. + include_cache_status: Check and return cache status for each file. + When False (default), external storage files return None for cache_status. + + Returns: + ListFilesResponse with data (list of FilesetFileOutput) and cache_status property. + """ + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + response = client.list_files( + workspace=ws, + name=fileset, + query_params=list_query_params(path, include_cache_status=include_cache_status), + ).data() + return _filter_listed(list(response.data), path) + + +def download( + client: FilesClient, + *, + remote_path: str | list[str] = "", + local_path: str, + fileset: str | None = None, + workspace: str | None = None, + callback: Callback | None = None, + max_workers: int | None = None, + filesystem: FilesetFileSystem | None = None, +) -> None: + """Download files from a fileset to a local path. + + Args: + client: Typed Files client. + remote_path: Path(s) within the fileset to download. Can be: + - A single path (str): Full path (e.g., "workspace/fileset#data/"), + relative path (e.g., "data/"), or glob pattern (e.g., "*.json"). + - A list of paths (list[str]): Multiple specific file paths to download. + When using a list, fileset and workspace must be provided explicitly. + Defaults to "" (root of fileset). + local_path: Local destination path (directory). + fileset: Fileset name. If not provided, inferred from remote_path (str only). + workspace: Workspace name. If not provided, inferred from remote_path + or uses the client's default workspace. + callback: Optional progress callback (e.g., RichProgressCallback). + max_workers: Maximum number of concurrent file transfers. + filesystem: Filesystem to transfer through; defaults to one built on *client*. + """ + fs = _fs(client, filesystem) + callback = callback or DEFAULT_CALLBACK + + if isinstance(remote_path, list): + if not remote_path: + return + ws = workspace or client.workspace + if fileset is None: + raise ValueError("fileset must be provided when remote_path is a list.") + if ws is None: + raise ValueError("workspace must be provided when remote_path is a list.") + rpaths, lpaths = _pairs_for(remote_path, workspace=ws, fileset=fileset, local_path=local_path) + fs.get(rpath=rpaths, lpath=lpaths, callback=callback) + return + + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + if has_magic(path): + matching = list_files(client, remote_path=path, fileset=fileset, workspace=ws) + if not matching.data: + return + rpaths, lpaths = _pairs_for( + [f.path for f in matching.data], workspace=ws, fileset=fileset, local_path=local_path + ) + fs.get(rpath=rpaths, lpath=lpaths, callback=callback) + return + + fs.get( + rpath=build_fileset_ref(path, workspace=ws, fileset=fileset), + lpath=local_path, + recursive=True, + callback=callback, + ) + + +def upload( + client: FilesClient, + *, + local_path: str, + remote_path: str = "", + fileset: str | None = None, + workspace: str | None = None, + callback: Callback | None = None, + max_workers: int | None = None, + fileset_auto_create: bool = False, + filesystem: FilesetFileSystem | None = None, +) -> FilesetOutput: + """Upload files from a local path to a fileset. + + Args: + client: Typed Files client. + local_path: Local source path (file or directory). A trailing slash on a + directory uploads its contents rather than the directory itself. + remote_path: Path within the fileset to upload to. Can be a full path + (e.g., "workspace/fileset#data/" or "fileset#data/") if fileset is not provided, + or a relative path (e.g., "data/") if fileset is provided. + Defaults to "" (root of fileset). + fileset: Fileset name. If not provided, inferred from remote_path. + workspace: Workspace name. If not provided, inferred from remote_path + or uses the client's default workspace. + callback: Optional progress callback (e.g., RichProgressCallback). + max_workers: Maximum number of concurrent file transfers. + fileset_auto_create: If True, create the fileset if it doesn't exist. + When no fileset is specified (neither as param nor in remote_path), + a unique name is generated (e.g., "fileset-a1b2c3d4"). + filesystem: Filesystem to transfer through; defaults to one built on *client*. + + Returns: + FilesetOutput: The fileset that was uploaded to. Check ``fileset.name`` to see + the generated name when using fileset_auto_create without specifying + a fileset. + """ + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + fileset = _resolve_upload_fileset(fileset, fileset_auto_create=fileset_auto_create) + + fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) + if fileset_auto_create: + client.create_fileset(workspace=ws, body=CreateFilesetRequest(name=fileset), exist_ok=True) + + _fs(client, filesystem).put( + lpath=local_path, rpath=fileset_ref, recursive=True, callback=callback or DEFAULT_CALLBACK + ) + + return client.get_fileset(name=fileset, workspace=ws).data() + + +def delete( + client: FilesClient, + *, + remote_path: str, + fileset: str | None = None, + workspace: str | None = None, + filesystem: FilesetFileSystem | None = None, +) -> None: + """Delete a file from a fileset. + + Args: + client: Typed Files client. + remote_path: Path of the file to delete. Can be a full path + (e.g., "workspace/fileset#data/file.txt") if fileset is not provided, + or a relative path (e.g., "data/file.txt") if fileset is provided. + fileset: Fileset name. If not provided, inferred from remote_path. + workspace: Workspace name. If not provided, inferred from remote_path + or uses the client's default workspace. + filesystem: Filesystem to delete through; defaults to one built on *client*. + """ + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + _fs(client, filesystem).rm(build_fileset_ref(path, workspace=ws, fileset=fileset)) + + +# --------------------------------------------------------------------------- +# Async API +# --------------------------------------------------------------------------- + + +def _async_fs(client: AsyncFilesClient, filesystem: AsyncFilesetFileSystem | None) -> AsyncFilesetFileSystem: + return filesystem if filesystem is not None else AsyncFilesetFileSystem(client=client) + + +async def async_list_files( + client: AsyncFilesClient, + *, + remote_path: str = "", + fileset: str | None = None, + workspace: str | None = None, + include_cache_status: bool = False, +) -> ListFilesResponse: + """Async twin of :func:`list_files`.""" + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + response = await client.list_files( + workspace=ws, + name=fileset, + query_params=list_query_params(path, include_cache_status=include_cache_status), + ) + return _filter_listed(list(response.data().data), path) + + +async def async_download( + client: AsyncFilesClient, + *, + remote_path: str | list[str] = "", + local_path: str, + fileset: str | None = None, + workspace: str | None = None, + callback: Callback | None = None, + max_workers: int | None = None, + filesystem: AsyncFilesetFileSystem | None = None, +) -> None: + """Async twin of :func:`download`.""" + fs = _async_fs(client, filesystem) + cb = callback or DEFAULT_CALLBACK + + if isinstance(remote_path, list): + if not remote_path: + return + ws = workspace or client.workspace + if fileset is None: + raise ValueError("fileset must be provided when remote_path is a list.") + if ws is None: + raise ValueError("workspace must be provided when remote_path is a list.") + rpaths, lpaths = _pairs_for(remote_path, workspace=ws, fileset=fileset, local_path=local_path) + await fs._get(rpaths, lpaths, batch_size=max_workers, callback=cb) + return + + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + if has_magic(path): + matching = await async_list_files(client, remote_path=path, fileset=fileset, workspace=ws) + if not matching.data: + return + rpaths, lpaths = _pairs_for( + [f.path for f in matching.data], workspace=ws, fileset=fileset, local_path=local_path + ) + await fs._get(rpaths, lpaths, batch_size=max_workers, callback=cb) + return + + await fs._get( + build_fileset_ref(path, workspace=ws, fileset=fileset), + local_path, + recursive=True, + batch_size=max_workers, + callback=cb, + ) + + +async def async_upload( + client: AsyncFilesClient, + *, + local_path: str, + remote_path: str = "", + fileset: str | None = None, + workspace: str | None = None, + callback: Callback | None = None, + max_workers: int | None = None, + fileset_auto_create: bool = False, + filesystem: AsyncFilesetFileSystem | None = None, +) -> FilesetOutput: + """Async twin of :func:`upload`.""" + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + fileset = _resolve_upload_fileset(fileset, fileset_auto_create=fileset_auto_create) + + fileset_ref = build_fileset_ref(path, workspace=ws, fileset=fileset) + if fileset_auto_create: + await client.create_fileset(workspace=ws, body=CreateFilesetRequest(name=fileset), exist_ok=True) + + await _async_fs(client, filesystem)._put( + lpath=local_path, rpath=fileset_ref, recursive=True, **_async_transfer_kwargs(callback, max_workers) + ) + + return (await client.get_fileset(name=fileset, workspace=ws)).data() + + +async def async_delete( + client: AsyncFilesClient, + *, + remote_path: str, + fileset: str | None = None, + workspace: str | None = None, + filesystem: AsyncFilesetFileSystem | None = None, +) -> None: + """Async twin of :func:`delete`.""" + ws, fileset, path = _resolve_target( + remote_path, fileset=fileset, workspace=workspace, client_workspace=client.workspace + ) + if fileset is None: + raise ValueError(FILESET_REQUIRED_MESSAGE) + + await _async_fs(client, filesystem)._rm(build_fileset_ref(path, workspace=ws, fileset=fileset)) diff --git a/packages/filesets/tests/test_transfer.py b/packages/filesets/tests/test_transfer.py new file mode 100644 index 0000000000..d7726d0f2a --- /dev/null +++ b/packages/filesets/tests/test_transfer.py @@ -0,0 +1,544 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for the Stainless-free transfer helpers in ``filesets.transfer``. + +A small in-memory fileset server answers both the sync client and the async +client the filesystem builds from it, so the tests pin the request sequence, +paths, query params, and bodies that each helper produces. +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from pathlib import Path + +import httpx +import pytest +from filesets import transfer +from filesets.filesystem.filesystem import FilesetFileSystem +from filesets.transfer import ListFilesResponse +from nemo_platform_plugin.client.errors import NotFoundError +from nemo_platform_plugin.files.client import AsyncFilesClient, FilesClient +from nemo_platform_plugin.files.types import CacheStatus, FilesetFileOutput +from starlette.requests import Request +from starlette.responses import Response +from starlette.types import Receive, Scope, Send + +BASE = "http://test" +WORKSPACE = "default" + + +def _fileset_json(workspace: str, name: str) -> dict: + return { + "id": f"id-{name}", + "name": name, + "workspace": workspace, + "description": "", + "purpose": "generic", + "storage": {"type": "local", "path": f"/data/{name}"}, + "metadata": {}, + "custom_fields": {}, + "project": "", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", + } + + +@dataclass +class FakeFilesServer: + """In-memory fileset store keyed by ``(workspace, fileset)``; records every request.""" + + filesets: dict[tuple[str, str], dict[str, bytes]] = field(default_factory=dict) + requests: list[httpx.Request] = field(default_factory=list) + + def calls(self) -> list[tuple[str, str]]: + return [(request.method, request.url.path) for request in self.requests] + + def __call__(self, request: httpx.Request) -> httpx.Response: + request.read() + self.requests.append(request) + parts = request.url.path.lstrip("/").split("/") + # apis/files/v2/workspaces/{ws}/filesets[/{name}[/files | /-/{path}]] + workspace = parts[4] + if len(parts) == 6: + if request.method == "POST": + name = json.loads(request.content)["name"] + if (workspace, name) in self.filesets: + return httpx.Response(409, json={"detail": "exists"}) + self.filesets[(workspace, name)] = {} + return httpx.Response(201, json=_fileset_json(workspace, name)) + return httpx.Response(405) + name = parts[6] + files = self.filesets.get((workspace, name)) + if files is None: + return httpx.Response(404, json={"detail": f"Fileset '{name}' not found"}) + if len(parts) == 7: + return httpx.Response(200, json=_fileset_json(workspace, name)) + if parts[7] == "files": + prefix = request.url.params.get("path", "") + data = [ + { + "file_ref": f"{workspace}/{name}#{path}", + "file_url": f"/apis/files/v2/workspaces/{workspace}/filesets/{name}/-/{path}", + "path": path, + "size": len(content), + } + for path, content in sorted(files.items()) + if path.startswith(prefix) + ] + return httpx.Response(200, json={"data": data}) + path = "/".join(parts[8:]) + file_json = { + "file_ref": f"{workspace}/{name}#{path}", + "file_url": request.url.path, + "path": path, + "size": len(files.get(path, b"")), + } + if request.method == "PUT": + files[path] = request.content + return httpx.Response(200, json={**file_json, "size": len(request.content)}) + if path not in files: + return httpx.Response(404, json={"detail": "File not found"}) + if request.method == "GET": + return httpx.Response(200, content=files[path], headers={"content-length": str(len(files[path]))}) + if request.method == "DELETE": + del files[path] + return httpx.Response(200, json=file_json) + return httpx.Response(405) + + async def asgi(self, scope: Scope, receive: Receive, send: Send) -> None: + incoming = Request(scope, receive) + body = await incoming.body() + response = self(httpx.Request(incoming.method, str(incoming.url), headers=incoming.headers.raw, content=body)) + await Response(content=response.content, status_code=response.status_code, headers=dict(response.headers))( + scope, receive, send + ) + + +class _HttpClient(httpx.Client): + def __init__(self, server: FakeFilesServer) -> None: + super().__init__(transport=httpx.MockTransport(server), base_url=BASE) + self._server = server + + @property + def asgi_app(self): + return self._server.asgi + + +@pytest.fixture +def server() -> FakeFilesServer: + return FakeFilesServer() + + +@pytest.fixture +def client(server: FakeFilesServer) -> FilesClient: + return FilesClient(base_url=BASE, workspace=WORKSPACE, http_client=_HttpClient(server)) + + +@pytest.fixture +def async_client(server: FakeFilesServer) -> AsyncFilesClient: + return AsyncFilesClient( + base_url=BASE, workspace=WORKSPACE, http_client=httpx.AsyncClient(transport=httpx.ASGITransport(server.asgi)) + ) + + +def _put(request: httpx.Request) -> tuple[str, str]: + return request.method, request.url.path + + +# --------------------------------------------------------------------------- +# helpers +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + ("filepath", "pattern", "expected"), + [ + ("train.json", "*.json", True), + ("subdir/nested.json", "*.json", False), + ("subdir/nested.json", "subdir/*.json", True), + ("subdir/nested.json", "*/*.json", True), + ("a/b/c.txt", "b/*.txt", True), + ("a/b/c.txt", "*.md", False), + ], +) +def test_matches_glob(filepath: str, pattern: str, expected: bool) -> None: + assert transfer.matches_glob(filepath, pattern) is expected + + +def test_generate_fileset_name_is_unique_and_prefixed() -> None: + names = {transfer.generate_fileset_name() for _ in range(5)} + assert len(names) == 5 + assert all(name.startswith("fileset-") and len(name) == len("fileset-") + 8 for name in names) + + +def _file(path: str, cache_status: CacheStatus | None) -> FilesetFileOutput: + return FilesetFileOutput(file_ref=f"ws/fs#{path}", file_url="/x", path=path, size=1, cache_status=cache_status) + + +@pytest.mark.parametrize( + ("statuses", "expected"), + [ + ([], None), + ([None, None], None), + ([CacheStatus.CACHED, CacheStatus.CACHING], CacheStatus.CACHING), + ([CacheStatus.CACHED, CacheStatus.NOT_CACHED], CacheStatus.NOT_CACHED), + ([CacheStatus.CACHED, CacheStatus.CACHED], CacheStatus.CACHED), + ([CacheStatus.NOT_CACHEABLE, CacheStatus.NOT_CACHEABLE], CacheStatus.NOT_CACHEABLE), + ([CacheStatus.CACHED, CacheStatus.NOT_CACHEABLE], CacheStatus.CACHED), + ], +) +def test_list_files_response_cache_status(statuses: list[CacheStatus | None], expected: CacheStatus | None) -> None: + response = ListFilesResponse(data=[_file(f"f{i}", status) for i, status in enumerate(statuses)]) + assert response.cache_status == expected + + +# --------------------------------------------------------------------------- +# list_files +# --------------------------------------------------------------------------- + + +def test_list_files_root(server: FakeFilesServer, client: FilesClient) -> None: + server.filesets[(WORKSPACE, "fs")] = {"a.txt": b"a", "d/b.txt": b"bb"} + + response = transfer.list_files(client, fileset="fs") + + assert [(f.path, f.size) for f in response.data] == [("a.txt", 1), ("d/b.txt", 2)] + assert server.calls() == [("GET", "/apis/files/v2/workspaces/default/filesets/fs/files")] + assert dict(server.requests[0].url.params) == {} + + +def test_list_files_prefix_is_sent_as_path_param(server: FakeFilesServer, client: FilesClient) -> None: + server.filesets[(WORKSPACE, "fs")] = {"a.txt": b"a", "d/b.txt": b"bb"} + + response = transfer.list_files(client, fileset="fs", remote_path="d/", include_cache_status=True) + + assert [f.path for f in response.data] == ["d/b.txt"] + assert dict(server.requests[0].url.params) == {"path": "d/", "include_cache_status": "true"} + + +def test_list_files_glob_filters_client_side(server: FakeFilesServer, client: FilesClient) -> None: + server.filesets[(WORKSPACE, "fs")] = {"a.json": b"a", "b.txt": b"b", "d/c.json": b"c"} + + response = transfer.list_files(client, fileset="fs", remote_path="*.json") + + assert [f.path for f in response.data] == ["a.json"] + assert dict(server.requests[0].url.params) == {} + + +def test_list_files_parses_fileset_ref_in_remote_path(server: FakeFilesServer, client: FilesClient) -> None: + server.filesets[("other", "fs")] = {"d/x.txt": b"x"} + + response = transfer.list_files(client, remote_path="other/fs#d/") + + assert [f.path for f in response.data] == ["d/x.txt"] + assert server.calls() == [("GET", "/apis/files/v2/workspaces/other/filesets/fs/files")] + + +def test_list_files_requires_fileset(client: FilesClient) -> None: + with pytest.raises(ValueError, match="Fileset must be specified"): + transfer.list_files(client, remote_path="d/") + + +def test_list_files_missing_fileset_raises_not_found(client: FilesClient) -> None: + with pytest.raises(NotFoundError): + transfer.list_files(client, fileset="nope") + + +# --------------------------------------------------------------------------- +# upload +# --------------------------------------------------------------------------- + + +def test_upload_single_file(server: FakeFilesServer, client: FilesClient, tmp_path: Path) -> None: + server.filesets[(WORKSPACE, "fs")] = {} + local = tmp_path / "a.txt" + local.write_bytes(b"hello") + + result = transfer.upload(client, local_path=str(local), fileset="fs", remote_path="data/") + + assert result.name == "fs" + assert server.filesets[(WORKSPACE, "fs")] == {"data/a.txt": b"hello"} + put = next(r for r in server.requests if r.method == "PUT") + assert put.url.path == "/apis/files/v2/workspaces/default/filesets/fs/-/data/a.txt" + assert put.headers["content-length"] == "5" + assert server.calls()[-1] == ("GET", "/apis/files/v2/workspaces/default/filesets/fs") + + +def test_upload_directory_keeps_name_without_trailing_slash( + server: FakeFilesServer, client: FilesClient, tmp_path: Path +) -> None: + server.filesets[(WORKSPACE, "fs")] = {} + src = tmp_path / "src" + (src / "n").mkdir(parents=True) + (src / "a.txt").write_bytes(b"a") + (src / "n" / "b.txt").write_bytes(b"b") + + transfer.upload(client, local_path=str(src), fileset="fs") + + assert server.filesets[(WORKSPACE, "fs")] == {"src/a.txt": b"a", "src/n/b.txt": b"b"} + + +def test_upload_directory_contents_with_trailing_slash( + server: FakeFilesServer, client: FilesClient, tmp_path: Path +) -> None: + server.filesets[(WORKSPACE, "fs")] = {} + src = tmp_path / "src" + (src / "n").mkdir(parents=True) + (src / "a.txt").write_bytes(b"a") + (src / "n" / "b.txt").write_bytes(b"b") + + transfer.upload(client, local_path=f"{src}/", fileset="fs", remote_path="up/") + + assert server.filesets[(WORKSPACE, "fs")] == {"up/a.txt": b"a", "up/n/b.txt": b"b"} + + +def test_upload_fileset_from_remote_path_ref(server: FakeFilesServer, client: FilesClient, tmp_path: Path) -> None: + server.filesets[("other", "fs")] = {} + local = tmp_path / "a.txt" + local.write_bytes(b"x") + + result = transfer.upload(client, local_path=str(local), remote_path="other/fs#dir/") + + assert result.workspace == "other" + assert server.filesets[("other", "fs")] == {"dir/a.txt": b"x"} + + +def test_upload_requires_fileset_without_auto_create(client: FilesClient, tmp_path: Path) -> None: + with pytest.raises(ValueError, match="fileset_auto_create is False"): + transfer.upload(client, local_path=str(tmp_path)) + + +def test_upload_auto_create_named_fileset_is_idempotent( + server: FakeFilesServer, client: FilesClient, tmp_path: Path +) -> None: + local = tmp_path / "a.txt" + local.write_bytes(b"x") + + first = transfer.upload(client, local_path=str(local), fileset="new", fileset_auto_create=True) + second = transfer.upload(client, local_path=str(local), fileset="new", fileset_auto_create=True) + + assert first.name == second.name == "new" + posts = [r for r in server.requests if r.method == "POST"] + assert [json.loads(r.content) for r in posts] == [{"name": "new"}, {"name": "new"}] + # The second create 409s and is resolved by re-fetching the fileset (exist_ok). + assert server.filesets[(WORKSPACE, "new")] == {"a.txt": b"x"} + + +def test_upload_auto_create_generates_name(server: FakeFilesServer, client: FilesClient, tmp_path: Path) -> None: + local = tmp_path / "a.txt" + local.write_bytes(b"x") + + result = transfer.upload(client, local_path=str(local), fileset_auto_create=True) + + assert result.name.startswith("fileset-") + assert server.filesets[(WORKSPACE, result.name)] == {"a.txt": b"x"} + assert server.calls()[0] == ("POST", "/apis/files/v2/workspaces/default/filesets") + + +def test_upload_uses_supplied_filesystem(server: FakeFilesServer, client: FilesClient, tmp_path: Path) -> None: + server.filesets[(WORKSPACE, "fs")] = {} + local = tmp_path / "a.txt" + local.write_bytes(b"x") + fs = FilesetFileSystem(client=client) + + transfer.upload(client, local_path=str(local), fileset="fs", filesystem=fs) + + assert server.filesets[(WORKSPACE, "fs")] == {"a.txt": b"x"} + + +# --------------------------------------------------------------------------- +# download +# --------------------------------------------------------------------------- + + +@pytest.fixture +def populated(server: FakeFilesServer) -> dict[str, bytes]: + files = {"a/file1.txt": b"content1", "a/b/file2.txt": b"content2", "a/b/file3.txt": b"content3", "r.json": b"{}"} + server.filesets[(WORKSPACE, "fs")] = dict(files) + return files + + +def test_download_single_file_into_existing_directory( + populated: dict[str, bytes], client: FilesClient, tmp_path: Path +) -> None: + out = tmp_path / "out" + out.mkdir() + + transfer.download(client, fileset="fs", remote_path="a/b/file2.txt", local_path=str(out)) + + assert (out / "file2.txt").read_bytes() == b"content2" + + +def test_download_single_file_to_exact_path(populated: dict[str, bytes], client: FilesClient, tmp_path: Path) -> None: + dest = tmp_path / "renamed.txt" + + transfer.download(client, fileset="fs", remote_path="a/b/file2.txt", local_path=str(dest)) + + assert dest.read_bytes() == b"content2" + + +def test_download_directory_contents_with_trailing_slash( + populated: dict[str, bytes], client: FilesClient, tmp_path: Path +) -> None: + out = tmp_path / "out" + + transfer.download(client, fileset="fs", remote_path="a/", local_path=f"{out}/") + + assert (out / "file1.txt").read_bytes() == b"content1" + assert (out / "b" / "file2.txt").read_bytes() == b"content2" + assert (out / "b" / "file3.txt").read_bytes() == b"content3" + + +def test_download_directory_keeps_name_without_trailing_slash( + populated: dict[str, bytes], client: FilesClient, tmp_path: Path +) -> None: + out = tmp_path / "out" + + transfer.download(client, fileset="fs", remote_path="a/b", local_path=str(out)) + + assert (out / "b" / "file2.txt").read_bytes() == b"content2" + assert (out / "b" / "file3.txt").read_bytes() == b"content3" + assert not (out / "file1.txt").exists() + + +def test_download_fileset_root_copies_contents( + populated: dict[str, bytes], client: FilesClient, tmp_path: Path +) -> None: + out = tmp_path / "out" + + transfer.download(client, fileset="fs", local_path=str(out)) + + assert sorted(p.relative_to(out).as_posix() for p in out.rglob("*") if p.is_file()) == sorted(populated) + + +def test_download_glob_preserves_relative_paths( + populated: dict[str, bytes], server: FakeFilesServer, client: FilesClient, tmp_path: Path +) -> None: + out = tmp_path / "out" + + transfer.download(client, fileset="fs", remote_path="a/b/*.txt", local_path=str(out)) + + assert (out / "a" / "b" / "file2.txt").read_bytes() == b"content2" + assert (out / "a" / "b" / "file3.txt").read_bytes() == b"content3" + assert not (out / "a" / "file1.txt").exists() + downloads = sorted(r.url.path for r in server.requests if r.method == "GET" and "/-/" in r.url.path) + assert downloads == [ + "/apis/files/v2/workspaces/default/filesets/fs/-/a/b/file2.txt", + "/apis/files/v2/workspaces/default/filesets/fs/-/a/b/file3.txt", + ] + + +def test_download_glob_without_matches_is_noop( + populated: dict[str, bytes], server: FakeFilesServer, client: FilesClient, tmp_path: Path +) -> None: + transfer.download(client, fileset="fs", remote_path="*.parquet", local_path=str(tmp_path)) + + assert server.calls() == [("GET", "/apis/files/v2/workspaces/default/filesets/fs/files")] + + +def test_download_list_of_paths(populated: dict[str, bytes], client: FilesClient, tmp_path: Path) -> None: + out = tmp_path / "out" + + transfer.download(client, fileset="fs", remote_path=["a/file1.txt", "r.json"], local_path=str(out)) + + assert (out / "a" / "file1.txt").read_bytes() == b"content1" + assert (out / "r.json").read_bytes() == b"{}" + + +def test_download_list_requires_fileset_and_workspace(server: FakeFilesServer, tmp_path: Path) -> None: + no_workspace = FilesClient(base_url=BASE, http_client=_HttpClient(server)) + + with pytest.raises(ValueError, match="fileset must be provided"): + transfer.download(no_workspace, remote_path=["a"], local_path=str(tmp_path)) + with pytest.raises(ValueError, match="workspace must be provided"): + transfer.download(no_workspace, fileset="fs", remote_path=["a"], local_path=str(tmp_path)) + + +def test_download_empty_list_is_noop(server: FakeFilesServer, client: FilesClient, tmp_path: Path) -> None: + transfer.download(client, fileset="fs", remote_path=[], local_path=str(tmp_path)) + + assert server.requests == [] + + +def test_download_requires_fileset(client: FilesClient, tmp_path: Path) -> None: + with pytest.raises(ValueError, match="Fileset must be specified"): + transfer.download(client, remote_path="a/", local_path=str(tmp_path)) + + +def test_download_missing_fileset_raises_not_found(client: FilesClient, tmp_path: Path) -> None: + with pytest.raises(NotFoundError): + transfer.download(client, fileset="nope", local_path=str(tmp_path)) + + +# --------------------------------------------------------------------------- +# delete +# --------------------------------------------------------------------------- + + +def test_delete_file(populated: dict[str, bytes], server: FakeFilesServer, client: FilesClient) -> None: + transfer.delete(client, fileset="fs", remote_path="a/b/file2.txt") + + assert server.calls() == [("DELETE", "/apis/files/v2/workspaces/default/filesets/fs/-/a/b/file2.txt")] + assert "a/b/file2.txt" not in server.filesets[(WORKSPACE, "fs")] + + +def test_delete_with_fileset_ref(populated: dict[str, bytes], server: FakeFilesServer, client: FilesClient) -> None: + transfer.delete(client, remote_path="default/fs#r.json") + + assert server.calls() == [("DELETE", "/apis/files/v2/workspaces/default/filesets/fs/-/r.json")] + + +def test_delete_requires_fileset(client: FilesClient) -> None: + with pytest.raises(ValueError, match="Fileset must be specified"): + transfer.delete(client, remote_path="a.txt") + + +def test_delete_missing_file_raises_not_found(populated: dict[str, bytes], client: FilesClient) -> None: + with pytest.raises(NotFoundError): + transfer.delete(client, fileset="fs", remote_path="nope.txt") + + +# --------------------------------------------------------------------------- +# async twins +# --------------------------------------------------------------------------- + + +async def test_async_roundtrip(server: FakeFilesServer, async_client: AsyncFilesClient, tmp_path: Path) -> None: + src = tmp_path / "src" + src.mkdir() + (src / "a.txt").write_bytes(b"a") + (src / "b.json").write_bytes(b"{}") + + result = await transfer.async_upload( + async_client, local_path=f"{src}/", fileset="fs", remote_path="up/", fileset_auto_create=True + ) + assert result.name == "fs" + assert server.filesets[(WORKSPACE, "fs")] == {"up/a.txt": b"a", "up/b.json": b"{}"} + + listed = await transfer.async_list_files(async_client, fileset="fs", remote_path="up/*.json") + assert [f.path for f in listed.data] == ["up/b.json"] + + out = tmp_path / "out" + await transfer.async_download(async_client, fileset="fs", remote_path="up/", local_path=f"{out}/") + assert (out / "a.txt").read_bytes() == b"a" + assert (out / "b.json").read_bytes() == b"{}" + + await transfer.async_download(async_client, fileset="fs", remote_path=["up/a.txt"], local_path=str(out / "l")) + assert (out / "l" / "up" / "a.txt").read_bytes() == b"a" + + await transfer.async_delete(async_client, fileset="fs", remote_path="up/a.txt") + assert server.filesets[(WORKSPACE, "fs")] == {"up/b.json": b"{}"} + + +async def test_async_requires_fileset(async_client: AsyncFilesClient, tmp_path: Path) -> None: + with pytest.raises(ValueError, match="Fileset must be specified"): + await transfer.async_list_files(async_client, remote_path="x/") + with pytest.raises(ValueError, match="Fileset must be specified"): + await transfer.async_download(async_client, remote_path="x/", local_path=str(tmp_path)) + with pytest.raises(ValueError, match="Fileset must be specified"): + await transfer.async_delete(async_client, remote_path="x") + with pytest.raises(ValueError, match="fileset_auto_create is False"): + await transfer.async_upload(async_client, local_path=str(tmp_path)) diff --git a/packages/models/src/models/resources.py b/packages/models/src/models/resources.py index 7ce619680b..e6c950319e 100644 --- a/packages/models/src/models/resources.py +++ b/packages/models/src/models/resources.py @@ -7,7 +7,6 @@ import logging import time from collections.abc import Awaitable, Callable -from dataclasses import dataclass from datetime import datetime from typing import TypeVar @@ -17,75 +16,20 @@ from nemo_platform.types.inference import ModelDeployment, ModelProvider from nemo_platform.types.inference.gateway.openai.v1 import OpenAIModelResp from nemo_platform.types.models import ModelEntity +from nemo_platform_plugin.models.refs import ( + ResolvedModelReference, + first_provider_ref, + model_entity_route_openai_url, + parse_workspace_name_ref, + resolved_model_reference, + warn_provider_host_url_resolution_failure, +) _T = TypeVar("_T") _TRANSIENT_GATEWAY_STATUS_CODES = {429, 502, 503, 504} _logger = logging.getLogger(__name__) -@dataclass(frozen=True, slots=True) -class ResolvedModelReference: - """Inference route details for a workspace-qualified model reference.""" - - url: str - name: str - host_url: str | None - - -def parse_workspace_name_ref(ref: str, *, label: str, expected_format: str = "workspace/name") -> tuple[str, str]: - """Parse a strict workspace-qualified reference.""" - workspace, separator, name = ref.partition("/") - if separator != "/" or not workspace or not name or "/" in name: - raise ValueError(f"{label} must be in format '{expected_format}'") - return workspace, name - - -def first_provider_ref(model_providers: list[str] | None) -> tuple[str, str, str] | None: - if not model_providers: - return None - - provider_ref = model_providers[0] - try: - provider_workspace, provider_name = parse_workspace_name_ref(provider_ref, label="Provider reference") - except ValueError: - _logger.warning("Invalid provider reference format", extra={"provider_ref": provider_ref}) - return None - return provider_ref, provider_workspace, provider_name - - -def model_entity_route_openai_url(*, base_url: str, workspace: str, name: str) -> str: - """OpenAI SDK-compatible URL for a model-entity proxy route.""" - return f"{base_url.rstrip('/')}/apis/inference-gateway/v2/workspaces/{workspace}/model/{name}/-/v1" - - -def resolved_model_reference( - *, - base_url: str, - name: str, - route_workspace: str, - route_model_name: str, - host_url: str | None, -) -> ResolvedModelReference: - """Build route details for a resolved model entity.""" - return ResolvedModelReference( - url=model_entity_route_openai_url(base_url=base_url, workspace=route_workspace, name=route_model_name), - name=name, - host_url=host_url, - ) - - -def warn_provider_host_url_resolution_failure( - provider_ref: str, - exc: Exception, - *, - not_found_error_type: type[Exception], -) -> None: - if isinstance(exc, not_found_error_type): - _logger.warning("Provider not found during host_url resolution", extra={"provider_ref": provider_ref}) - return - _logger.warning("Failed to resolve provider host_url", extra={"provider_ref": provider_ref}, exc_info=True) - - def _seconds_since_creation(entry_timestamp: datetime | str | None, created_at: datetime | None) -> int | None: """Seconds from deployment creation to the entry timestamp. Returns None if either is missing or not comparable.""" if created_at is None or entry_timestamp is None: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py index cb94196798..463fa17a5f 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py @@ -1,26 +1,27 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Adapter to create a :class:`NemoClient` from an existing :class:`NeMoPlatform`. +"""Adapter to create a typed client from an existing platform client. -This bridges the legacy ``NeMoPlatform`` SDK with the new typed client, -allowing plugins registered via ``NemoPluginSDKResources`` to use the -new endpoint/client infrastructure internally. +Accepts either a legacy ``NeMoPlatform`` SDK instance or a :class:`NemoClient` +/ :class:`AsyncNemoClient`, so plugins registered via ``NemoPluginSDKResources`` +can use the typed endpoint/client infrastructure regardless of which platform +client the caller holds. Usage:: from nemo_platform_plugin.client.adapter import client_from_platform - def make_sync_resource(platform: NeMoPlatform) -> NemoClient: + def make_sync_resource(platform: object) -> NemoClient: return client_from_platform(platform, NemoClient) """ from __future__ import annotations -from typing import TypeVar, overload +from collections.abc import Mapping +from typing import Protocol, TypeVar, cast, overload, runtime_checkable import httpx -from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_plugin.client.types import RetryPolicy @@ -28,70 +29,132 @@ def make_sync_resource(platform: NeMoPlatform) -> NemoClient: AsyncT = TypeVar("AsyncT", bound=AsyncNemoClient) +@runtime_checkable +class PlatformClient(Protocol): + """Structural shape shared by every platform handle :func:`client_from_platform` accepts. + + Satisfied by :class:`NemoClient` / :class:`AsyncNemoClient` and by the + generated ``NeMoPlatform`` / ``AsyncNeMoPlatform`` SDK classes. Use it to + annotate ``sdk`` / ``async_sdk`` parameters that are only forwarded to + :func:`client_from_platform`, so the annotating module does not need to + import the generated SDK. + """ + + @property + def base_url(self) -> str | httpx.URL: ... + + @property + def workspace(self) -> str | None: ... + + +class _PlatformClient(Protocol): + base_url: str | httpx.URL + workspace: str | None + max_retries: int + timeout: float | httpx.Timeout | None + _custom_headers: Mapping[str, str] + _client: httpx.Client | httpx.AsyncClient + + def _prepare_url(self, url: str) -> httpx.URL: ... + + +def platform_default_headers(platform: object) -> dict[str, str]: + """Return a copy of the default headers *platform* sends on every request. + + Reads ``default_headers`` off a :class:`NemoClient` / :class:`AsyncNemoClient` + and ``_custom_headers`` off a generated ``NeMoPlatform`` SDK instance, so + callers forwarding identity headers through a non-SDK HTTP client do not + need to know which platform handle they were given. + """ + if isinstance(platform, (NemoClient, AsyncNemoClient)): + return dict(platform.default_headers) + return dict(cast(_PlatformClient, platform)._custom_headers) + + @overload -def client_from_platform(platform: NeMoPlatform, client_cls: type[SyncT]) -> SyncT: ... +def client_from_platform(platform: object, client_cls: type[SyncT]) -> SyncT: ... @overload -def client_from_platform(platform: AsyncNeMoPlatform, client_cls: type[AsyncT]) -> AsyncT: ... +def client_from_platform(platform: object, client_cls: type[AsyncT]) -> AsyncT: ... def client_from_platform( - platform: NeMoPlatform | AsyncNeMoPlatform, + platform: object, client_cls: type[NemoClient] | type[AsyncNemoClient], ) -> NemoClient | AsyncNemoClient: - """Create a typed client sharing a generated platform SDK's transport. + """Create a typed client sharing a platform client's transport. + + When *platform* is already a :class:`NemoClient` or :class:`AsyncNemoClient` + the typed client is derived with ``client_cls.from_client`` and shares its + auth, headers, retry policy, and transport. Otherwise *platform* is treated + as a generated ``NeMoPlatform`` SDK instance. The overloads preserve the sync/async pairing between platform and client. """ + if isinstance(platform, AsyncNemoClient): + if not issubclass(client_cls, AsyncNemoClient): + raise TypeError("AsyncNemoClient requires an AsyncNemoClient class") + if isinstance(platform, client_cls): + return platform + return client_cls.from_client(platform) + if isinstance(platform, NemoClient): + if not issubclass(client_cls, NemoClient): + raise TypeError("NemoClient requires a NemoClient class") + if isinstance(platform, client_cls): + return platform + return client_cls.from_client(platform) + + platform_client = cast(_PlatformClient, platform) + # Prefer _custom_headers (set via with_options/set_default_headers), # fall back to the httpx client's actual headers (set at construction, # e.g. TestClient(headers={...})), filtering out httpx defaults. - headers = platform._custom_headers + headers = platform_client._custom_headers if not headers: _skip = {"accept", "accept-encoding", "connection", "user-agent", "host"} - headers = {k: v for k, v in platform._client.headers.items() if k.lower() not in _skip} + headers = {k: v for k, v in platform_client._client.headers.items() if k.lower() not in _skip} retry = RetryPolicy( - max_retries=platform.max_retries, + max_retries=platform_client.max_retries, retryable_status_codes=(408, 409, 429), retry_all_server_errors=True, respect_retry_decision_headers=True, respect_retry_after_headers=True, ) - url_resolver = platform._prepare_url + url_resolver = platform_client._prepare_url # Carry the platform's timeout across as a per-request override. The shared # httpx client keeps whatever timeout it was built with, so a caller's # ``platform.with_options(timeout=...)`` would otherwise be silently dropped # on the way to the typed client — the httpx client it hands over is the # *same* transport instance, with the *original* timeout still on it. - timeout = platform.timeout + timeout = platform_client.timeout if timeout is None: # ``None`` on the platform means "no timeout at all", but the typed # client reads None as "defer to the transport". Say the same thing in # the form httpx itself uses, so the override survives. timeout = httpx.Timeout(None) - if isinstance(platform, AsyncNeMoPlatform): + if isinstance(platform_client._client, httpx.AsyncClient): if not issubclass(client_cls, AsyncNemoClient): raise TypeError("AsyncNeMoPlatform requires an AsyncNemoClient class") return client_cls( - base_url=str(platform.base_url).rstrip("/"), - workspace=platform.workspace, + base_url=str(platform_client.base_url).rstrip("/"), + workspace=platform_client.workspace, default_headers=headers or None, timeout=timeout, retry=retry, - http_client=platform._client, + http_client=platform_client._client, url_resolver=url_resolver, ) if not issubclass(client_cls, NemoClient): raise TypeError("NeMoPlatform requires a NemoClient class") return client_cls( - base_url=str(platform.base_url).rstrip("/"), - workspace=platform.workspace, + base_url=str(platform_client.base_url).rstrip("/"), + workspace=platform_client.workspace, default_headers=headers or None, timeout=timeout, retry=retry, - http_client=platform._client, + http_client=platform_client._client, url_resolver=url_resolver, ) 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 18042bb49f..978ef0b90f 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 @@ -19,6 +19,7 @@ import asyncio import copy import email.utils +import inspect import json import logging import os @@ -285,6 +286,10 @@ def _should_retry( else: # A response only arrives once the body has gone out on the wire. body_is_spent = True + if response.status_code == 409 and (request.client_options or {}).get("exist_ok"): + # The caller declared the conflict an expected outcome that send() + # resolves by fetching the existing entity; retrying it is wasted work. + return None decision = response.headers.get("x-should-retry") if policy.respect_retry_decision_headers else None if policy.respect_retry_decision_headers and response.status_code < 400: return None @@ -486,6 +491,15 @@ def _request_headers(self, request: PreparedRequest) -> dict[str, str] | None: headers.update(request.extra_headers) return headers or None + def _needs_auth_header(self, headers: Mapping[str, str] | None) -> bool: + """Whether this attempt must carry a freshly resolved bearer token. + + Explicit ``Authorization`` values (per-call ``headers=`` or client defaults) + stay authoritative; otherwise the token provider is consulted on every + HTTP attempt so retries and later pages never replay an expired token. + """ + return self._auth is not None and not _has_header(headers, _AUTHORIZATION_HEADER) + def _is_binary(self, request: PreparedRequest) -> bool: return request.response_type is BinaryContent @@ -627,6 +641,30 @@ def agent_hardener(self) -> NemoClient | AsyncNemoClient: def inference(self: NemoClient | AsyncNemoClient) -> _InferenceNamespace: return _InferenceNamespace(self) + def __getattr__(self, name: str) -> Any: + """Resolve ``nemo.sdk`` plugin resource namespaces as client attributes. + + Only reached when normal attribute lookup fails. Plugin resources are + built with the sync or async factory matching this client and cached on + the instance, so ``client.example`` resolves the same object each time. + """ + if name.startswith("_"): + raise AttributeError(f"'{type(self).__name__}' object has no attribute {name!r}") + + from nemo_platform_plugin.discovery import discover_sdk + + resources = discover_sdk().get(name) + if resources is None: + raise AttributeError(f"'{type(self).__name__}' object has no attribute {name!r}") + + factory = resources.async_resource if isinstance(self, AsyncNemoClient) else resources.sync_resource + if factory is None: + raise AttributeError(f"'{type(self).__name__}' object has no attribute {name!r}") + + instance = factory(self) + self.__dict__[name] = instance + return instance + def _resolve_query_params(self, request: PreparedRequest) -> dict[str, str | int | bool] | None: """Filter out None values and JSON-serialize dicts/lists in query params.""" if request.query_params is None: @@ -783,14 +821,8 @@ def send( if headers: request = request.with_headers(headers) - # Inject auth header if a TokenProvider is configured. - # NOTE: If a 401 occurs despite this, a future enhancement could - # call provider.force_refresh() and retry once. The proactive - # refresh margin (60s) makes this unlikely in practice. - if self._auth: - token = self._auth.get_access_token() - request = request.with_headers({"Authorization": f"Bearer {token}"}) - + # The bearer token is resolved per HTTP attempt (see _authorized_headers), + # not once per logical request, so retries and later pages use a live token. url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) @@ -824,6 +856,15 @@ def send( body = _parse_response_body(request.response_type, raw) return NemoResponse(http_response=raw, body=body, request=request) + def _authorized_headers(self, headers: dict[str, str] | None) -> dict[str, str] | None: + """Return *headers* with a bearer token resolved for this attempt when the provider owns auth.""" + if self._auth is None or not self._needs_auth_header(headers): + return headers + token = self._auth.get_access_token() + if inspect.isawaitable(token): + raise TypeError("Async token provider used on a synchronous client; use AsyncNemoClient.") + return {**(headers or {}), _AUTHORIZATION_HEADER: f"Bearer {token}"} + def _request_with_retry( self, request: PreparedRequest, @@ -836,7 +877,11 @@ def _request_with_retry( last_response: httpx.Response | None = None for attempt in range(retry.max_retries + 1 if retry else 1): try: - kwargs: dict = {"content": request.content, "headers": headers, "params": params} + kwargs: dict = { + "content": request.content, + "headers": self._authorized_headers(headers), + "params": params, + } if self._timeout is not None: kwargs["timeout"] = self._timeout raw = self._http.request(request.method, url, **kwargs) @@ -870,7 +915,11 @@ def _stream_with_retry( for attempt in range(retry.max_retries + 1 if retry else 1): yielded = False try: - kwargs: dict = {"content": request.content, "headers": headers, "params": params} + kwargs: dict = { + "content": request.content, + "headers": self._authorized_headers(headers), + "params": params, + } if self._timeout is not None: kwargs["timeout"] = self._timeout with self._http.stream(request.method, url, **kwargs) as raw: @@ -1040,9 +1089,6 @@ async def send( if headers: request = request.with_headers(headers) - if self._auth: - request = request.with_headers({"Authorization": f"Bearer {await resolve_token_async(self._auth)}"}) - url = self._resolve_path(request) req_headers = self._request_headers(request) params = self._resolve_query_params(request) @@ -1076,6 +1122,13 @@ async def send( body = _parse_response_body(request.response_type, raw) return NemoResponse(http_response=raw, body=body, request=request) + async def _authorized_headers(self, headers: dict[str, str] | None) -> dict[str, str] | None: + """Return *headers* with a bearer token resolved for this attempt when the provider owns auth.""" + if self._auth is None or not self._needs_auth_header(headers): + return headers + token = await resolve_token_async(self._auth) + return {**(headers or {}), _AUTHORIZATION_HEADER: f"Bearer {token}"} + async def _request_with_retry( self, request: PreparedRequest, @@ -1088,7 +1141,11 @@ async def _request_with_retry( last_response: httpx.Response | None = None for attempt in range(retry.max_retries + 1 if retry else 1): try: - kwargs: dict = {"content": request.content, "headers": headers, "params": params} + kwargs: dict = { + "content": request.content, + "headers": await self._authorized_headers(headers), + "params": params, + } if self._timeout is not None: kwargs["timeout"] = self._timeout raw = await self._http.request(request.method, url, **kwargs) @@ -1122,7 +1179,11 @@ async def _stream_with_retry( for attempt in range(retry.max_retries + 1 if retry else 1): yielded = False try: - kwargs: dict = {"content": request.content, "headers": headers, "params": params} + kwargs: dict = { + "content": request.content, + "headers": await self._authorized_headers(headers), + "params": params, + } if self._timeout is not None: kwargs["timeout"] = self._timeout async with self._http.stream(request.method, url, **kwargs) as raw: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/endpoints.py index b1443cb51a..535051f9c4 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/endpoints.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/endpoints.py @@ -23,6 +23,7 @@ OtlpExportLogsResponse, OtlpLogQueryRequest, UpdateFilesetRequest, + UploadOtlpLogsQueryParams, ) from nemo_platform_plugin.jobs.schemas import PlatformJobLogPage @@ -102,7 +103,11 @@ def delete_file(*, workspace: str | None = None, name: str, path: str) -> Filese @post("/apis/files/v2/workspaces/{workspace}/filesets/{name}/otlp/v1/logs") @abstractmethod def upload_otlp_logs( - *, workspace: str | None = None, name: str, content: bytes | Iterable[bytes] | AsyncIterable[bytes] + *, + workspace: str | None = None, + name: str, + content: bytes | Iterable[bytes] | AsyncIterable[bytes], + query_params: UploadOtlpLogsQueryParams | None = None, ) -> OtlpExportLogsResponse: ... diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/types.py index c32bd1c890..44de08bdcc 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/files/types.py @@ -160,6 +160,10 @@ class ListFilesQueryParams(TypedDict, total=False): include_cache_status: NotRequired[bool] +class UploadOtlpLogsQueryParams(TypedDict, total=False): + artifact_base_path: NotRequired[str] + + # --------------------------------------------------------------------------- # OTLP types # --------------------------------------------------------------------------- diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/guardrail/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/guardrail/types.py index 6922334bb0..41c55435a8 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/guardrail/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/guardrail/types.py @@ -29,7 +29,7 @@ class GuardrailConfig(BaseModel): workspace: str project: str | None = None description: str | None = None - data: dict[str, Any] = Field(default_factory=dict, description="Guardrail configuration data") + data: dict[str, Any] | None = Field(default=None, description="Guardrail configuration data") id: str created_at: datetime created_by: str | None = None diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/endpoints.py index 6f9e88799f..f8b1d3463c 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/endpoints.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/endpoints.py @@ -22,18 +22,13 @@ @get("/apis/auth/v2/iam/role-bindings") @abstractmethod -def list_role_bindings( - *, - query_params: ListRoleBindingsQueryParams = {"page": 1, "page_size": 10, "sort": "created_at"}, -) -> Paginated[RoleBinding]: ... +def list_role_bindings(*, query_params: ListRoleBindingsQueryParams | None = None) -> Paginated[RoleBinding]: ... @post("/apis/auth/v2/iam/role-bindings") @abstractmethod def create_role_binding( - *, - body: RoleBindingInput, - query_params: RolePropagationQueryParams = {"wait_role_propagation": True}, + *, body: RoleBindingInput, query_params: RolePropagationQueryParams | None = None ) -> RoleBinding: ... @@ -45,7 +40,7 @@ def get_role_binding(*, name: str) -> RoleBinding: ... @delete("/apis/auth/v2/iam/role-bindings/{name}") @abstractmethod def revoke_role_binding( - *, name: str, query_params: RolePropagationQueryParams = {"wait_role_propagation": True} + *, name: str, query_params: RolePropagationQueryParams | None = None ) -> RoleBindingDeleteResponse: ... diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/types.py index 8c05075d3c..cfbdd7054d 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/iam/types.py @@ -83,6 +83,7 @@ class AuthzErrorResponse(BaseModel): "page": NotRequired[int], "page_size": NotRequired[int], "sort": NotRequired[str], + "filter": NotRequired[str], "filter[principal]": NotRequired[str], "filter[principal][$eq]": NotRequired[str], "filter[principal][$like]": NotRequired[str], diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/client.py index f8bff69315..78ba604dc0 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/client.py @@ -1,9 +1,17 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Typed clients for narrow Inference Gateway provider proxy calls.""" +"""Typed sync and async clients for the Inference Gateway proxy routes. -from __future__ import annotations +Usage:: + + client = InferenceGatewayClient(base_url="...", workspace="default") + client.provider_ready(name="nvidia-build") + completion = client.openai_post( + trailing_uri="v1/chat/completions", + body=JsonBody({"model": "default/my-model", "messages": [...]}), + ).data() +""" import json @@ -27,12 +35,30 @@ def _decode_provider_proxy_body(content: bytes) -> object: return text -class _InferenceGatewayProviderMethods: +class _InferenceGatewayMethods: + provider_get = method(endpoints.provider_get) + provider_post = method(endpoints.provider_post) + provider_put = method(endpoints.provider_put) + provider_patch = method(endpoints.provider_patch) + provider_delete = method(endpoints.provider_delete) + provider_ready = method(endpoints.provider_ready) get_provider_models_raw = method(endpoints.get_provider_models_raw) + stream_provider = method(endpoints.stream_provider) + model_get = method(endpoints.model_get) + model_post = method(endpoints.model_post) + model_put = method(endpoints.model_put) + model_patch = method(endpoints.model_patch) + model_delete = method(endpoints.model_delete) + stream_model = method(endpoints.stream_model) + openai_get = method(endpoints.openai_get) + openai_post = method(endpoints.openai_post) + stream_openai = method(endpoints.stream_openai) + list_openai_models = method(endpoints.list_openai_models) + get_openai_model = method(endpoints.get_openai_model) -class InferenceGatewayProviderClient(_InferenceGatewayProviderMethods, NemoClient): - """Sync client for Inference Gateway provider proxy reads.""" +class InferenceGatewayClient(_InferenceGatewayMethods, NemoClient): + """Sync client for the Inference Gateway proxy routes.""" def get_provider_models( self, @@ -47,8 +73,8 @@ def get_provider_models( return _decode_provider_proxy_body(response.read()) -class AsyncInferenceGatewayProviderClient(_InferenceGatewayProviderMethods, AsyncNemoClient): - """Async client for Inference Gateway provider proxy reads.""" +class AsyncInferenceGatewayClient(_InferenceGatewayMethods, AsyncNemoClient): + """Async client for the Inference Gateway proxy routes.""" async def get_provider_models( self, @@ -61,3 +87,11 @@ async def get_provider_models( client = self.with_options(timeout=timeout) if timeout is not None else self response = await client.get_provider_models_raw(workspace=workspace, name=name) return _decode_provider_proxy_body(await response.read()) + + +class InferenceGatewayProviderClient(InferenceGatewayClient): + """Sync client for the narrow provider proxy reads the Models reconciler makes.""" + + +class AsyncInferenceGatewayProviderClient(AsyncInferenceGatewayClient): + """Async client for the narrow provider proxy reads the Models reconciler makes.""" diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/endpoints.py index d3af3a2ad1..37f07280b4 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/endpoints.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/endpoints.py @@ -1,18 +1,136 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Typed endpoint definitions for narrow Inference Gateway provider calls.""" +"""Typed endpoint definitions for the Inference Gateway proxy routes. + +Three proxy surfaces share the ``/-/{trailing_uri}`` convention, where +``trailing_uri`` is the provider-relative path such as ``v1/chat/completions``: + +- ``provider``: route straight to a registered ModelProvider by name. +- ``model``: route through a model entity (``workspace/name``). +- ``openai``: the workspace's OpenAI-compatible surface, where the body's + ``model`` field selects the VirtualModel. + +``stream_openai`` and ``stream_model`` return the raw SSE byte stream so callers +can decode events (including ``event: error`` frames) themselves. +""" from __future__ import annotations from abc import abstractmethod +from typing import Any -from nemo_platform_plugin.client.endpoint import get +from nemo_platform_plugin.client.endpoint import delete, get, patch, post, put from nemo_platform_plugin.client.types import BinaryContent +from nemo_platform_plugin.inference_gateway.types import JsonBody, OpenAIModel, OpenAIModelList, ProviderReadyResponse + +_BASE = "/apis/inference-gateway/v2/workspaces/{workspace}" + +# --------------------------------------------------------------------------- +# Provider proxy +# --------------------------------------------------------------------------- + + +@get(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def provider_get(*, workspace: str | None = None, name: str, trailing_uri: str) -> Any: ... -_PROVIDER = "/apis/inference-gateway/v2/workspaces/{workspace}/provider/{name}/-" + +@post(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def provider_post(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... + + +@put(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def provider_put(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... + + +@patch(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def provider_patch(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... -@get(_PROVIDER + "/v1/models") +@delete(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def provider_delete(*, workspace: str | None = None, name: str, trailing_uri: str) -> None: ... + + +@post(f"{_BASE}/provider/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def stream_provider(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> BinaryContent: ... + + +@get(f"{_BASE}/provider/{{name}}/-/v1/models") @abstractmethod def get_provider_models_raw(*, workspace: str | None = None, name: str) -> BinaryContent: ... + + +@get(f"{_BASE}/provider/{{name}}/ready") +@abstractmethod +def provider_ready(*, workspace: str | None = None, name: str) -> ProviderReadyResponse: ... + + +# --------------------------------------------------------------------------- +# Model entity proxy +# --------------------------------------------------------------------------- + + +@get(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def model_get(*, workspace: str | None = None, name: str, trailing_uri: str) -> Any: ... + + +@post(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def model_post(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... + + +@put(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def model_put(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... + + +@patch(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def model_patch(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> Any: ... + + +@delete(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def model_delete(*, workspace: str | None = None, name: str, trailing_uri: str) -> None: ... + + +@post(f"{_BASE}/model/{{name}}/-/{{trailing_uri}}") +@abstractmethod +def stream_model(*, workspace: str | None = None, name: str, trailing_uri: str, body: JsonBody) -> BinaryContent: ... + + +# --------------------------------------------------------------------------- +# OpenAI-compatible surface +# --------------------------------------------------------------------------- + + +@get(f"{_BASE}/openai/-/{{trailing_uri}}") +@abstractmethod +def openai_get(*, workspace: str | None = None, trailing_uri: str) -> Any: ... + + +@post(f"{_BASE}/openai/-/{{trailing_uri}}") +@abstractmethod +def openai_post(*, workspace: str | None = None, trailing_uri: str, body: JsonBody) -> Any: ... + + +@post(f"{_BASE}/openai/-/{{trailing_uri}}") +@abstractmethod +def stream_openai(*, workspace: str | None = None, trailing_uri: str, body: JsonBody) -> BinaryContent: ... + + +@get(f"{_BASE}/openai/-/v1/models") +@abstractmethod +def list_openai_models(*, workspace: str | None = None) -> OpenAIModelList: ... + + +@get(f"{_BASE}/openai/-/v1/models/{{name}}") +@abstractmethod +def get_openai_model(*, workspace: str | None = None, name: str) -> OpenAIModel: ... diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/types.py new file mode 100644 index 0000000000..81c0d4dfe6 --- /dev/null +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/inference_gateway/types.py @@ -0,0 +1,49 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Wire shapes for the Inference Gateway proxy routes. + +The gateway forwards arbitrary OpenAI-compatible (or provider-native) payloads, +so request and response bodies are open JSON objects rather than fixed models. +""" + +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field, RootModel + +JsonObject = dict[str, Any] + + +class JsonBody(RootModel[JsonObject]): + """Arbitrary JSON object request body for proxied inference calls.""" + + +class OpenAIModel(BaseModel): + """One entry of an OpenAI-compatible ``/v1/models`` listing.""" + + model_config = ConfigDict(extra="allow") + + id: str + object: str = "model" + created: int | None = None + owned_by: str | None = None + + +class OpenAIModelList(BaseModel): + """OpenAI-compatible ``/v1/models`` response.""" + + model_config = ConfigDict(extra="allow") + + data: list[OpenAIModel] = Field(default_factory=list) + object: str = "list" + + +class ProviderReadyResponse(BaseModel): + """Gateway readiness payload for a registered provider.""" + + model_config = ConfigDict(extra="allow") + + workspace: str | None = None + name: str | None = None diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py index 31a7821595..848abbffdb 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py @@ -7,7 +7,7 @@ from collections.abc import Callable from dataclasses import dataclass -from typing import Generic, TypeVar +from typing import Any, Generic, TypeVar from nemo_platform import AsyncNeMoPlatform, NeMoPlatform @@ -17,14 +17,17 @@ @dataclass(frozen=True, slots=True) class NemoPluginSDKResources(Generic[SyncResourceT, AsyncResourceT]): - """Container for plugin SDK resources exposed on legacy platform SDK owners. + """Container for plugin SDK resources exposed as platform client namespaces. - Typed clients should expose resources through explicit typed APIs instead - of consuming this dynamic legacy ``nemo.sdk`` entry-point surface. + Each factory receives the owning platform client (a ``NeMoPlatform`` or a + :class:`~nemo_platform_plugin.client.client.NemoClient`, sync or async) and + returns the plugin's resource object. Typed clients should expose resources + through explicit typed APIs instead of consuming this dynamic ``nemo.sdk`` + entry-point surface. """ - sync_resource: Callable[[NeMoPlatform], SyncResourceT] | None = None - async_resource: Callable[[AsyncNeMoPlatform], AsyncResourceT] | None = None + sync_resource: Callable[[Any], SyncResourceT] | None = None + async_resource: Callable[[Any], AsyncResourceT] | None = None def __post_init__(self) -> None: if self.sync_resource is None and self.async_resource is None: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/virtual_models/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/virtual_models/endpoints.py index eda8a30735..59a96e541f 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/virtual_models/endpoints.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/virtual_models/endpoints.py @@ -8,7 +8,7 @@ from abc import abstractmethod from nemo_platform_plugin.client.endpoint import delete, get, patch, post -from nemo_platform_plugin.client.types import Paginated +from nemo_platform_plugin.client.types import Paginated, PreparedRequest from nemo_platform_plugin.virtual_models.types import ( CreateVirtualModelRequest, DeleteVirtualModelQueryParams, @@ -20,9 +20,23 @@ _VIRTUAL_MODELS = "/apis/inference-gateway/v2/workspaces/{workspace}/virtual-models" -@post(_VIRTUAL_MODELS) +@get(_VIRTUAL_MODELS + "/{name}") @abstractmethod -def create_virtual_model(*, workspace: str | None = None, body: CreateVirtualModelRequest) -> VirtualModel: ... +def get_virtual_model(*, workspace: str | None = None, name: str) -> VirtualModel: ... + + +def _get_virtual_model_on_conflict( + body: CreateVirtualModelRequest, workspace: str | None +) -> PreparedRequest[VirtualModel]: + """Retrieve request replayed when ``create_virtual_model(exist_ok=True)`` 409s.""" + return get_virtual_model(name=body.name, workspace=workspace) + + +@post(_VIRTUAL_MODELS, get_on_conflict=_get_virtual_model_on_conflict) +@abstractmethod +def create_virtual_model( + *, workspace: str | None = None, body: CreateVirtualModelRequest, exist_ok: bool = False +) -> VirtualModel: ... @get(_VIRTUAL_MODELS) @@ -32,11 +46,6 @@ def list_virtual_models( ) -> Paginated[VirtualModel]: ... -@get(_VIRTUAL_MODELS + "/{name}") -@abstractmethod -def get_virtual_model(*, workspace: str | None = None, name: str) -> VirtualModel: ... - - @patch(_VIRTUAL_MODELS + "/{name}") @abstractmethod def update_virtual_model( diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/endpoints.py index 055ae2c25d..5d3003aa96 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/endpoints.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/endpoints.py @@ -24,6 +24,7 @@ Workspace, WorkspaceMember, WorkspaceMemberListResponse, + WorkspaceMemberQueryParams, ) # --------------------------------------------------------------------------- @@ -73,21 +74,32 @@ def delete_workspace(*, name: str) -> DeleteResponse: ... @get("/apis/entities/v2/workspaces/{workspace}/members") @abstractmethod -def list_workspace_members(*, workspace: str) -> WorkspaceMemberListResponse: ... +def list_workspace_members(*, workspace: str | None = None) -> WorkspaceMemberListResponse: ... @post("/apis/entities/v2/workspaces/{workspace}/members") @abstractmethod -def create_workspace_member(*, workspace: str, body: CreateWorkspaceMemberRequest) -> WorkspaceMember: ... +def create_workspace_member( + *, + workspace: str | None = None, + body: CreateWorkspaceMemberRequest, + query_params: WorkspaceMemberQueryParams | None = None, +) -> WorkspaceMember: ... @put("/apis/entities/v2/workspaces/{workspace}/members/{principal_id}") @abstractmethod def update_workspace_member( - *, workspace: str, principal_id: str, body: UpdateWorkspaceMemberRequest + *, + workspace: str | None = None, + principal_id: str, + body: UpdateWorkspaceMemberRequest, + query_params: WorkspaceMemberQueryParams | None = None, ) -> WorkspaceMember: ... @delete("/apis/entities/v2/workspaces/{workspace}/members/{principal_id}") @abstractmethod -def delete_workspace_member(*, workspace: str, principal_id: str) -> DeleteResponse: ... +def delete_workspace_member( + *, workspace: str | None = None, principal_id: str, query_params: WorkspaceMemberQueryParams | None = None +) -> DeleteResponse: ... diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/types.py index 261ff0fc8e..41f02af1e4 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/workspaces/types.py @@ -94,3 +94,9 @@ class ListWorkspacesQueryParams(TypedDict, total=False): class CreateWorkspaceQueryParams(TypedDict, total=False): wait_role_propagation: NotRequired[bool] + + +class WorkspaceMemberQueryParams(TypedDict, total=False): + """Query parameters shared by the member create/update/delete endpoints.""" + + wait_role_propagation: NotRequired[bool] diff --git a/packages/nemo_platform_plugin/tests/client/test_adapter.py b/packages/nemo_platform_plugin/tests/client/test_adapter.py index 5119b39162..9affdca398 100644 --- a/packages/nemo_platform_plugin/tests/client/test_adapter.py +++ b/packages/nemo_platform_plugin/tests/client/test_adapter.py @@ -8,7 +8,8 @@ import httpx import pytest from nemo_platform import AsyncNeMoPlatform, NeMoPlatform -from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.adapter import PlatformClient, client_from_platform, platform_default_headers +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_plugin.client.types import RetryPolicy from nemo_platform_plugin.jobs import endpoints from nemo_platform_plugin.jobs.client import AsyncJobsClient, JobsClient @@ -163,3 +164,61 @@ def test_client_from_platform_carries_disabled_timeout() -> None: # Not the transport's 60s: httpx reads an all-None Timeout as "wait forever". assert client._timeout == httpx.Timeout(None) + + +def test_client_from_platform_accepts_a_typed_client_and_shares_its_transport() -> None: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + base = NemoClient(base_url="http://test", workspace="ws", default_headers={"X-A": "1"}, http_client=http_client) + + client = client_from_platform(base, JobsClient) + + assert isinstance(client, JobsClient) + assert client._http is http_client + assert client.workspace == "ws" + assert client.default_headers == {"X-A": "1"} + assert client_from_platform(client, JobsClient) is client + + +@pytest.mark.asyncio +async def test_client_from_platform_accepts_an_async_typed_client() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + base = AsyncNemoClient(base_url="http://test", workspace="ws", http_client=http_client) + + client = client_from_platform(base, AsyncJobsClient) + + assert isinstance(client, AsyncJobsClient) + assert client._http is http_client + with pytest.raises(TypeError): + client_from_platform(base, JobsClient) + + +def test_platform_client_protocol_matches_both_platform_handle_shapes() -> None: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + generated = NeMoPlatform(base_url="http://test", workspace="ws", http_client=http_client) + typed = NemoClient(base_url="http://test", workspace="ws", http_client=http_client) + + assert isinstance(generated, PlatformClient) + assert isinstance(typed, PlatformClient) + assert isinstance(AsyncNemoClient(base_url="http://test", http_client=httpx.AsyncClient()), PlatformClient) + assert not isinstance(object(), PlatformClient) + + +def test_platform_default_headers_reads_typed_client_headers() -> None: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + client = NemoClient(base_url="http://test", default_headers={"X-NMP-Internal": "true"}, http_client=http_client) + + headers = platform_default_headers(client) + + assert headers == {"X-NMP-Internal": "true"} + headers["mutated"] = "yes" + assert client.default_headers == {"X-NMP-Internal": "true"} + + +def test_platform_default_headers_reads_generated_sdk_custom_headers() -> None: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + platform = NeMoPlatform( + base_url="http://test", default_headers={"X-NMP-Principal-Id": "service:agents"}, http_client=http_client + ) + + assert platform_default_headers(platform) == {"X-NMP-Principal-Id": "service:agents"} + assert platform_default_headers(AsyncNemoClient(base_url="http://test", http_client=httpx.AsyncClient())) == {} diff --git a/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py b/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py new file mode 100644 index 0000000000..bea22614a1 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py @@ -0,0 +1,241 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""The bearer token is resolved for every HTTP attempt, not once per logical request. + +A token provider refreshes proactively, but that only helps if it is consulted +each time bytes go on the wire. Pages after the first and retries after a +backoff must therefore carry whatever the provider returns at that moment. +""" + +from __future__ import annotations + +from typing import Any +from urllib.parse import parse_qs + +import httpx +import pytest +from nemo_platform_plugin.client.auth import TokenProviderAuth +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.endpoint import get +from nemo_platform_plugin.client.types import BinaryContent, Paginated, RetryPolicy +from pydantic import BaseModel + +BASE = "http://test" + + +class Item(BaseModel): + name: str + + +@get("/apis/test/v2/items") +def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: ... + + +@get("/apis/test/v2/items/{name}") +def get_item(*, name: str) -> Item: ... + + +@get("/apis/test/v2/download") +def download() -> BinaryContent: ... + + +class RotatingProvider: + """Returns ``token-1``, ``token-2``, ... on successive calls.""" + + def __init__(self) -> None: + self.calls = 0 + + def get_access_token(self) -> str: + self.calls += 1 + return f"token-{self.calls}" + + +class AsyncRotatingProvider: + def __init__(self) -> None: + self.calls = 0 + + async def get_access_token(self) -> str: + self.calls += 1 + return f"token-{self.calls}" + + +def _page_body(page: int, total_pages: int) -> dict[str, Any]: + return { + "data": [{"name": f"item-{page}"}], + "pagination": { + "page": page, + "page_size": 1, + "current_page_size": 1, + "total_pages": total_pages, + "total_results": total_pages, + }, + } + + +def _paging_handler(seen: list[str]) -> Any: + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request.headers["Authorization"]) + page = int(parse_qs(request.url.query.decode()).get("page", ["1"])[0]) + return httpx.Response(200, json=_page_body(page, 3)) + + return handler + + +def _retry_handler(seen: list[str], failures: int) -> Any: + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request.headers["Authorization"]) + if len(seen) <= failures: + return httpx.Response(503, json={"detail": "busy"}) + return httpx.Response(200, json={"name": "alice"}) + + return handler + + +NO_BACKOFF = RetryPolicy(max_retries=2, backoff_base=0.0) + + +# --------------------------------------------------------------------------- +# sync +# --------------------------------------------------------------------------- + + +def test_each_page_uses_a_freshly_resolved_token() -> None: + seen: list[str] = [] + client = NemoClient( + base_url=BASE, + auth=RotatingProvider(), + http_client=httpx.Client(transport=httpx.MockTransport(_paging_handler(seen))), + ) + + names = [item.name for item in client.send(list_items()).items()] + + assert names == ["item-1", "item-2", "item-3"] + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] + + +def test_each_retry_attempt_uses_a_freshly_resolved_token() -> None: + seen: list[str] = [] + client = NemoClient( + base_url=BASE, + auth=RotatingProvider(), + retry=NO_BACKOFF, + http_client=httpx.Client(transport=httpx.MockTransport(_retry_handler(seen, failures=2))), + ) + + assert client.send(get_item(name="alice")).body.name == "alice" + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] + + +def test_stream_retry_attempts_use_a_freshly_resolved_token() -> None: + """Binary downloads go through the streaming path, which has its own attempt loop.""" + seen: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request.headers["Authorization"]) + if len(seen) == 1: + raise httpx.ConnectError("no route to host", request=request) + return httpx.Response(200, stream=httpx.ByteStream(b"payload")) + + client = NemoClient( + base_url=BASE, + auth=RotatingProvider(), + retry=NO_BACKOFF, + http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + with client.send(download()).stream() as chunks: + assert b"".join(chunks) == b"payload" + assert seen == ["Bearer token-1", "Bearer token-2"] + + +def test_explicit_authorization_header_is_not_overridden() -> None: + seen: list[str] = [] + provider = RotatingProvider() + client = NemoClient( + base_url=BASE, + auth=provider, + http_client=httpx.Client(transport=httpx.MockTransport(_paging_handler(seen))), + ) + + list(client.send(list_items(), headers={"Authorization": "Bearer pinned"}).items()) + + assert seen == ["Bearer pinned"] * 3 + assert provider.calls == 0 + + +def test_transport_auth_hook_does_not_double_resolve() -> None: + """Clients built with TokenProviderAuth on the transport still resolve once per attempt.""" + seen: list[str] = [] + provider = RotatingProvider() + client = NemoClient( + base_url=BASE, + auth=provider, + http_client=httpx.Client( + transport=httpx.MockTransport(_paging_handler(seen)), auth=TokenProviderAuth(provider) + ), + ) + + list(client.send(list_items()).items()) + + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] + assert provider.calls == 3 + + +def test_sync_client_rejects_async_provider_at_send_time() -> None: + client = NemoClient( + base_url=BASE, + auth=AsyncRotatingProvider(), # type: ignore[arg-type] + http_client=httpx.Client(transport=httpx.MockTransport(_paging_handler([]))), + ) + + with pytest.raises(TypeError, match="Async token provider"): + client.send(get_item(name="alice")) + + +# --------------------------------------------------------------------------- +# async +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_async_each_page_uses_a_freshly_resolved_token() -> None: + seen: list[str] = [] + client = AsyncNemoClient( + base_url=BASE, + auth=AsyncRotatingProvider(), + http_client=httpx.AsyncClient(transport=httpx.MockTransport(_paging_handler(seen))), + ) + + names = [item.name async for item in (await client.send(list_items())).items()] + + assert names == ["item-1", "item-2", "item-3"] + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] + + +@pytest.mark.asyncio +async def test_async_each_retry_attempt_uses_a_freshly_resolved_token() -> None: + seen: list[str] = [] + client = AsyncNemoClient( + base_url=BASE, + auth=AsyncRotatingProvider(), + retry=NO_BACKOFF, + http_client=httpx.AsyncClient(transport=httpx.MockTransport(_retry_handler(seen, failures=2))), + ) + + assert (await client.send(get_item(name="alice"))).body.name == "alice" + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] + + +@pytest.mark.asyncio +async def test_async_client_accepts_sync_provider() -> None: + seen: list[str] = [] + client = AsyncNemoClient( + base_url=BASE, + auth=RotatingProvider(), + http_client=httpx.AsyncClient(transport=httpx.MockTransport(_paging_handler(seen))), + ) + + [item async for item in (await client.send(list_items())).items()] + + assert seen == ["Bearer token-1", "Bearer token-2", "Bearer token-3"] 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 b4c0020b3d..8ca2ce7363 100644 --- a/packages/nemo_platform_plugin/tests/client/test_client_options.py +++ b/packages/nemo_platform_plugin/tests/client/test_client_options.py @@ -203,6 +203,39 @@ def test_409_with_exist_ok_replays_get_and_returns_entity(self) -> None: assert mock.request.call_count == 2 assert mock.request.call_args_list[1].args[0] == "GET" + def test_409_with_exist_ok_is_not_retried_before_resolving(self) -> None: + """A retry policy that lists 409 must not replay the POST when the caller opted into exist_ok.""" + mock = _mock_http( + _resp(409, {"detail": "Item 'alice' already exists"}), + _resp(200, {"id": 7, "name": "alice"}, http_method="GET", url=f"{BASE}/apis/test/v2/items/alice"), + ) + client = NemoClient( + base_url=BASE, + http_client=mock, + retry=RetryPolicy(max_retries=2, retryable_status_codes=(409,), backoff_base=0.0), + ) + + resp = client.send(CREATE_ITEM(ItemRequest(name="alice"), exist_ok=True)) + + assert resp.body is not None and resp.body.id == 7 + assert [call.args[0] for call in mock.request.call_args_list] == ["POST", "GET"] + + def test_409_without_exist_ok_is_still_retried_when_the_policy_says_so(self) -> None: + mock = _mock_http( + _resp(409, {"detail": "Item 'alice' already exists"}), + _resp(201, {"id": 1, "name": "alice"}), + ) + client = NemoClient( + base_url=BASE, + http_client=mock, + retry=RetryPolicy(max_retries=2, retryable_status_codes=(409,), backoff_base=0.0), + ) + + resp = client.send(CREATE_ITEM(ItemRequest(name="alice"))) + + assert resp.http_response.status_code == 201 + assert [call.args[0] for call in mock.request.call_args_list] == ["POST", "POST"] + def test_409_without_exist_ok_raises_conflict(self) -> None: mock = _mock_http(_resp(409, {"detail": "Item 'alice' already exists"})) client = NemoClient(base_url=BASE, http_client=mock) diff --git a/packages/nemo_platform_plugin/tests/files/test_endpoints.py b/packages/nemo_platform_plugin/tests/files/test_endpoints.py index c3f72fc060..104f764eda 100644 --- a/packages/nemo_platform_plugin/tests/files/test_endpoints.py +++ b/packages/nemo_platform_plugin/tests/files/test_endpoints.py @@ -5,6 +5,7 @@ from __future__ import annotations +import json from typing import get_origin from nemo_platform_plugin.client.types import BinaryContent, Paginated, PreparedRequest @@ -14,6 +15,8 @@ FilesetFileOutput, FilesetOutput, ListFilesetFilesResponse, + OtlpExportLogsResponse, + OtlpLogQueryRequest, UpdateFilesetRequest, ) @@ -146,3 +149,33 @@ def test_update_fileset_excludes_unset_fields() -> None: assert content == {"description": "updated"} assert "purpose" not in content assert "metadata" not in content + + +def test_upload_otlp_logs() -> None: + prepared = endpoints.upload_otlp_logs(workspace="default", name="my-fileset", content=b'{"resourceLogs": []}') + + assert prepared.method == "POST" + assert prepared.path_template == "/apis/files/v2/workspaces/{workspace}/filesets/{name}/otlp/v1/logs" + assert prepared.path_params == {"workspace": "default", "name": "my-fileset"} + assert prepared.content == b'{"resourceLogs": []}' + assert prepared.query_params is None + assert prepared.response_type is OtlpExportLogsResponse + + +def test_upload_otlp_logs_with_artifact_base_path() -> None: + prepared = endpoints.upload_otlp_logs( + workspace="default", name="my-fileset", content=b"{}", query_params={"artifact_base_path": "logs/run-1"} + ) + + assert prepared.query_params == {"artifact_base_path": "logs/run-1"} + + +def test_query_otlp_logs() -> None: + body = OtlpLogQueryRequest(limit=10, artifact_base_path="logs/run-1") + prepared = endpoints.query_otlp_logs(workspace="default", name="my-fileset", body=body) + + assert prepared.method == "POST" + assert prepared.path_template == "/apis/files/v2/workspaces/{workspace}/filesets/{name}/otlp/v1/logs/query" + assert prepared.path_params == {"workspace": "default", "name": "my-fileset"} + assert json.loads(prepared.content) == {"limit": 10, "artifact_base_path": "logs/run-1"} + assert prepared.content_type == "application/json" diff --git a/packages/nemo_platform_plugin/tests/guardrail/test_endpoints.py b/packages/nemo_platform_plugin/tests/guardrail/test_endpoints.py new file mode 100644 index 0000000000..69d0e6ff6d --- /dev/null +++ b/packages/nemo_platform_plugin/tests/guardrail/test_endpoints.py @@ -0,0 +1,183 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for Guardrails service endpoint definitions and response types.""" + +from __future__ import annotations + +import json +from typing import Any, get_origin + +from nemo_platform_plugin.client.types import Paginated, PreparedRequest +from nemo_platform_plugin.entities.types import DeleteResponse +from nemo_platform_plugin.guardrail import endpoints +from nemo_platform_plugin.guardrail.types import ( + CreateGuardrailConfigRequest, + GuardrailCheckRequest, + GuardrailCheckResponse, + GuardrailConfig, + UpdateGuardrailConfigRequest, +) + +CONFIGS = "/apis/guardrails/v2/workspaces/{workspace}/configs" + + +def _json_content(prepared: PreparedRequest[Any]) -> Any: + assert isinstance(prepared.content, bytes) + return json.loads(prepared.content) + + +def test_get_guardrail_config() -> None: + prepared = endpoints.get_guardrail_config(workspace="default", name="safety") + + assert isinstance(prepared, PreparedRequest) + assert prepared.method == "GET" + assert prepared.path_template == f"{CONFIGS}/{{name}}" + assert prepared.path_params == {"workspace": "default", "name": "safety"} + assert prepared.content is None + assert prepared.response_type is GuardrailConfig + + +def test_get_guardrail_config_workspace_optional() -> None: + prepared = endpoints.get_guardrail_config(name="safety") + + assert prepared.path_params == {"name": "safety"} + + +def test_list_guardrail_configs() -> None: + prepared = endpoints.list_guardrail_configs(workspace="default") + + assert prepared.method == "GET" + assert prepared.path_template == CONFIGS + assert prepared.path_params == {"workspace": "default"} + assert prepared.query_params is None + assert get_origin(prepared.response_type) is Paginated + + +def test_list_guardrail_configs_with_query_params() -> None: + prepared = endpoints.list_guardrail_configs( + workspace="default", query_params={"page": 2, "page_size": 5, "sort": "-name", "filter": '{"name": "x"}'} + ) + + assert prepared.query_params == {"page": 2, "page_size": 5, "sort": "-name", "filter": '{"name": "x"}'} + + +def test_create_guardrail_config() -> None: + body = CreateGuardrailConfigRequest(name="safety", description="d", data={"rails": {}}) + prepared = endpoints.create_guardrail_config(workspace="default", body=body) + + assert prepared.method == "POST" + assert prepared.path_template == CONFIGS + assert prepared.path_params == {"workspace": "default"} + assert prepared.content_type == "application/json" + assert _json_content(prepared) == {"name": "safety", "description": "d", "data": {"rails": {}}} + assert prepared.response_type is GuardrailConfig + + +def test_create_guardrail_config_only_sends_set_fields() -> None: + prepared = endpoints.create_guardrail_config(workspace="default", body=CreateGuardrailConfigRequest(name="s")) + + assert _json_content(prepared) == {"name": "s"} + + +def test_create_guardrail_config_exist_ok_builds_retrieve_request() -> None: + body = CreateGuardrailConfigRequest(name="safety") + prepared = endpoints.create_guardrail_config(workspace="default", body=body, exist_ok=True) + + assert prepared.client_options == {"exist_ok": True} + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.method == "GET" + assert prepared.on_conflict_get.path_template == f"{CONFIGS}/{{name}}" + assert prepared.on_conflict_get.path_params == {"workspace": "default", "name": "safety"} + + +def test_update_guardrail_config() -> None: + body = UpdateGuardrailConfigRequest(description="changed") + prepared = endpoints.update_guardrail_config(workspace="default", name="safety", body=body) + + assert prepared.method == "PATCH" + assert prepared.path_template == f"{CONFIGS}/{{name}}" + assert prepared.path_params == {"workspace": "default", "name": "safety"} + assert _json_content(prepared) == {"description": "changed"} + assert prepared.response_type is GuardrailConfig + + +def test_update_guardrail_config_empty_body() -> None: + prepared = endpoints.update_guardrail_config( + workspace="default", name="safety", body=UpdateGuardrailConfigRequest() + ) + + assert _json_content(prepared) == {} + + +def test_delete_guardrail_config() -> None: + prepared = endpoints.delete_guardrail_config(workspace="default", name="safety") + + assert prepared.method == "DELETE" + assert prepared.path_template == f"{CONFIGS}/{{name}}" + assert prepared.path_params == {"workspace": "default", "name": "safety"} + assert prepared.content is None + assert prepared.response_type is DeleteResponse + + +def test_check_guardrail() -> None: + body = GuardrailCheckRequest( + model="m", messages=[{"role": "user", "content": "hi"}], guardrails={"config_id": "default/safety"} + ) + prepared = endpoints.check_guardrail(workspace="default", body=body) + + assert prepared.method == "POST" + assert prepared.path_template == "/apis/guardrails/v2/workspaces/{workspace}/checks" + assert prepared.path_params == {"workspace": "default"} + assert _json_content(prepared) == { + "model": "m", + "messages": [{"role": "user", "content": "hi"}], + "guardrails": {"config_id": "default/safety"}, + } + assert prepared.response_type is GuardrailCheckResponse + + +def test_check_guardrail_passes_through_extra_sampling_params() -> None: + body = GuardrailCheckRequest.model_validate({"model": "m", "messages": [], "seed": 7, "stop": ["\n"]}) + prepared = endpoints.check_guardrail(workspace="default", body=body) + + assert _json_content(prepared) == {"model": "m", "messages": [], "seed": 7, "stop": ["\n"]} + + +def test_guardrail_config_accepts_null_data() -> None: + """The create route returns ``data: null`` for a config created without data.""" + config = GuardrailConfig.model_validate( + { + "name": "safety", + "workspace": "default", + "project": None, + "description": None, + "data": None, + "id": "guardrail-config-1", + "created_at": "2026-01-01T00:00:00", + "created_by": "service:guardrails", + "updated_at": "2026-01-01T00:00:00", + "updated_by": "service:guardrails", + "entity_id": "guardrail-config-1", + "parent": None, + "db_version": 1, + } + ) + + assert config.data is None + assert config.name == "safety" + + +def test_guardrail_config_accepts_omitted_data() -> None: + """GET/list use ``response_model_exclude_none`` and omit ``data`` entirely.""" + config = GuardrailConfig.model_validate( + { + "name": "safety", + "workspace": "default", + "id": "guardrail-config-1", + "created_at": "2026-01-01T00:00:00", + "updated_at": "2026-01-01T00:00:00", + } + ) + + assert config.data is None diff --git a/packages/nemo_platform_plugin/tests/iam/test_client.py b/packages/nemo_platform_plugin/tests/iam/test_client.py index 45cf581533..ab545e1f6c 100644 --- a/packages/nemo_platform_plugin/tests/iam/test_client.py +++ b/packages/nemo_platform_plugin/tests/iam/test_client.py @@ -57,7 +57,7 @@ def test_sync_role_binding_and_authz_responses() -> None: assert created.name == "rb-123" assert decision.result == {"allowed": True} - assert mock_http.request.call_args_list[0].kwargs["params"] == {"wait_role_propagation": True} + assert mock_http.request.call_args_list[0].kwargs["params"] is None def test_role_binding_pagination_dtos_and_filter_encoding() -> None: @@ -186,9 +186,9 @@ async def test_async_iam_request_dispatch_parity() -> None: ("POST", f"{BASE}/apis/auth/v2/authz/allow"), ] assert calls[0].kwargs["params"] == {"filter[role]": "Viewer"} - assert calls[1].kwargs["params"] == {"wait_role_propagation": True} + assert calls[1].kwargs["params"] is None assert json.loads(calls[1].kwargs["content"]) == body.model_dump() - assert calls[2].kwargs["params"] == {"wait_role_propagation": True} + assert calls[2].kwargs["params"] is None assert json.loads(calls[3].kwargs["content"]) == {"input": {"principal_id": "user@example.com"}} diff --git a/packages/nemo_platform_plugin/tests/iam/test_endpoints.py b/packages/nemo_platform_plugin/tests/iam/test_endpoints.py index 9f4b20d191..64d5ac5437 100644 --- a/packages/nemo_platform_plugin/tests/iam/test_endpoints.py +++ b/packages/nemo_platform_plugin/tests/iam/test_endpoints.py @@ -30,19 +30,29 @@ def test_role_binding_endpoint_contracts() -> None: assert listed.query_params == {"filter[principal][$like]": "service:%"} assert get_origin(listed.response_type) is Paginated assert created.method == "POST" - assert created.query_params == {"wait_role_propagation": True} + assert created.query_params is None assert created.response_type is RoleBinding assert fetched.path_params == {"name": "rb-123"} assert fetched.response_type is RoleBinding assert revoked.method == "DELETE" - assert revoked.query_params == {"wait_role_propagation": True} + assert revoked.query_params is None assert revoked.response_type is RoleBindingDeleteResponse -def test_list_role_bindings_uses_server_query_defaults() -> None: - prepared = endpoints.list_role_bindings() - - assert prepared.query_params == {"page": 1, "page_size": 10, "sort": "created_at"} +def test_role_binding_endpoints_send_only_supplied_query_params() -> None: + """Server-side defaults (page, page_size, sort, wait_role_propagation) are not replicated client-side.""" + assert endpoints.list_role_bindings().query_params is None + assert endpoints.list_role_bindings(query_params={"page": 2, "filter": 'role:"Viewer"'}).query_params == { + "page": 2, + "filter": 'role:"Viewer"', + } + assert endpoints.create_role_binding( + body=RoleBindingInput(principal="user@example.com", role="Viewer"), + query_params={"wait_role_propagation": False}, + ).query_params == {"wait_role_propagation": False} + assert endpoints.revoke_role_binding(name="rb-123", query_params={"wait_role_propagation": False}).query_params == { + "wait_role_propagation": False + } def test_authz_and_bundle_endpoint_contracts() -> None: diff --git a/packages/nemo_platform_plugin/tests/inference_gateway/test_endpoints.py b/packages/nemo_platform_plugin/tests/inference_gateway/test_endpoints.py new file mode 100644 index 0000000000..65f99eed3d --- /dev/null +++ b/packages/nemo_platform_plugin/tests/inference_gateway/test_endpoints.py @@ -0,0 +1,222 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Inference Gateway proxy endpoints: prepared-request shape and on-the-wire path encoding.""" + +from __future__ import annotations + +import json +from typing import Any + +import httpx +import pytest +from nemo_platform_plugin.client.client import NemoClient +from nemo_platform_plugin.client.types import BinaryContent, PreparedRequest +from nemo_platform_plugin.inference_gateway import endpoints +from nemo_platform_plugin.inference_gateway.client import InferenceGatewayClient +from nemo_platform_plugin.inference_gateway.types import JsonBody, OpenAIModel, OpenAIModelList, ProviderReadyResponse + +BASE = "/apis/inference-gateway/v2/workspaces/{workspace}" + + +def _json_body(prepared: PreparedRequest) -> dict: + assert isinstance(prepared.content, bytes) + return json.loads(prepared.content) + + +# --------------------------------------------------------------------------- +# prepared requests +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + ("endpoint", "http_method", "surface"), + [ + (endpoints.provider_get, "GET", "provider"), + (endpoints.provider_post, "POST", "provider"), + (endpoints.provider_put, "PUT", "provider"), + (endpoints.provider_patch, "PATCH", "provider"), + (endpoints.provider_delete, "DELETE", "provider"), + (endpoints.model_get, "GET", "model"), + (endpoints.model_post, "POST", "model"), + (endpoints.model_put, "PUT", "model"), + (endpoints.model_patch, "PATCH", "model"), + (endpoints.model_delete, "DELETE", "model"), + ], +) +def test_named_proxy_routes(endpoint: Any, http_method: str, surface: str) -> None: + kwargs: dict[str, Any] = {"workspace": "ws", "name": "nvidia-build", "trailing_uri": "v1/chat/completions"} + if http_method in {"POST", "PUT", "PATCH"}: + kwargs["body"] = JsonBody({"model": "m"}) + + prepared = endpoint(**kwargs) + + assert prepared.method == http_method + assert prepared.path_template == f"{BASE}/{surface}/{{name}}/-/{{trailing_uri}}" + assert prepared.path_params == {"workspace": "ws", "name": "nvidia-build", "trailing_uri": "v1/chat/completions"} + if "body" in kwargs: + assert _json_body(prepared) == {"model": "m"} + assert prepared.content_type == "application/json" + else: + assert prepared.content is None + + +@pytest.mark.parametrize( + ("endpoint", "http_method"), + [(endpoints.openai_get, "GET"), (endpoints.openai_post, "POST")], +) +def test_openai_proxy_routes(endpoint: Any, http_method: str) -> None: + kwargs: dict[str, Any] = {"workspace": "ws", "trailing_uri": "v1/models"} + if http_method == "POST": + kwargs["body"] = JsonBody({"model": "default/vm", "messages": []}) + + prepared = endpoint(**kwargs) + + assert prepared.method == http_method + assert prepared.path_template == f"{BASE}/openai/-/{{trailing_uri}}" + assert prepared.path_params == {"workspace": "ws", "trailing_uri": "v1/models"} + assert prepared.response_type is Any + + +@pytest.mark.parametrize("endpoint", [endpoints.stream_provider, endpoints.stream_model]) +def test_named_stream_routes_return_raw_bytes(endpoint: Any) -> None: + prepared = endpoint(name="n", trailing_uri="v1/chat/completions", body=JsonBody({"stream": True})) + + assert prepared.method == "POST" + assert prepared.response_type is BinaryContent + assert _json_body(prepared) == {"stream": True} + + +def test_stream_openai_returns_raw_bytes() -> None: + prepared = endpoints.stream_openai(trailing_uri="v1/chat/completions", body=JsonBody({"stream": True})) + + assert prepared.method == "POST" + assert prepared.path_template == f"{BASE}/openai/-/{{trailing_uri}}" + assert prepared.response_type is BinaryContent + + +def test_provider_ready() -> None: + prepared = endpoints.provider_ready(name="nvidia-build") + + assert prepared.method == "GET" + assert prepared.path_template == f"{BASE}/provider/{{name}}/ready" + assert prepared.path_params == {"name": "nvidia-build"} + assert prepared.response_type is ProviderReadyResponse + + +def test_openai_model_listing_routes() -> None: + listing = endpoints.list_openai_models(workspace="ws") + single = endpoints.get_openai_model(workspace="ws", name="default/vm") + + assert listing.path_template == f"{BASE}/openai/-/v1/models" + assert listing.response_type is OpenAIModelList + assert single.path_template == f"{BASE}/openai/-/v1/models/{{name}}" + assert single.path_params == {"workspace": "ws", "name": "default/vm"} + assert single.response_type is OpenAIModel + + +def test_workspace_defaults_to_the_client_workspace() -> None: + prepared = endpoints.openai_get(trailing_uri="v1/models") + + assert "workspace" not in prepared.path_params + + +def test_json_body_serializes_arbitrary_objects() -> None: + prepared = endpoints.openai_post(trailing_uri="v1/embeddings", body=JsonBody({"input": ["a", "b"], "n": 1})) + + assert _json_body(prepared) == {"input": ["a", "b"], "n": 1} + + +# --------------------------------------------------------------------------- +# on the wire +# --------------------------------------------------------------------------- + + +class _Recorder: + def __init__(self, response: httpx.Response | None = None) -> None: + self.requests: list[httpx.Request] = [] + self.response = response or httpx.Response(200, json={"ok": True}) + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return self.response + + +def _client(recorder: _Recorder) -> InferenceGatewayClient: + return InferenceGatewayClient( + base_url="http://test", + workspace="default", + http_client=httpx.Client(transport=httpx.MockTransport(recorder)), + ) + + +def test_trailing_uri_slashes_are_percent_encoded_on_the_wire() -> None: + """The gateway declares ``{trailing_uri:path}`` and decodes before routing. + + Every path parameter is encoded with ``quote(safe="")``, so the caller-supplied + provider path travels as one segment. This pins the wire form the server is + known to accept; ``request.url.path`` would decode it and pass vacuously. + """ + recorder = _Recorder() + + _client(recorder).openai_post(trailing_uri="v1/chat/completions", body=JsonBody({"model": "m"})) + + assert recorder.requests[0].url.raw_path == ( + b"/apis/inference-gateway/v2/workspaces/default/openai/-/v1%2Fchat%2Fcompletions" + ) + + +def test_named_route_encodes_both_name_and_trailing_uri() -> None: + recorder = _Recorder() + + _client(recorder).provider_get(name="my provider", trailing_uri="v1/models") + + assert recorder.requests[0].method == "GET" + assert recorder.requests[0].url.raw_path == ( + b"/apis/inference-gateway/v2/workspaces/default/provider/my%20provider/-/v1%2Fmodels" + ) + + +def test_untyped_proxy_response_is_returned_as_parsed_json() -> None: + recorder = _Recorder(httpx.Response(200, json={"choices": [{"message": {"content": "hi"}}]})) + + data = _client(recorder).openai_post(trailing_uri="v1/chat/completions", body=JsonBody({"model": "m"})).data() + + assert data == {"choices": [{"message": {"content": "hi"}}]} + + +def test_stream_openai_yields_the_raw_sse_bytes() -> None: + sse = b'data: {"choices":[{"delta":{"content":"h"}}]}\n\ndata: [DONE]\n\n' + recorder = _Recorder( + httpx.Response(200, stream=httpx.ByteStream(sse), headers={"content-type": "text/event-stream"}) + ) + + response = _client(recorder).stream_openai(trailing_uri="v1/chat/completions", body=JsonBody({"stream": True})) + with response.stream() as chunks: + assert b"".join(chunks) == sse + + assert recorder.requests[0].headers["content-type"] == "application/json" + assert json.loads(recorder.requests[0].content) == {"stream": True} + + +def test_list_openai_models_is_typed_and_keeps_extra_fields() -> None: + recorder = _Recorder( + httpx.Response(200, json={"object": "list", "data": [{"id": "default/vm", "object": "model", "extra": 1}]}) + ) + + models = _client(recorder).list_openai_models().data() + + assert isinstance(models, OpenAIModelList) + assert models.data[0].id == "default/vm" + assert models.data[0].model_dump()["extra"] == 1 + + +def test_from_client_shares_transport_and_workspace() -> None: + recorder = _Recorder(httpx.Response(200, json={"ready": True})) + base = NemoClient( + base_url="http://test", workspace="ws", http_client=httpx.Client(transport=httpx.MockTransport(recorder)) + ) + + InferenceGatewayClient.from_client(base).provider_ready(name="p") + + assert recorder.requests[0].url.path == "/apis/inference-gateway/v2/workspaces/ws/provider/p/ready" diff --git a/packages/nemo_platform_plugin/tests/virtual_models/test_endpoints.py b/packages/nemo_platform_plugin/tests/virtual_models/test_endpoints.py index 218c83fe0b..a61f319ca2 100644 --- a/packages/nemo_platform_plugin/tests/virtual_models/test_endpoints.py +++ b/packages/nemo_platform_plugin/tests/virtual_models/test_endpoints.py @@ -22,6 +22,8 @@ VirtualModelInferenceConfig, ) +PATH = "/apis/inference-gateway/v2/workspaces/{workspace}/virtual-models" + def _json_body(prepared: PreparedRequest) -> dict[str, object]: assert isinstance(prepared.content, bytes) @@ -46,6 +48,34 @@ def test_create_omits_unset_fields_and_nested_nones() -> None: } +def test_create_prebuilds_conflict_retrieve_for_exist_ok() -> None: + """``exist_ok`` replays the GET for the same name on a 409; off by default.""" + prepared = endpoints.create_virtual_model(workspace="default", body=CreateVirtualModelRequest(name="router")) + + assert prepared.client_options == {"exist_ok": False} + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.method == "GET" + assert prepared.on_conflict_get.path_template == PATH + "/{name}" + assert prepared.on_conflict_get.path_params == {"workspace": "default", "name": "router"} + + prepared = endpoints.create_virtual_model(body=CreateVirtualModelRequest(name="router"), exist_ok=True) + assert prepared.client_options == {"exist_ok": True} + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.path_params == {"name": "router"} + + +def test_delete_sends_expected_db_version_as_query_param() -> None: + prepared = endpoints.delete_virtual_model(workspace="default", name="router") + assert prepared.method == "DELETE" + assert prepared.path_template == PATH + "/{name}" + assert prepared.query_params is None + + prepared = endpoints.delete_virtual_model( + workspace="default", name="router", query_params={"expected_db_version": 4} + ) + assert prepared.query_params == {"expected_db_version": 4} + + def test_update_distinguishes_explicit_null_and_empty_list_from_unset() -> None: """PATCH must send only what the caller set, so omitted fields stay unchanged.""" body = UpdateVirtualModelRequest(default_model_entity=None, request_middleware=[]) diff --git a/packages/nemo_platform_plugin/tests/workspaces/test_client.py b/packages/nemo_platform_plugin/tests/workspaces/test_client.py new file mode 100644 index 0000000000..d5514495cd --- /dev/null +++ b/packages/nemo_platform_plugin/tests/workspaces/test_client.py @@ -0,0 +1,174 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for WorkspacesClient / AsyncWorkspacesClient over a recording httpx transport.""" + +from __future__ import annotations + +import json + +import httpx +import pytest +from nemo_platform_plugin.client.errors import ConflictError, NotFoundError +from nemo_platform_plugin.workspaces.client import AsyncWorkspacesClient, WorkspacesClient +from nemo_platform_plugin.workspaces.types import ( + CreateWorkspaceMemberRequest, + CreateWorkspaceRequest, + UpdateWorkspaceMemberRequest, +) + +BASE = "http://test:8000" + +WORKSPACE = { + "id": "8a4d8f1e-0f7f-4b8a-9d5f-2f4b6c8e1a2b", + "name": "ml-team", + "description": None, + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", +} +MEMBER = {"principal": "user@example.com", "roles": ["Editor"], "granted_at": None, "granted_by": "system"} + + +class Recorder: + def __init__(self, responses: list[httpx.Response]) -> None: + self.responses = responses + self.requests: list[httpx.Request] = [] + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return self.responses.pop(0) + + @property + def last(self) -> httpx.Request: + return self.requests[-1] + + +def make_client(recorder: Recorder, *, workspace: str | None = "default") -> WorkspacesClient: + return WorkspacesClient( + base_url=BASE, workspace=workspace, http_client=httpx.Client(transport=httpx.MockTransport(recorder)) + ) + + +def test_create_workspace_sends_query_params() -> None: + recorder = Recorder([httpx.Response(201, json=WORKSPACE)]) + client = make_client(recorder) + + workspace = client.create_workspace( + body=CreateWorkspaceRequest(name="ml-team"), query_params={"wait_role_propagation": False} + ).data() + + assert workspace.name == "ml-team" + assert recorder.last.method == "POST" + assert recorder.last.url.path == "/apis/entities/v2/workspaces" + assert dict(recorder.last.url.params) == {"wait_role_propagation": "false"} + assert json.loads(recorder.last.content) == {"name": "ml-team"} + + +def test_create_workspace_exist_ok_replays_get_on_conflict() -> None: + recorder = Recorder([httpx.Response(409, json={"detail": "exists"}), httpx.Response(200, json=WORKSPACE)]) + client = make_client(recorder) + + workspace = client.create_workspace(body=CreateWorkspaceRequest(name="ml-team"), exist_ok=True).data() + + assert workspace.name == "ml-team" + assert [(r.method, r.url.path) for r in recorder.requests] == [ + ("POST", "/apis/entities/v2/workspaces"), + ("GET", "/apis/entities/v2/workspaces/ml-team"), + ] + + +def test_create_workspace_conflict_raises_without_exist_ok() -> None: + recorder = Recorder([httpx.Response(409, json={"detail": "exists"})]) + client = make_client(recorder) + + with pytest.raises(ConflictError): + client.create_workspace(body=CreateWorkspaceRequest(name="ml-team")) + + +def test_get_workspace_not_found_raises() -> None: + recorder = Recorder([httpx.Response(404, json={"detail": "Workspace 'missing' not found"})]) + client = make_client(recorder) + + with pytest.raises(NotFoundError) as exc: + client.get_workspace(name="missing") + assert exc.value.status_code == 404 + + +def test_list_workspace_members_uses_client_default_workspace() -> None: + recorder = Recorder([httpx.Response(200, json={"data": [MEMBER]})]) + client = make_client(recorder) + + members = client.list_workspace_members().data() + + assert [m.principal for m in members.data] == ["user@example.com"] + assert recorder.last.url.path == "/apis/entities/v2/workspaces/default/members" + + +def test_list_workspace_members_without_workspace_raises() -> None: + recorder = Recorder([]) + client = make_client(recorder, workspace=None) + + with pytest.raises(ValueError, match="Missing path parameter 'workspace'"): + client.list_workspace_members() + assert recorder.requests == [] + + +def test_create_workspace_member_sends_body_and_query_params() -> None: + recorder = Recorder([httpx.Response(201, json=MEMBER)]) + client = make_client(recorder) + + member = client.create_workspace_member( + workspace="ml-team", + body=CreateWorkspaceMemberRequest(principal="user@example.com", roles=["Editor"]), + query_params={"wait_role_propagation": False}, + ).data() + + assert member.principal == "user@example.com" + assert recorder.last.method == "POST" + assert recorder.last.url.path == "/apis/entities/v2/workspaces/ml-team/members" + assert dict(recorder.last.url.params) == {"wait_role_propagation": "false"} + assert json.loads(recorder.last.content) == {"principal": "user@example.com", "roles": ["Editor"]} + + +def test_update_workspace_member_encodes_principal_in_path() -> None: + recorder = Recorder([httpx.Response(200, json={**MEMBER, "roles": ["Viewer"]})]) + client = make_client(recorder) + + member = client.update_workspace_member( + principal_id="user@example.com", body=UpdateWorkspaceMemberRequest(roles=["Viewer"]) + ).data() + + assert member.roles == ["Viewer"] + assert recorder.last.method == "PUT" + assert recorder.last.url.raw_path == b"/apis/entities/v2/workspaces/default/members/user%40example.com" + assert dict(recorder.last.url.params) == {} + assert json.loads(recorder.last.content) == {"roles": ["Viewer"]} + + +def test_delete_workspace_member_sends_query_params() -> None: + recorder = Recorder([httpx.Response(200, json={"message": "deleted", "id": "user-123", "deleted_at": None})]) + client = make_client(recorder) + + client.delete_workspace_member( + workspace="ml-team", principal_id="user-123", query_params={"wait_role_propagation": True} + ) + + assert recorder.last.method == "DELETE" + assert recorder.last.url.path == "/apis/entities/v2/workspaces/ml-team/members/user-123" + assert dict(recorder.last.url.params) == {"wait_role_propagation": "true"} + + +async def test_async_client_create_member() -> None: + recorder = Recorder([httpx.Response(201, json=MEMBER)]) + client = AsyncWorkspacesClient( + base_url=BASE, workspace="default", http_client=httpx.AsyncClient(transport=httpx.MockTransport(recorder)) + ) + + response = await client.create_workspace_member( + body=CreateWorkspaceMemberRequest(principal="user@example.com"), query_params={"wait_role_propagation": False} + ) + + assert response.data().principal == "user@example.com" + assert recorder.last.url.path == "/apis/entities/v2/workspaces/default/members" + assert dict(recorder.last.url.params) == {"wait_role_propagation": "false"} + assert json.loads(recorder.last.content) == {"principal": "user@example.com"} diff --git a/packages/nemo_platform_plugin/tests/workspaces/test_endpoints.py b/packages/nemo_platform_plugin/tests/workspaces/test_endpoints.py new file mode 100644 index 0000000000..85ff567772 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/workspaces/test_endpoints.py @@ -0,0 +1,169 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for Workspaces service endpoint definitions.""" + +from __future__ import annotations + +import json +from typing import get_origin + +from nemo_platform_plugin.client.types import Paginated, PreparedRequest +from nemo_platform_plugin.entities.types import DeleteResponse +from nemo_platform_plugin.workspaces import endpoints +from nemo_platform_plugin.workspaces.types import ( + CreateWorkspaceMemberRequest, + CreateWorkspaceRequest, + UpdateWorkspaceMemberRequest, + UpdateWorkspaceRequest, + Workspace, + WorkspaceMember, + WorkspaceMemberListResponse, +) + + +def _json_body(prepared: PreparedRequest) -> dict: + """Decode a prepared request's JSON body (asserting it is present bytes).""" + assert isinstance(prepared.content, bytes) + return json.loads(prepared.content) + + +def test_get_workspace() -> None: + prepared = endpoints.get_workspace(name="ml-team") + + assert isinstance(prepared, PreparedRequest) + assert prepared.method == "GET" + assert prepared.path_template == "/apis/entities/v2/workspaces/{name}" + assert prepared.path_params == {"name": "ml-team"} + assert prepared.response_type is Workspace + + +def test_list_workspaces() -> None: + prepared = endpoints.list_workspaces() + + assert prepared.method == "GET" + assert prepared.path_template == "/apis/entities/v2/workspaces" + assert prepared.path_params == {} + assert prepared.query_params is None + assert get_origin(prepared.response_type) is Paginated + + +def test_list_workspaces_with_query_params() -> None: + prepared = endpoints.list_workspaces(query_params={"page": 2, "page_size": 5, "sort": "-name", "filter": "x"}) + + assert prepared.query_params == {"page": 2, "page_size": 5, "sort": "-name", "filter": "x"} + + +def test_create_workspace_serializes_only_set_fields() -> None: + prepared = endpoints.create_workspace(body=CreateWorkspaceRequest(name="ml-team")) + + assert prepared.method == "POST" + assert prepared.path_template == "/apis/entities/v2/workspaces" + assert prepared.content_type == "application/json" + assert _json_body(prepared) == {"name": "ml-team"} + assert prepared.query_params is None + assert prepared.response_type is Workspace + + +def test_create_workspace_with_query_params_and_exist_ok() -> None: + prepared = endpoints.create_workspace( + body=CreateWorkspaceRequest(name="ml-team", description="d"), + query_params={"wait_role_propagation": False}, + exist_ok=True, + ) + + assert _json_body(prepared) == {"name": "ml-team", "description": "d"} + assert prepared.query_params == {"wait_role_propagation": False} + assert prepared.client_options == {"exist_ok": True} + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.path_params == {"name": "ml-team"} + + +def test_update_workspace() -> None: + prepared = endpoints.update_workspace(name="ml-team", body=UpdateWorkspaceRequest(description="new")) + + assert prepared.method == "PUT" + assert prepared.path_template == "/apis/entities/v2/workspaces/{name}" + assert prepared.path_params == {"name": "ml-team"} + assert _json_body(prepared) == {"description": "new"} + assert prepared.response_type is Workspace + + +def test_delete_workspace() -> None: + prepared = endpoints.delete_workspace(name="ml-team") + + assert prepared.method == "DELETE" + assert prepared.path_params == {"name": "ml-team"} + assert prepared.content is None + assert prepared.response_type is DeleteResponse + + +def test_list_workspace_members() -> None: + prepared = endpoints.list_workspace_members(workspace="ml-team") + + assert prepared.method == "GET" + assert prepared.path_template == "/apis/entities/v2/workspaces/{workspace}/members" + assert prepared.path_params == {"workspace": "ml-team"} + assert prepared.response_type is WorkspaceMemberListResponse + + +def test_list_workspace_members_workspace_optional() -> None: + prepared = endpoints.list_workspace_members() + + assert prepared.path_params == {} + + +def test_create_workspace_member() -> None: + body = CreateWorkspaceMemberRequest(principal="user@example.com", roles=["Viewer"]) + prepared = endpoints.create_workspace_member(workspace="ml-team", body=body) + + assert prepared.method == "POST" + assert prepared.path_template == "/apis/entities/v2/workspaces/{workspace}/members" + assert prepared.path_params == {"workspace": "ml-team"} + assert _json_body(prepared) == {"principal": "user@example.com", "roles": ["Viewer"]} + assert prepared.query_params is None + assert prepared.response_type is WorkspaceMember + + +def test_create_workspace_member_omits_default_roles() -> None: + prepared = endpoints.create_workspace_member(body=CreateWorkspaceMemberRequest(principal="user@example.com")) + + assert prepared.path_params == {} + assert _json_body(prepared) == {"principal": "user@example.com"} + + +def test_create_workspace_member_query_params() -> None: + prepared = endpoints.create_workspace_member( + workspace="ml-team", + body=CreateWorkspaceMemberRequest(principal="user@example.com"), + query_params={"wait_role_propagation": False}, + ) + + assert prepared.query_params == {"wait_role_propagation": False} + + +def test_update_workspace_member() -> None: + prepared = endpoints.update_workspace_member( + workspace="ml-team", + principal_id="user@example.com", + body=UpdateWorkspaceMemberRequest(roles=["Viewer", "Editor"]), + query_params={"wait_role_propagation": True}, + ) + + assert prepared.method == "PUT" + assert prepared.path_template == "/apis/entities/v2/workspaces/{workspace}/members/{principal_id}" + assert prepared.path_params == {"workspace": "ml-team", "principal_id": "user@example.com"} + assert _json_body(prepared) == {"roles": ["Viewer", "Editor"]} + assert prepared.query_params == {"wait_role_propagation": True} + assert prepared.response_type is WorkspaceMember + + +def test_delete_workspace_member() -> None: + prepared = endpoints.delete_workspace_member(principal_id="user-123", query_params={"wait_role_propagation": False}) + + assert prepared.method == "DELETE" + assert prepared.path_template == "/apis/entities/v2/workspaces/{workspace}/members/{principal_id}" + assert prepared.path_params == {"principal_id": "user-123"} + assert prepared.query_params == {"wait_role_propagation": False} + assert prepared.content is None + assert prepared.response_type is DeleteResponse From b2fa8ad667af5b2a989ede743688232b72f87226 Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Thu, 10 Sep 2026 17:21:23 -0400 Subject: [PATCH 2/8] feat(cli): typed-client seam and shared helpers for hand-written commands Additive plumbing so command groups can be rewritten on the typed clients one PR at a time while the generated commands keep working: - client/bootstrap.py resolves config, OIDC discovery and token providers without the generated SDK and builds NemoClient/AsyncNemoClient with the same retry policy, TLS verification and 5 s connect cap the SDK used; factory.py layers the NeMoPlatform constructors on top of it. - CLIContext.typed_client(XClient) / async_typed_client derive a service client that shares the platform client's transport and auth; get_workspace exposes the configured default. - pagination gains collect_offset_pages / collect_cursor_pages for typed paginated responses, carrying the server's envelope fields (sort, filter, grouped_by) into JSON output; fetch_all_pages stays for generated commands. - errors maps typed-client and pydantic errors alongside the SDK ones (validation and unknown --input-data keys exit 2), and recognises every unresolved-workspace message. - stdin_utils.build_request_body validates --input-data into a request model and rejects unknown keys instead of dropping them. - code_generator renders typed-client snippets (request models as constructor calls, SecretStr masked, RootModel positional); generated commands are routed to legacy_code_generator until they are replaced. - waiters accept either platform handle, poll with a monotonic clock and normalise str-enum statuses; formatters unwrap NemoResponse; the version flag reads distribution metadata; command groups no longer advertise shell completion (the root app owns it). Signed-off-by: Max Dubrinsky --- .../src/nemo_platform_ext/cli/app.py | 4 +- .../cli/core/autocomplete.py | 5 +- .../cli/core/code_generator.py | 476 +++++------ .../src/nemo_platform_ext/cli/core/context.py | 27 + .../src/nemo_platform_ext/cli/core/errors.py | 45 +- .../nemo_platform_ext/cli/core/formatters.py | 9 + .../cli/core/help_formatter.py | 3 + .../cli/core/legacy_code_generator.py | 385 +++++++++ .../nemo_platform_ext/cli/core/pagination.py | 185 ++++- .../nemo_platform_ext/cli/core/stdin_utils.py | 57 +- .../src/nemo_platform_ext/cli/core/waiters.py | 62 +- .../nemo_platform_ext/cli/telemetry/emit.py | 5 +- .../src/nemo_platform_ext/cli/version.py | 24 + .../src/nemo_platform_ext/client/bootstrap.py | 780 ++++++++++++++++++ .../src/nemo_platform_ext/client/factory.py | 639 ++------------ .../tests/cli/core/test_code_generator.py | 372 ++++----- .../tests/cli/core/test_context.py | 29 + .../tests/cli/core/test_errors.py | 77 ++ .../tests/cli/core/test_help_formatter.py | 15 + .../cli/core/test_legacy_code_generator.py | 284 +++++++ .../tests/cli/core/test_pagination.py | 284 ++++++- .../tests/cli/core/test_stdin_utils.py | 61 ++ .../tests/cli/core/test_waiters.py | 157 ++-- .../tests/cli/telemetry/test_job_events.py | 22 +- .../tests/client/test_bootstrap_builders.py | 253 ++++++ .../tests/client/test_client.py | 62 +- .../nmp_common/tests/sdk_factory/test_sdk.py | 4 +- 27 files changed, 3071 insertions(+), 1255 deletions(-) create mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py create mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/version.py create mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/client/bootstrap.py create mode 100644 packages/nemo_platform_ext/tests/cli/core/test_legacy_code_generator.py create mode 100644 packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/app.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/app.py index 38b47571a3..b89bb6e278 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/app.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/app.py @@ -330,9 +330,9 @@ def main( def _version_callback(value: bool) -> None: """Print version information and exit.""" if value: - import nemo_platform + from nemo_platform_ext.cli.version import client_version - typer.echo(f"nemo version {nemo_platform.__version__}") + typer.echo(f"nemo version {client_version()}") raise typer.Exit() diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/autocomplete.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/autocomplete.py index ad75595a88..49d843f8dc 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/autocomplete.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/autocomplete.py @@ -34,8 +34,9 @@ def autocomplete_model_entity(ctx: Context, incomplete: str) -> list[tuple[str, # Suppress logging during autocomplete logging.getLogger().setLevel(logging.CRITICAL) - client = state.get_client() - models = client.inference.gateway.openai.v1.models.list(workspace=workspace) + from nemo_platform_plugin.inference_gateway.client import InferenceGatewayClient + + models = state.typed_client(InferenceGatewayClient).list_openai_models(workspace=workspace).data() if models.data: results: list[tuple[str, str]] = [] diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/code_generator.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/code_generator.py index 41d3d58fb7..64b585162c 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/code_generator.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/code_generator.py @@ -1,193 +1,260 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Code generation utilities for the NeMo CLI.""" +"""Code generation for ``--output-format code``. + +Renders the typed-client Python equivalent of a CLI invocation: the service +client import, the client construction, the method call with the same keyword +arguments (request models rendered as constructor calls), and how to consume +the response. Optional wait/watch lifecycle blocks are appended for commands +that support ``--wait`` / ``--watch``. +""" from __future__ import annotations import json +from collections.abc import Mapping +from datetime import date, datetime +from enum import Enum from textwrap import dedent -from typing import Any +from typing import Any, Literal + +from pydantic import BaseModel, RootModel, SecretBytes, SecretStr from nemo_platform_ext.cli.core.context import CLIContext +ResultKind = Literal["entity", "list", "none", "binary"] + _INFERENCE_DEPLOYMENT_LIFECYCLE = "inference_deployment" -_PLATFORM_JOB_LIFECYCLE = "platform_job" -_LIFECYCLE_TYPES_WITH_DEADLINES = {_INFERENCE_DEPLOYMENT_LIFECYCLE} -_LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING = {_INFERENCE_DEPLOYMENT_LIFECYCLE} def handle_code_generation( - resource_path: list[str], + client_cls: type | list[str], method: str, - sdk_kwargs: dict[str, Any], - output_format: str, + kwargs: Mapping[str, Any], + output_format: str | None, context: CLIContext, + *, + result: ResultKind = "entity", wait_config: dict[str, Any] | None = None, wait_options: dict[str, Any] | None = None, watch_config: dict[str, Any] | None = None, watch_options: dict[str, Any] | None = None, ) -> bool: - """ - Check if in code generation mode and generate code if needed. + """Print generated code and return True when *output_format* is ``code``. Args: - resource_path: Path to the resource (e.g., ["models"], ["customization", "configs"]) - method: Method name (e.g., "list", "create", "retrieve") - sdk_kwargs: Arguments to pass to the SDK method - output_format: Output format - - Returns: - True if code was generated (don't execute command), False otherwise + client_cls: Typed service client class the command uses (e.g. ``SecretsClient``), + or a Stainless resource path (``["models"]``) from a generated command. + method: Client method name (e.g. ``"create_secret"``). + kwargs: Keyword arguments passed to the method. Pydantic models, enums, + dicts, and primitives are rendered as Python source. + output_format: Resolved output format. + context: CLI context, used for the base URL. + result: How the response is consumed in the snippet. """ - if output_format == "code": - base_url = context.get_base_url("http://localhost:8080") - code = generate_python_code( - resource_path=resource_path, - method=method, - args=sdk_kwargs, - base_url=base_url, + if output_format != "code": + return False + + if isinstance(client_cls, list): + # Generated commands hand over a Stainless resource path; they render + # through the legacy generator until they are replaced. + from nemo_platform_ext.cli.core.legacy_code_generator import handle_code_generation as legacy + + return legacy( + client_cls, + method, + dict(kwargs), + output_format, + context, wait_config=wait_config, wait_options=wait_options, watch_config=watch_config, watch_options=watch_options, ) - formatted_code = format_code_output(code, language="python") - print(formatted_code) - return True - return False + code = generate_python_code( + client_cls, + method, + kwargs, + base_url=context.get_base_url("http://localhost:8080"), + result=result, + wait_config=wait_config, + wait_options=wait_options, + watch_config=watch_config, + watch_options=watch_options, + ) + print(format_code_output(code, language="python")) + return True def generate_python_code( - resource_path: list[str], + client_cls: type, method: str, - args: dict[str, Any], + kwargs: Mapping[str, Any], + *, base_url: str | None = None, + result: ResultKind = "entity", wait_config: dict[str, Any] | None = None, wait_options: dict[str, Any] | None = None, watch_config: dict[str, Any] | None = None, watch_options: dict[str, Any] | None = None, ) -> str: - """ - Generate Python SDK code equivalent to a CLI command. - - Args: - resource_path: Path to the resource (e.g., ["models"], ["customization", "configs"]) - method: Method name (e.g., "list", "create", "retrieve") - args: Dictionary of arguments to pass to the method - base_url: Base URL for the client (if specified) - - Returns: - Python code string - """ - lines = [] - + """Generate the typed-client Python code equivalent to a CLI command.""" if wait_config and watch_config: raise ValueError("Only one of wait_config or watch_config may be provided") lifecycle_config = watch_config or wait_config lifecycle_options = watch_options if watch_config else wait_options + lifecycle_mode = "watch" if watch_config else "wait" if wait_config else None lifecycle_type = lifecycle_config.get("type") if lifecycle_config else None - lifecycle_mode = "watch" if watch_config else "wait" if wait_config else None + imports = _ImportCollector() + imports.add(client_cls) + rendered_args = [f"{key}={_render_value(value, imports)}" for key, value in kwargs.items() if value is not None] - if _lifecycle_uses_deadline(lifecycle_type, lifecycle_mode): + lines: list[str] = [] + if lifecycle_type == _INFERENCE_DEPLOYMENT_LIFECYCLE: lines.append("import time") - if _lifecycle_uses_status_error_handling(lifecycle_type, lifecycle_mode): - lines.append( - "from nemo_platform import APIConnectionError, APIStatusError, APITimeoutError, NeMoPlatform, NotFoundError" - ) - else: - lines.append("from nemo_platform import NeMoPlatform") - if lifecycle_type == _PLATFORM_JOB_LIFECYCLE: - lines.append("from nemo_platform_plugin.client.adapter import client_from_platform") - lines.append("from nemo_platform_plugin.jobs.client import JobsClient") - if lifecycle_mode == "wait": - lines.append("from nemo_platform_plugin.jobs.watch_types import JobStatusEvent, JobWatchTimeoutError") + imports.add_name("nemo_platform_plugin.client.errors", "NemoHTTPError") + imports.add_name("nemo_platform_plugin.client.errors", "NemoTransportError") + imports.add_name("nemo_platform_plugin.client.errors", "NotFoundError") + imports.add_name("nemo_platform_plugin.inference_gateway.client", "InferenceGatewayClient") + lines.extend(imports.render()) lines.append("") if base_url: - lines.append(f"client = NeMoPlatform(base_url={_format_python_literal(base_url)})") + lines.append(f"client = {client_cls.__name__}(base_url={_format_python_literal(base_url)})") else: - lines.append("client = NeMoPlatform()") + lines.append(f"client = {client_cls.__name__}.from_config()") lines.append("") - if lifecycle_type == _PLATFORM_JOB_LIFECYCLE: - lines.append("jobs_client = client_from_platform(client, JobsClient)") - lines.append("") - resource_chain = "client." + ".".join(resource_path) - _append_method_call(lines, resource_chain, method, _format_method_args(args)) + _append_method_call(lines, "client", method, rendered_args) + lines.extend(_render_result(result)) if lifecycle_config: - lines.extend( - [ - "", - _render_lifecycle_code( - resource_path, - args, - lifecycle_config, - lifecycle_options or {}, - mode=lifecycle_mode, - ), - ] - ) - - if lifecycle_type != _PLATFORM_JOB_LIFECYCLE: lines.append("") - lines.append("print(response)") + lines.append( + _render_lifecycle_code( + kwargs, + lifecycle_config, + lifecycle_options or {}, + mode=lifecycle_mode, + ) + ) return "\n".join(lines) -def _format_method_args(args: dict[str, Any]) -> list[str]: - return [f"{key}={_format_python_literal(value)}" for key, value in args.items() if value is not None] +def format_code_output(code: str, language: str = "python") -> str: + """Syntax-highlight generated code when writing to a terminal; plain text otherwise.""" + from rich.console import Console + from rich.syntax import Syntax + + from nemo_platform_ext.cli.core.api import is_tty + + if not is_tty(): + return code + + console = Console() + syntax = Syntax( + code, + language, + line_numbers=False, + background_color="black", + padding=(1, 2), + ) + + with console.capture() as capture: + console.print(syntax) + + return capture.get() -def _append_method_call(lines: list[str], resource_chain: str, method: str, formatted_args: list[str]) -> None: +class _ImportCollector: + """Collects ``from module import Name`` lines for types used in the snippet.""" + + def __init__(self) -> None: + self._names: dict[str, set[str]] = {} + + def add(self, cls: type) -> None: + self.add_name(cls.__module__, cls.__name__) + + def add_name(self, module: str, name: str) -> None: + self._names.setdefault(module, set()).add(name) + + def render(self) -> list[str]: + return [f"from {module} import {', '.join(sorted(names))}" for module, names in sorted(self._names.items())] + + +def _render_value(value: Any, imports: _ImportCollector) -> str: + """Render *value* as Python source, registering imports for models and enums. + + Secret fields are rendered as a masked placeholder so generated code never + embeds credentials. + """ + if isinstance(value, (SecretStr, SecretBytes)): + return _format_python_literal("***") + if isinstance(value, RootModel): + # The payload is the root value, whatever its shape (model, dict, list, scalar). + imports.add(type(value)) + return f"{type(value).__name__}({_render_value(value.root, imports)})" + if isinstance(value, BaseModel): + imports.add(type(value)) + fields = value.model_dump(exclude_unset=True) + rendered = ", ".join( + f"{name}={_render_value(getattr(value, name, field_value), imports)}" + for name, field_value in fields.items() + ) + return f"{type(value).__name__}({rendered})" + if isinstance(value, Enum): + imports.add(type(value)) + return f"{type(value).__name__}.{value.name}" + if isinstance(value, dict): + items = ", ".join(f"{_format_python_literal(k)}: {_render_value(v, imports)}" for k, v in value.items()) + return "{" + items + "}" + if isinstance(value, (list, tuple)): + items = ", ".join(_render_value(v, imports) for v in value) + return f"[{items}]" + return _format_python_literal(value) + + +def _render_result(result: ResultKind) -> list[str]: + if result == "entity": + return ["", "print(response.data())"] + if result == "list": + return ["", "for item in response.page().items:", " print(item)"] + if result == "binary": + return ["", "with response.stream() as chunks:", " for chunk in chunks:", " ..."] + return [] + + +def _append_method_call(lines: list[str], target: str, method: str, formatted_args: list[str]) -> None: if not formatted_args: - lines.append(f"response = {resource_chain}.{method}()") + lines.append(f"response = {target}.{method}()") return if len(formatted_args) <= 3 and all(len(arg) <= 40 for arg in formatted_args): - lines.append(f"response = {resource_chain}.{method}({', '.join(formatted_args)})") + lines.append(f"response = {target}.{method}({', '.join(formatted_args)})") return - lines.append(f"response = {resource_chain}.{method}(") + lines.append(f"response = {target}.{method}(") for i, arg in enumerate(formatted_args): comma = "," if i < len(formatted_args) - 1 else "" lines.append(f" {arg}{comma}") lines.append(")") -def _format_keyword_args(args: dict[str, Any], keys: list[str]) -> str: - formatted_args = [] - for key in keys: - value = args.get(key) - if value is None: - continue - formatted_args.append(f"{key}={_format_python_literal(value)}") - return ", " + ", ".join(formatted_args) if formatted_args else "" - - def _format_python_literal(value: Any) -> str: if isinstance(value, str): return json.dumps(value) + if isinstance(value, (datetime, date)): + # Pydantic accepts ISO-8601 strings for date/datetime fields, and the + # string round-trips into a model without a datetime import. + return json.dumps(value.isoformat()) return repr(value) -def _lifecycle_uses_deadline(lifecycle_type: object, mode: str | None) -> bool: - return lifecycle_type in _LIFECYCLE_TYPES_WITH_DEADLINES and not ( - lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" - ) - - -def _lifecycle_uses_status_error_handling(lifecycle_type: object, mode: str | None) -> bool: - return lifecycle_type in _LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING and not ( - lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" - ) - - def _require_timeout(timeout: Any, lifecycle_type: object, mode: str | None) -> Any: if timeout is not None: return timeout @@ -195,103 +262,50 @@ def _require_timeout(timeout: Any, lifecycle_type: object, mode: str | None) -> raise ValueError(f"{mode_label}{lifecycle_type!r} lifecycle code generation requires timeout") -def _require_resource_label(lifecycle_config: dict[str, Any], lifecycle_type: object, mode: str | None) -> str: - mode_label = f"{mode} " if mode else "" - try: - resource_label = lifecycle_config["resource_label"] - except KeyError as exc: - raise ValueError( - f"{mode_label}{lifecycle_type!r} lifecycle code generation requires a non-empty resource_label" - ) from exc - if not isinstance(resource_label, str) or not resource_label.strip(): - raise ValueError( - f"{mode_label}{lifecycle_type!r} lifecycle code generation requires a non-empty resource_label" - ) - return resource_label - - def _render_lifecycle_code( - resource_path: list[str], - args: dict[str, Any], + args: Mapping[str, Any], lifecycle_config: dict[str, Any], lifecycle_options: dict[str, Any], *, mode: str | None, ) -> str: lifecycle_type = lifecycle_config.get("type") - timeout = lifecycle_options.get("timeout") + if lifecycle_type != _INFERENCE_DEPLOYMENT_LIFECYCLE: + raise ValueError(f"Unsupported lifecycle config type: {lifecycle_type!r}") + + timeout = _require_timeout(lifecycle_options.get("timeout"), lifecycle_type, mode) poll_interval = lifecycle_options.get("poll_interval", 3) - resource_chain = "client." + ".".join(resource_path) - status_kwargs = _format_keyword_args(args, ["workspace"]) - resource_name = 'getattr(response, "name", None)' - if args.get("name") is not None: - resource_name = f"{resource_name} or {_format_python_literal(args['name'])}" + resource_name = 'getattr(response.data(), "name", None)' + body = args.get("body") + body_name = getattr(body, "name", None) if body is not None else None + if body_name is not None: + resource_name = f"{resource_name} or {_format_python_literal(body_name)}" + workspace_literal = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" + flag = "--watch" if mode == "watch" else "--wait" prelude = dedent( f""" resource_name = {resource_name} if not resource_name: - raise RuntimeError("Unable to determine created resource name for --wait") + raise RuntimeError("Unable to determine created resource name for {flag}") + deadline = time.monotonic() + {timeout} """ ).strip() - if mode == "watch": - prelude = prelude.replace("--wait", "--watch") - - if lifecycle_type == _INFERENCE_DEPLOYMENT_LIFECYCLE: - timeout = _require_timeout(timeout, lifecycle_type, mode) - prelude = "\n".join([prelude, f"deadline = time.monotonic() + {timeout}"]) - workspace_literal = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" - return "\n\n".join( - [ - prelude, - _render_inference_deployment_wait_code( - resource_chain, - status_kwargs, - workspace_literal, - poll_interval, - ), - ] - ) - - if lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch": - return "\n\n".join( - [ - prelude, - _render_platform_job_watch_code(args, timeout, poll_interval), - ] - ) - if lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "wait": - timeout = _require_timeout(timeout, lifecycle_type, mode) - resource_label = _require_resource_label(lifecycle_config, lifecycle_type, mode) - return "\n\n".join( - [ - prelude, - _render_platform_job_wait_code(args, timeout, poll_interval, resource_label), - ] - ) - - raise ValueError(f"Unsupported lifecycle config type: {lifecycle_type!r}") + return "\n\n".join([prelude, _render_inference_deployment_wait_code(workspace_literal, poll_interval)]) -def _render_inference_deployment_wait_code( - resource_chain: str, - status_kwargs: str, - workspace_literal: str, - poll_interval: int, -) -> str: +def _render_inference_deployment_wait_code(workspace_literal: str, poll_interval: int) -> str: return dedent( f""" while True: - deployment = {resource_chain}.retrieve(resource_name{status_kwargs}) - history = getattr(deployment, "status_history", None) - status = history[-1].status if history else deployment.status + deployment = client.get_deployment(name=resource_name, workspace={workspace_literal}).data() + history = deployment.status_history + status = (history[-1].status if history else deployment.status).value if status == "READY": - response = deployment provider_name = resource_name provider_workspace = {workspace_literal} - model_provider_id = getattr(deployment, "model_provider_id", None) - if model_provider_id: - provider_workspace, _, provider_name = model_provider_id.partition("/") + if deployment.model_provider_id: + provider_workspace, _, provider_name = deployment.model_provider_id.partition("/") if not provider_workspace or not provider_name: provider_workspace = {workspace_literal} provider_name = resource_name @@ -303,18 +317,14 @@ def _render_inference_deployment_wait_code( raise TimeoutError(f"Timed out waiting for deployment {{resource_name!r}} to become READY") time.sleep(min({poll_interval}, remaining)) + gateway = InferenceGatewayClient.from_client(client) while True: try: - if provider_workspace is None: - client.inference.gateway.provider.ready(provider_name) - else: - client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace) + gateway.provider_ready(name=provider_name, workspace=provider_workspace) break - except NotFoundError: - pass - except (APIConnectionError, APITimeoutError): + except (NotFoundError, NemoTransportError): pass - except APIStatusError as exc: + except NemoHTTPError as exc: if exc.status_code not in {{429, 502, 503, 504}}: raise remaining = deadline - time.monotonic() @@ -323,91 +333,3 @@ def _render_inference_deployment_wait_code( time.sleep(min({poll_interval}, remaining)) """ ).strip() - - -def _render_platform_job_watch_code( - args: dict[str, Any], - timeout: int | None, - poll_interval: int, -) -> str: - workspace = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" - - return dedent( - f""" - for event in jobs_client.watch_job( - resource_name, - workspace={workspace}, - timeout={timeout}, - poll_interval={poll_interval}, - ): - print(event) - """ - ).strip() - - -def _render_platform_job_wait_code( - args: dict[str, Any], - timeout: int, - poll_interval: int, - resource_label: str, -) -> str: - workspace = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" - - return dedent( - f""" - resource_label = {_format_python_literal(resource_label)} - try: - for event in jobs_client.watch_job( - resource_name, - workspace={workspace}, - timeout={timeout}, - poll_interval={poll_interval}, - include_logs=False, - ): - if not isinstance(event, JobStatusEvent): - continue - if not event.terminal: - continue - if event.successful: - break - raise RuntimeError( - f"{{resource_label.title()}} {{resource_name!r}} ended with status {{event.status!r}}" - ) - except JobWatchTimeoutError as exc: - raise TimeoutError(f"Timed out waiting for {{resource_label}} {{resource_name!r}} to complete") from exc - """ - ).strip() - - -def format_code_output(code: str, language: str = "python") -> str: - """ - Format code output with syntax highlighting. - - Args: - code: Code string to format - language: Programming language - - Returns: - Formatted code string - """ - from rich.console import Console - from rich.syntax import Syntax - - from nemo_platform_ext.cli.core.api import is_tty - - if not is_tty(): - return code - - console = Console() - syntax = Syntax( - code, - language, - line_numbers=False, - background_color="black", - padding=(1, 2), - ) - - with console.capture() as capture: - console.print(syntax) - - return capture.get() diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/context.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/context.py index d636645a9d..b19b0484ac 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/context.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/context.py @@ -17,10 +17,14 @@ if typing.TYPE_CHECKING: from nemo_platform import AsyncNeMoPlatform, NeMoPlatform + from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_ext.config.config import ConfigParams, Context from nemo_platform_ext.quickstart import QuickstartConfig +TypedClientT = typing.TypeVar("TypedClientT", bound="NemoClient") +AsyncTypedClientT = typing.TypeVar("AsyncTypedClientT", bound="AsyncNemoClient") + logger = logging.getLogger("nemo_platform_ext.cli") @@ -31,6 +35,10 @@ class CLIContext: Holds CLI overrides (via ConfigParams) and lazy-loads SDK config. Priority resolution is handled by SDK Config: CLI > env_var > config file > default. + + Hand-written commands derive service clients from the platform client with + :meth:`typed_client` (for example ``state.typed_client(SecretsClient)``); the + typed client shares the CLI's auth and transport. """ # CLI overrides passed to SDK Config.load() @@ -157,6 +165,25 @@ def get_async_client(self, timeout: float = 60.0) -> AsyncNeMoPlatform: ) return self._async_client + def typed_client(self, client_cls: type[TypedClientT], timeout: float = 60.0) -> TypedClientT: + """Return a service client of *client_cls* sharing the CLI client's transport and auth.""" + from nemo_platform_plugin.client.adapter import client_from_platform + + return client_from_platform(self.get_client(timeout=timeout), client_cls) + + def async_typed_client(self, client_cls: type[AsyncTypedClientT], timeout: float = 60.0) -> AsyncTypedClientT: + """Async twin of :meth:`typed_client`.""" + from nemo_platform_plugin.client.adapter import client_from_platform + + return client_from_platform(self.get_async_client(timeout=timeout), client_cls) + + def get_workspace(self) -> str | None: + """Return the configured default workspace, if any.""" + try: + return self.get_sdk_context().workspace + except Exception: + return None + def get_output_format( self, override: OutputFormat | None = None, diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/errors.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/errors.py index 285dcdc3a0..53d7336e38 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/errors.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/errors.py @@ -12,6 +12,7 @@ import click import httpx import typer +from pydantic import ValidationError REMOTE_ERROR_EXIT_CODE = 3 @@ -33,6 +34,16 @@ def __init__( super().__init__(f"Missing required fields: {missing_str}") +class UnknownInputFieldsError(Exception): + """Raised when ``--input-data`` / ``--input-file`` carries keys the request does not define.""" + + def __init__(self, unknown_fields: list[str], command_name: str, known_fields: list[str]): + self.unknown_fields = unknown_fields + self.command_name = command_name + self.known_fields = known_fields + super().__init__(f"Unknown fields: {', '.join(unknown_fields)}") + + class InvalidSearchPatternError(Exception): """Raised when --filter is given JSON that fails to parse or an otherwise unusable value.""" @@ -58,6 +69,19 @@ def _build_list_cmd(ctx: click.Context | None, prog: str) -> str | None: return " ".join([prog, *parts, "list"]) +_MISSING_WORKSPACE_MARKERS = ( + "Missing workspace argument", + "Missing path parameter 'workspace'", + "workspace must be provided", +) + + +def _is_missing_workspace_error(error: ValueError) -> bool: + """Return whether *error* is a client reporting an unresolved workspace.""" + message = str(error) + return any(marker in message for marker in _MISSING_WORKSPACE_MARKERS) + + def _format_api_error(error: object) -> str: """Extract a clean error message from an API error.""" if hasattr(error, "body") and error.body is not None: @@ -69,6 +93,9 @@ def _format_api_error(error: object) -> str: message = body.get("message") if message: return str(message) + detail = getattr(error, "detail", None) + if isinstance(detail, str) and detail: + return detail message = getattr(error, "message", None) if isinstance(message, str) and message: return message @@ -309,11 +336,18 @@ def handle_exception(error: Exception, ctx: click.Context | None = None) -> None console.print(f"[bold red]API response error:[/] {_format_api_error(error)}") _print_api_request_context(console, error) raise typer.Exit(code=REMOTE_ERROR_EXIT_CODE) - elif isinstance(error, APIError): + elif isinstance(error, (APIError, plugin_errors.NemoClientError)): console.print(f"[bold red]API error:[/] {_format_api_error(error)}") _print_api_request_context(console, error) raise typer.Exit(code=REMOTE_ERROR_EXIT_CODE) - elif isinstance(error, ValueError) and "Missing workspace argument" in str(error): + elif isinstance(error, ValidationError): + console.print("[bold red]Invalid input:[/]") + for detail in error.errors(): + location = ".".join(str(part) for part in detail.get("loc", ())) or "input" + console.print(f" [yellow]{location}[/] {detail.get('msg', 'invalid value')}") + console.print("[yellow]Hint:[/] Check your input values. Run with [cyan]--help[/] to see required options.") + raise typer.Exit(code=2) + elif isinstance(error, ValueError) and _is_missing_workspace_error(error): console.print("[bold red]Missing workspace:[/] No workspace configured for this command.") console.print( f"[yellow]Hint:[/] Run [cyan]{prog} config set --workspace [/] or use the [cyan]--workspace[/] option." @@ -338,6 +372,13 @@ def handle_exception(error: Exception, ctx: click.Context | None = None) -> None console.print() console.print("[yellow]Hint:[/] Provide via CLI flags or [cyan]--input-file[/]/[cyan]--input-data[/].") raise typer.Exit(code=2) + elif isinstance(error, UnknownInputFieldsError): + console.print(f"[bold bright_green]Usage:[/] {prog} [GLOBAL OPTIONS] {error.command_name} [OPTIONS]") + console.print(f"Try [cyan]{prog} {error.command_name} --help[/] for help.") + console.print() + console.print(f"[bold red]Error:[/] Unknown input fields: {', '.join(error.unknown_fields)}") + console.print(f"[yellow]Hint:[/] Accepted fields: {', '.join(error.known_fields)}.") + raise typer.Exit(code=2) elif isinstance(error, InvalidSearchPatternError): if error.parse_error: console.print(f"[bold red]Error:[/] Invalid filter JSON: {error.parse_error}") diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/formatters.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/formatters.py index 20fdcfaced..99c7f32ec6 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/formatters.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/formatters.py @@ -16,6 +16,7 @@ import click import yaml +from nemo_platform_plugin.client.response import NemoResponse from rich.console import Console from rich.syntax import Syntax from rich.table import Table @@ -113,6 +114,13 @@ def _iter_items_from_response(data: Any) -> Iterator[Any]: return +def unwrap_response(data: Any) -> Any: + """Return the parsed body of a typed entity response, or *data* unchanged.""" + if isinstance(data, NemoResponse): + return data.data() + return data + + def _to_dict_items(items: list[Any]) -> list[dict[str, Any]]: """Convert items to dicts (Pydantic models, dicts, or fallback to {value: str}).""" result = [] @@ -634,6 +642,7 @@ def format_output( """ from nemo_platform_ext.cli.core.table_config import resolve_and_validate_columns, validate_output_columns + data = unwrap_response(data) timestamp_format = timestamp_format or "iso" if stream: diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/help_formatter.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/help_formatter.py index cc2d13ece0..495ad3f688 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/help_formatter.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/help_formatter.py @@ -817,6 +817,9 @@ def create_typer_app(**kwargs) -> Typer: kwargs.setdefault("cls", NmpGroup) kwargs.setdefault("no_args_is_help", True) + # Shell completion is owned by the root ``nemo`` app; command groups (including + # plugin-hosted roots mounted by the lazy loader) must not advertise it again. + kwargs.setdefault("add_completion", False) kwargs["context_settings"] = _context_settings_with_help(kwargs.get("context_settings")) return typer.Typer(**kwargs) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py new file mode 100644 index 0000000000..4c0e778be1 --- /dev/null +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py @@ -0,0 +1,385 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Code generation for ``--output-format code`` on the generated (Stainless SDK) commands. + +Kept only while ``cli/commands/api/`` exists; hand-written commands use +:mod:`nemo_platform_ext.cli.core.code_generator`, which dispatches here when +handed a resource path instead of a typed client class. +""" + +from __future__ import annotations + +import json +from textwrap import dedent +from typing import Any + +from nemo_platform_ext.cli.core.code_generator import format_code_output +from nemo_platform_ext.cli.core.context import CLIContext + +_INFERENCE_DEPLOYMENT_LIFECYCLE = "inference_deployment" +_PLATFORM_JOB_LIFECYCLE = "platform_job" +_LIFECYCLE_TYPES_WITH_DEADLINES = {_INFERENCE_DEPLOYMENT_LIFECYCLE} +_LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING = {_INFERENCE_DEPLOYMENT_LIFECYCLE} + + +def handle_code_generation( + resource_path: list[str], + method: str, + sdk_kwargs: dict[str, Any], + output_format: str, + context: CLIContext, + wait_config: dict[str, Any] | None = None, + wait_options: dict[str, Any] | None = None, + watch_config: dict[str, Any] | None = None, + watch_options: dict[str, Any] | None = None, +) -> bool: + """ + Check if in code generation mode and generate code if needed. + + Args: + resource_path: Path to the resource (e.g., ["models"], ["customization", "configs"]) + method: Method name (e.g., "list", "create", "retrieve") + sdk_kwargs: Arguments to pass to the SDK method + output_format: Output format + + Returns: + True if code was generated (don't execute command), False otherwise + """ + if output_format == "code": + base_url = context.get_base_url("http://localhost:8080") + code = generate_python_code( + resource_path=resource_path, + method=method, + args=sdk_kwargs, + base_url=base_url, + wait_config=wait_config, + wait_options=wait_options, + watch_config=watch_config, + watch_options=watch_options, + ) + formatted_code = format_code_output(code, language="python") + print(formatted_code) + return True + + return False + + +def generate_python_code( + resource_path: list[str], + method: str, + args: dict[str, Any], + base_url: str | None = None, + wait_config: dict[str, Any] | None = None, + wait_options: dict[str, Any] | None = None, + watch_config: dict[str, Any] | None = None, + watch_options: dict[str, Any] | None = None, +) -> str: + """ + Generate Python SDK code equivalent to a CLI command. + + Args: + resource_path: Path to the resource (e.g., ["models"], ["customization", "configs"]) + method: Method name (e.g., "list", "create", "retrieve") + args: Dictionary of arguments to pass to the method + base_url: Base URL for the client (if specified) + + Returns: + Python code string + """ + lines = [] + + if wait_config and watch_config: + raise ValueError("Only one of wait_config or watch_config may be provided") + + lifecycle_config = watch_config or wait_config + lifecycle_options = watch_options if watch_config else wait_options + lifecycle_type = lifecycle_config.get("type") if lifecycle_config else None + + lifecycle_mode = "watch" if watch_config else "wait" if wait_config else None + + if _lifecycle_uses_deadline(lifecycle_type, lifecycle_mode): + lines.append("import time") + if _lifecycle_uses_status_error_handling(lifecycle_type, lifecycle_mode): + lines.append( + "from nemo_platform import APIConnectionError, APIStatusError, APITimeoutError, NeMoPlatform, NotFoundError" + ) + else: + lines.append("from nemo_platform import NeMoPlatform") + if lifecycle_type == _PLATFORM_JOB_LIFECYCLE: + lines.append("from nemo_platform_plugin.client.adapter import client_from_platform") + lines.append("from nemo_platform_plugin.jobs.client import JobsClient") + if lifecycle_mode == "wait": + lines.append("from nemo_platform_plugin.jobs.watch_types import JobStatusEvent, JobWatchTimeoutError") + lines.append("") + + if base_url: + lines.append(f"client = NeMoPlatform(base_url={_format_python_literal(base_url)})") + else: + lines.append("client = NeMoPlatform()") + lines.append("") + if lifecycle_type == _PLATFORM_JOB_LIFECYCLE: + lines.append("jobs_client = client_from_platform(client, JobsClient)") + lines.append("") + + resource_chain = "client." + ".".join(resource_path) + _append_method_call(lines, resource_chain, method, _format_method_args(args)) + + if lifecycle_config: + lines.extend( + [ + "", + _render_lifecycle_code( + resource_path, + args, + lifecycle_config, + lifecycle_options or {}, + mode=lifecycle_mode, + ), + ] + ) + + if lifecycle_type != _PLATFORM_JOB_LIFECYCLE: + lines.append("") + lines.append("print(response)") + + return "\n".join(lines) + + +def _format_method_args(args: dict[str, Any]) -> list[str]: + return [f"{key}={_format_python_literal(value)}" for key, value in args.items() if value is not None] + + +def _append_method_call(lines: list[str], resource_chain: str, method: str, formatted_args: list[str]) -> None: + if not formatted_args: + lines.append(f"response = {resource_chain}.{method}()") + return + + if len(formatted_args) <= 3 and all(len(arg) <= 40 for arg in formatted_args): + lines.append(f"response = {resource_chain}.{method}({', '.join(formatted_args)})") + return + + lines.append(f"response = {resource_chain}.{method}(") + for i, arg in enumerate(formatted_args): + comma = "," if i < len(formatted_args) - 1 else "" + lines.append(f" {arg}{comma}") + lines.append(")") + + +def _format_keyword_args(args: dict[str, Any], keys: list[str]) -> str: + formatted_args = [] + for key in keys: + value = args.get(key) + if value is None: + continue + formatted_args.append(f"{key}={_format_python_literal(value)}") + return ", " + ", ".join(formatted_args) if formatted_args else "" + + +def _format_python_literal(value: Any) -> str: + if isinstance(value, str): + return json.dumps(value) + return repr(value) + + +def _lifecycle_uses_deadline(lifecycle_type: object, mode: str | None) -> bool: + return lifecycle_type in _LIFECYCLE_TYPES_WITH_DEADLINES and not ( + lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" + ) + + +def _lifecycle_uses_status_error_handling(lifecycle_type: object, mode: str | None) -> bool: + return lifecycle_type in _LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING and not ( + lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" + ) + + +def _require_timeout(timeout: Any, lifecycle_type: object, mode: str | None) -> Any: + if timeout is not None: + return timeout + mode_label = f"{mode} " if mode else "" + raise ValueError(f"{mode_label}{lifecycle_type!r} lifecycle code generation requires timeout") + + +def _require_resource_label(lifecycle_config: dict[str, Any], lifecycle_type: object, mode: str | None) -> str: + mode_label = f"{mode} " if mode else "" + try: + resource_label = lifecycle_config["resource_label"] + except KeyError as exc: + raise ValueError( + f"{mode_label}{lifecycle_type!r} lifecycle code generation requires a non-empty resource_label" + ) from exc + if not isinstance(resource_label, str) or not resource_label.strip(): + raise ValueError( + f"{mode_label}{lifecycle_type!r} lifecycle code generation requires a non-empty resource_label" + ) + return resource_label + + +def _render_lifecycle_code( + resource_path: list[str], + args: dict[str, Any], + lifecycle_config: dict[str, Any], + lifecycle_options: dict[str, Any], + *, + mode: str | None, +) -> str: + lifecycle_type = lifecycle_config.get("type") + timeout = lifecycle_options.get("timeout") + poll_interval = lifecycle_options.get("poll_interval", 3) + resource_chain = "client." + ".".join(resource_path) + status_kwargs = _format_keyword_args(args, ["workspace"]) + resource_name = 'getattr(response, "name", None)' + if args.get("name") is not None: + resource_name = f"{resource_name} or {_format_python_literal(args['name'])}" + + prelude = dedent( + f""" + resource_name = {resource_name} + if not resource_name: + raise RuntimeError("Unable to determine created resource name for --wait") + """ + ).strip() + if mode == "watch": + prelude = prelude.replace("--wait", "--watch") + + if lifecycle_type == _INFERENCE_DEPLOYMENT_LIFECYCLE: + timeout = _require_timeout(timeout, lifecycle_type, mode) + prelude = "\n".join([prelude, f"deadline = time.monotonic() + {timeout}"]) + workspace_literal = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" + return "\n\n".join( + [ + prelude, + _render_inference_deployment_wait_code( + resource_chain, + status_kwargs, + workspace_literal, + poll_interval, + ), + ] + ) + + if lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch": + return "\n\n".join( + [ + prelude, + _render_platform_job_watch_code(args, timeout, poll_interval), + ] + ) + if lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "wait": + timeout = _require_timeout(timeout, lifecycle_type, mode) + resource_label = _require_resource_label(lifecycle_config, lifecycle_type, mode) + return "\n\n".join( + [ + prelude, + _render_platform_job_wait_code(args, timeout, poll_interval, resource_label), + ] + ) + + raise ValueError(f"Unsupported lifecycle config type: {lifecycle_type!r}") + + +def _render_inference_deployment_wait_code( + resource_chain: str, + status_kwargs: str, + workspace_literal: str, + poll_interval: int, +) -> str: + return dedent( + f""" + while True: + deployment = {resource_chain}.retrieve(resource_name{status_kwargs}) + history = getattr(deployment, "status_history", None) + status = history[-1].status if history else deployment.status + if status == "READY": + response = deployment + provider_name = resource_name + provider_workspace = {workspace_literal} + model_provider_id = getattr(deployment, "model_provider_id", None) + if model_provider_id: + provider_workspace, _, provider_name = model_provider_id.partition("/") + if not provider_workspace or not provider_name: + provider_workspace = {workspace_literal} + provider_name = resource_name + break + if status in {{"ERROR", "LOST"}}: + raise RuntimeError(f"Deployment {{resource_name!r}} ended with status {{status!r}}") + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"Timed out waiting for deployment {{resource_name!r}} to become READY") + time.sleep(min({poll_interval}, remaining)) + + while True: + try: + if provider_workspace is None: + client.inference.gateway.provider.ready(provider_name) + else: + client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace) + break + except NotFoundError: + pass + except (APIConnectionError, APITimeoutError): + pass + except APIStatusError as exc: + if exc.status_code not in {{429, 502, 503, 504}}: + raise + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"Timed out waiting for gateway readiness for {{resource_name!r}}") + time.sleep(min({poll_interval}, remaining)) + """ + ).strip() + + +def _render_platform_job_watch_code( + args: dict[str, Any], + timeout: int | None, + poll_interval: int, +) -> str: + workspace = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" + + return dedent( + f""" + for event in jobs_client.watch_job( + resource_name, + workspace={workspace}, + timeout={timeout}, + poll_interval={poll_interval}, + ): + print(event) + """ + ).strip() + + +def _render_platform_job_wait_code( + args: dict[str, Any], + timeout: int, + poll_interval: int, + resource_label: str, +) -> str: + workspace = _format_python_literal(args["workspace"]) if args.get("workspace") is not None else "None" + + return dedent( + f""" + resource_label = {_format_python_literal(resource_label)} + try: + for event in jobs_client.watch_job( + resource_name, + workspace={workspace}, + timeout={timeout}, + poll_interval={poll_interval}, + include_logs=False, + ): + if not isinstance(event, JobStatusEvent): + continue + if not event.terminal: + continue + if event.successful: + break + raise RuntimeError( + f"{{resource_label.title()}} {{resource_name!r}} ended with status {{event.status!r}}" + ) + except JobWatchTimeoutError as exc: + raise TimeoutError(f"Timed out waiting for {{resource_label}} {{resource_name!r}} to complete") from exc + """ + ).strip() diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py index bf3398699e..a1b4663f8f 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py @@ -1,16 +1,26 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Pagination utilities for the NeMo CLI.""" +"""Pagination helpers for CLI list commands. + +Typed clients return a :class:`~nemo_platform_plugin.client.response.NemoPaginatedResponse` +for list endpoints. The helpers here turn that into the page objects the +formatters render: a single page keeps the server's pagination metadata, and +``--all-pages`` collects every page into one synthetic response with the same +shape. +""" from __future__ import annotations import logging import typing +from collections.abc import Callable, Iterator from contextlib import contextmanager from enum import Enum -from typing import Any, Callable, Iterator +from types import SimpleNamespace, TracebackType +from typing import Any, cast +from nemo_platform_plugin.client.response import NemoPaginatedResponse from rich.progress import Progress, SpinnerColumn, TextColumn from nemo_platform_ext.cli.core.help_formatter import add_warning @@ -49,9 +59,17 @@ class AllPagesResponse: Used for page-number based pagination which has total_pages and total_results. """ - def __init__(self, data: list[Any], total_items: int, total_pages: int, page_size: int | None = None): + def __init__( + self, + data: list[Any], + total_items: int, + total_pages: int, + page_size: int | None = None, + envelope: dict[str, Any] | None = None, + ): self.data = data - self.sort = None + self.envelope = dict(envelope or {}) + self.sort = self.envelope.get("sort") # Create pagination info for all items self.pagination = type( @@ -80,7 +98,7 @@ def model_dump(self, mode: str = "json") -> dict[str, Any]: return { "data": serialized_data, - "sort": self.sort, + **self.envelope, "pagination": { "page": self.pagination.page, "page_size": self.pagination.page_size, @@ -124,6 +142,163 @@ def model_dump(self, mode: str = "json") -> dict[str, Any]: } +def _model_dump_item(item: Any, *, mode: str) -> Any: + if hasattr(item, "model_dump"): + return item.model_dump(mode=mode) + if isinstance(item, list): + return [_model_dump_item(child, mode=mode) for child in item] + if isinstance(item, dict): + return {key: _model_dump_item(value, mode=mode) for key, value in item.items()} + return item + + +def _envelope_fields(response: NemoPaginatedResponse[Any, Any]) -> dict[str, Any]: + """Return the first page's non-item envelope fields (``sort``, ``filter``, ``grouped_by``, ...) in wire order.""" + http_response = getattr(response, "http_response", None) + if http_response is None: + return {} + try: + body = http_response.json() + except ValueError: + return {} + if not isinstance(body, dict): + return {} + return {key: value for key, value in body.items() if key not in {"data", "pagination"}} + + +class OffsetPageResponse: + """One page of an offset-paginated list, keeping the server's envelope and pagination block.""" + + def __init__(self, items: list[Any], metadata: dict[str, Any], envelope: dict[str, Any] | None = None) -> None: + self.data = items + self.envelope = dict(envelope or {}) + self.sort = self.envelope.get("sort") + self.pagination = SimpleNamespace(**metadata) + + def model_dump(self, mode: str = "json") -> dict[str, Any]: + return { + "data": [_model_dump_item(item, mode=mode) for item in self.data], + **self.envelope, + "pagination": vars(self.pagination), + } + + +class CursorPageResponse: + """One page of a cursor-paginated list, keeping ``total`` and the page cursors.""" + + def __init__(self, items: list[Any], metadata: dict[str, Any]) -> None: + self.data = items + self.total = metadata.get("total", len(items)) + self.next_page = metadata.get("next_page") + self.prev_page = metadata.get("prev_page") + + def model_dump(self, mode: str = "json") -> dict[str, Any]: + return { + "data": [_model_dump_item(item, mode=mode) for item in self.data], + "total": self.total, + "next_page": self.next_page, + "prev_page": self.prev_page, + } + + +def collect_offset_pages( + response: NemoPaginatedResponse[Any, Any], + *, + all_pages: bool, + show_progress: bool = True, +) -> OffsetPageResponse | AllPagesResponse: + """Return the first page, or every page merged when *all_pages* is set.""" + envelope = _envelope_fields(response) + if not all_pages: + page = response.page() + return OffsetPageResponse(list(page.items), dict(page.metadata), envelope) + + items: list[Any] = [] + total_results = 0 + total_pages = 0 + page_size: int | None = None + with _progress(show_progress) as (progress, task): + for page in response.pages(): + items.extend(page.items) + metadata = dict(page.metadata) + total_results = int(metadata.get("total_results") or total_results) + total_pages = int(metadata.get("total_pages") or total_pages) + page_size = cast(int | None, metadata.get("page_size") or page_size) + current = int(metadata.get("page") or 0) + progress.update( + task, + total=total_pages or None, + completed=current, + description=f"Fetching pages... (page {current}/{total_pages})", + ) + + return AllPagesResponse( + data=items, + total_items=total_results or len(items), + total_pages=total_pages or 1, + page_size=page_size, + envelope=envelope, + ) + + +def collect_cursor_pages( + response: NemoPaginatedResponse[Any, Any], + *, + all_pages: bool, + limit: int | None = None, + show_progress: bool = True, +) -> CursorPageResponse | AllCursorPagesResponse: + """Return the first page, or every page merged when *all_pages* is set.""" + if not all_pages: + page = response.page() + return CursorPageResponse(list(page.items), dict(page.metadata)) + + items: list[Any] = [] + with _progress(show_progress) as (progress, task): + for page_num, page in enumerate(response.pages(), start=1): + items.extend(page.items) + progress.update(task, completed=page_num, description=f"Fetching pages... (page {page_num})") + return AllCursorPagesResponse(data=items, limit=limit) + + +def collect_pages( + response: NemoPaginatedResponse[Any, Any], + *, + all_pages: bool, + pagination_type: PaginationType = PaginationType.PAGE_NUMBER, + limit: int | None = None, + show_progress: bool = True, +) -> OffsetPageResponse | CursorPageResponse | AllPagesResponse | AllCursorPagesResponse: + """Dispatch to the offset or cursor collector based on *pagination_type*.""" + if pagination_type == PaginationType.CURSOR: + return collect_cursor_pages(response, all_pages=all_pages, limit=limit, show_progress=show_progress) + return collect_offset_pages(response, all_pages=all_pages, show_progress=show_progress) + + +class _progress: + """Context manager yielding ``(progress, task)`` for page-fetch feedback.""" + + def __init__(self, show_progress: bool) -> None: + self._progress = Progress( + SpinnerColumn(), + TextColumn("[progress.description]{task.description}"), + transient=True, + disable=not show_progress, + ) + + def __enter__(self) -> tuple[Progress, Any]: + progress = self._progress.__enter__() + return progress, progress.add_task("Fetching pages...", total=None) + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + self._progress.__exit__(exc_type, exc_value, traceback) + + def _fetch_all_pages_page_number( list_method: Callable[..., SyncDefaultPagination[Any]], progress: Progress, diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py index 5469f23eb2..fab0049cd5 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py @@ -7,10 +7,14 @@ import os import sys -from typing import Any +from collections.abc import Collection, Mapping +from typing import Any, TypeVar import yaml from click import UsageError +from pydantic import BaseModel, RootModel + +RequestModelT = TypeVar("RequestModelT", bound=BaseModel) def is_stdin_available() -> bool: @@ -165,6 +169,57 @@ def validate_required_fields( raise MissingRequiredFieldsError(missing, command_name, field_help) +def _field_model(model_cls: type[BaseModel]) -> type[BaseModel] | None: + """Return the model whose fields name the accepted payload keys, or None when keys are unconstrained.""" + if issubclass(model_cls, RootModel): + root_type = model_cls.model_fields["root"].annotation + if not (isinstance(root_type, type) and issubclass(root_type, BaseModel)): + return None + model_cls = root_type + if model_cls.model_config.get("extra") == "allow": + return None + return model_cls + + +def _accepted_field_names(model_cls: type[BaseModel]) -> list[str]: + names: list[str] = [] + for name, field in model_cls.model_fields.items(): + names.append(name) + if isinstance(field.alias, str): + names.append(field.alias) + if isinstance(field.validation_alias, str): + names.append(field.validation_alias) + return names + + +def build_request_body( + model_cls: type[RequestModelT], + payload: Mapping[str, Any], + *, + exclude: Collection[str] = (), + command_name: str = "this command", +) -> RequestModelT: + """Validate the user-supplied *payload* into a request body of *model_cls*. + + Keys in *exclude* are CLI-only (workspace, exist_ok, ...) and dropped. Any + remaining key the model does not define is rejected instead of silently + ignored, so a typo in ``--input-data`` cannot degrade into a no-op. Models + that declare ``extra="allow"`` (pass-through request shapes) and root models + wrapping unstructured payloads accept any key. Only the + keys present in *payload* are marked set on the returned model, so + ``exclude_unset`` serialization sends exactly what the user provided. + """ + from nemo_platform_ext.cli.core.errors import UnknownInputFieldsError + + body = {key: value for key, value in payload.items() if key not in exclude} + field_model = _field_model(model_cls) + if field_model is not None: + unknown = sorted(set(body) - set(_accepted_field_names(field_model))) + if unknown: + raise UnknownInputFieldsError(unknown, command_name, sorted(field_model.model_fields)) + return model_cls.model_validate(body) + + def read_payload(field_name: str, field_value: str) -> Any: """ Read payload for a specific field, supporting JSON and YAML strings. diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/waiters.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/waiters.py index 552cf7c34f..c8c426648d 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/waiters.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/waiters.py @@ -21,11 +21,14 @@ import time from collections.abc import Iterator from datetime import datetime, timezone +from enum import Enum from typing import Any -from nemo_platform import APIConnectionError, APIStatusError, APITimeoutError, NotFoundError +from nemo_platform_plugin.client.adapter import PlatformClient, client_from_platform +from nemo_platform_plugin.client.errors import NemoHTTPError, NemoTransportError, NotFoundError from nemo_platform_plugin.client.response import NemoPaginatedResponse, NemoResponse from nemo_platform_plugin.client.types import CursorPagination +from nemo_platform_plugin.inference_gateway.client import InferenceGatewayClient from nemo_platform_plugin.jobs.client import JobsWatchClient from nemo_platform_plugin.jobs.schemas import PlatformJobLog, PlatformJobStatusResponse from nemo_platform_plugin.jobs.types import JobLogsQueryParams @@ -35,6 +38,7 @@ JobWatchEvent, JobWatchTimeoutError, ) +from nemo_platform_plugin.models.client import ModelsClient from rich.console import Console from rich.live import Live from rich.text import Text @@ -137,6 +141,9 @@ def _seconds_since_creation(entry_timestamp: datetime | str | None, created_at: def _status_text(status: Any) -> str: + """Normalize a status to its wire string; typed models carry ``str`` enums.""" + if isinstance(status, Enum): + status = status.value return str(status or "") @@ -263,7 +270,7 @@ def __init__(self, *, start_time: float, timeout: int, poll_interval: int) -> No self.poll_interval = poll_interval def snapshot(self) -> tuple[str, int]: - return datetime.now().strftime("%H:%M:%S"), int(time.time() - self.start_time) + return datetime.now().strftime("%H:%M:%S"), int(time.monotonic() - self.start_time) def __rich__(self) -> Text: polling_time, wait_elapsed = self.snapshot() @@ -273,7 +280,7 @@ def __rich__(self) -> Text: def _sleep_until_next_poll(start_time: float, timeout: float, poll_interval: int) -> bool: if poll_interval <= 0: raise ValueError(f"_sleep_until_next_poll poll_interval must be greater than 0, got {poll_interval}") - remaining = timeout - (time.time() - start_time) + remaining = timeout - (time.monotonic() - start_time) if remaining <= 0: return False _pause(min(poll_interval, remaining)) @@ -296,7 +303,7 @@ def _print_transient_wait_error(live: Live, resource_label: str, error: Exceptio def wait_for_inference_deployment( - client: Any, + client: PlatformClient, name: str, *, workspace: str | None = None, @@ -307,10 +314,10 @@ def wait_for_inference_deployment( verbose: bool = True, ) -> bool: """Wait for an inference deployment to reach the requested status.""" - if workspace is None: - workspace = client._get_workspace_path_param() + models_client = client_from_platform(client, ModelsClient) + workspace = models_client.require_workspace(workspace) - start_time = time.time() + start_time = time.monotonic() last_history_len = 0 last_status = "" last_message = "" @@ -319,14 +326,14 @@ def wait_for_inference_deployment( console.print(f"[bold]Waiting for deployment '{name}' to reach status: {status}[/bold]\n") with Live(console=console, refresh_per_second=4, transient=True) as live: - while time.time() - start_time < timeout: - wait_elapsed = int(time.time() - start_time) + while time.monotonic() - start_time < timeout: + wait_elapsed = int(time.monotonic() - start_time) polling_time = datetime.now().strftime("%H:%M:%S") if verbose: live.update(_make_live_display(polling_time, timeout, poll_interval, wait_elapsed)) try: - deployment = client.inference.deployments.retrieve(name, workspace=workspace) + deployment = models_client.get_deployment(name=name, workspace=workspace).data() history = getattr(deployment, "status_history", None) created_at = getattr(deployment, "created_at", None) if history and len(history) > 0: @@ -363,7 +370,7 @@ def wait_for_inference_deployment( if verbose: console.print(f"\n[green]✓ Deployment reached {status} status![/green]") if status == "READY" and check_gateway: - remaining_timeout = timeout - (time.time() - start_time) + remaining_timeout = timeout - (time.monotonic() - start_time) if remaining_timeout <= 0: console.print("\n[red]✗ Timeout before gateway readiness check could complete[/red]") return False @@ -391,10 +398,10 @@ def wait_for_inference_deployment( live.stop() console.print("\n[red]✗ Deployment not found[/red]") return False - except (APIConnectionError, APITimeoutError) as exc: + except NemoTransportError as exc: if verbose: _print_transient_wait_error(live, "deployment status", exc) - except APIStatusError as exc: + except NemoHTTPError as exc: if exc.status_code not in _TRANSIENT_GATEWAY_STATUS_CODES: raise if verbose: @@ -403,7 +410,7 @@ def wait_for_inference_deployment( if not _sleep_until_next_poll(start_time, timeout, poll_interval): break - wait_elapsed = int(time.time() - start_time) + wait_elapsed = int(time.monotonic() - start_time) detail = f"Last status: {last_status}" if last_message: detail += f" - {last_message}" @@ -421,7 +428,8 @@ def wait_for_platform_job( poll_interval: int = 3, ) -> bool: """Wait for a platform job resource to complete.""" - start_time = time.time() + start_time = time.monotonic() + wall_start_time = time.time() last_status = "" jobs = _WatchedJobsClient(jobs_client) @@ -460,7 +468,10 @@ def wait_for_platform_job( if event.terminal: live.stop() _emit_job_run_event( - jobs.last_status, resource_label=resource_label, status=current_status, start_time=start_time + jobs.last_status, + resource_label=resource_label, + status=current_status, + start_time=wall_start_time, ) if event.successful: console.print(f"\n[green]✓ {resource_label.title()} completed![/green]") @@ -474,14 +485,14 @@ def wait_for_platform_job( except JobWatchTimeoutError: pass - wait_elapsed = int(time.time() - start_time) + wait_elapsed = int(time.monotonic() - start_time) detail = f"Last status: {last_status}" if last_status else "No status returned" console.print(f"\n[red]✗ Timeout after {wait_elapsed}s. {detail}[/red]") return False def wait_for_gateway( - client: Any, + client: PlatformClient, provider_name: str, workspace: str, timeout: float = 60, @@ -489,8 +500,9 @@ def wait_for_gateway( verbose: bool = True, ) -> bool: """Wait for the inference gateway to be able to route to a provider.""" - start_time = time.time() + start_time = time.monotonic() start_timestamp = datetime.now().strftime("%H:%M:%S") + gateway_client = client_from_platform(client, InferenceGatewayClient) if verbose: console.print(f"[bold]Waiting for gateway to be ready for provider '{provider_name}'[/bold]\n") @@ -504,23 +516,23 @@ def _make_gateway_display(polling_time: str, elapsed: int, status: str) -> Text: return text with Live(console=console, refresh_per_second=4, transient=True) as live: - while time.time() - start_time < timeout: - elapsed = int(time.time() - start_time) + while time.monotonic() - start_time < timeout: + elapsed = int(time.monotonic() - start_time) polling_time = datetime.now().strftime("%H:%M:%S") if verbose: live.update(_make_gateway_display(polling_time, elapsed, "Checking gateway...")) try: - client.inference.gateway.provider.ready(provider_name, workspace=workspace) + gateway_client.provider_ready(name=provider_name, workspace=workspace) live.stop() if verbose: console.print(f" [{polling_time}] ({elapsed}s) [green]Gateway is ready![/green]") return True except NotFoundError: pass - except (APIConnectionError, APITimeoutError): + except NemoTransportError: pass - except APIStatusError as exc: + except NemoHTTPError as exc: if exc.status_code in _TRANSIENT_GATEWAY_STATUS_CODES: pass else: @@ -531,6 +543,6 @@ def _make_gateway_display(polling_time: str, elapsed: int, status: str) -> Text: if not _sleep_until_next_poll(start_time, timeout, poll_interval): break - elapsed = int(time.time() - start_time) + elapsed = int(time.monotonic() - start_time) console.print(f"\n[red]✗ Gateway timeout after {elapsed}s[/red]") return False diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/telemetry/emit.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/telemetry/emit.py index bfd94a5dc0..2e77446640 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/telemetry/emit.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/telemetry/emit.py @@ -61,9 +61,10 @@ def telemetry_opted_in() -> bool: def _client_version() -> str: try: - import nemo_platform + from nemo_platform_ext.cli.version import UNKNOWN_VERSION, client_version - return nemo_platform.__version__ + resolved = client_version() + return "undefined" if resolved == UNKNOWN_VERSION else resolved except Exception: logger.debug("Could not resolve client version for telemetry", exc_info=True) return "undefined" diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/version.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/version.py new file mode 100644 index 0000000000..a55572e886 --- /dev/null +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/version.py @@ -0,0 +1,24 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Version reporting for the ``nemo`` CLI.""" + +from __future__ import annotations + +from importlib.metadata import PackageNotFoundError, version + +# Distributions that can carry the CLI, most specific first. The wrapper +# distribution and the SDK distribution share one release version. +_DISTRIBUTIONS = ("nemo-platform", "nemo-platform-sdk", "nemo-platform-ext") + +UNKNOWN_VERSION = "unknown" + + +def client_version() -> str: + """Return the installed release version of the CLI, or ``"unknown"``.""" + for distribution in _DISTRIBUTIONS: + try: + return version(distribution) + except PackageNotFoundError: + continue + return UNKNOWN_VERSION diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/client/bootstrap.py b/packages/nemo_platform_ext/src/nemo_platform_ext/client/bootstrap.py new file mode 100644 index 0000000000..40f323f288 --- /dev/null +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/client/bootstrap.py @@ -0,0 +1,780 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Config and auth bootstrap shared by NeMo Platform clients. + +This module bridges the nmp CLI config file (~/.config/nmp/config.yaml) and the +platform HTTP clients. When the active user is an OAuthUser, the bootstrap +wires up **transparent token refresh** so that every HTTP request made through a +client automatically carries a valid Bearer token — no manual token management +needed. It has no dependency on any generated SDK; ``factory.py`` layers the +``NeMoPlatform`` constructors on top, and :func:`build_nemo_client` / +:func:`build_async_nemo_client` build typed ``NemoClient`` instances directly. + +High-level flow +=============== + + build_nemo_client() / build_client_init_kwargs() + │ + ├─ resolve_bootstrap() + │ ├─ _resolve_client_context() → reads nmp config, resolves context + │ ├─ _discover_oidc_client_settings() → GET /apis/auth/discovery + │ └─ _get_or_create_provider() → reuses or creates OIDCTokenProvider + │ + └─ returns a resolved bootstrap with: + • base_url, workspace, default_headers + • a token provider consulted before every request + +Token refresh is **lazy** (on-demand), not proactive. No background threads are +started. The token is checked on each HTTP request; if it expires within +_TOKEN_REFRESH_MARGIN_SECONDS (60 s), a synchronous refresh_token grant is +performed inline before the request proceeds. + +Concurrency safety +================== + +Three layers prevent race conditions when multiple SDK instances or processes +share the same config file: + +1. **In-process thread lock** — ``OIDCTokenProvider._lock`` serializes concurrent + ``get_access_token()`` calls from different threads. +2. **Provider cache** — ``_TOKEN_PROVIDER_CACHE`` ensures multiple clients created + in the same process with the same (config_path, context) share a single + ``OIDCTokenProvider`` instance, avoiding redundant refreshes. +3. **Cross-process file lock** — An ``fcntl.flock``-based lock file next to the + config prevents multiple processes from racing on refresh and clobbering each + other's tokens. + +See also: ``architecture/docs/auth/sdk-cli-oauth.md`` for a full design doc. +""" + +import asyncio +import logging +import os +import threading +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Literal, Mapping, Protocol + +import httpx +from nemo_platform_plugin.client.auth import TokenProviderAuth +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.types import RetryPolicy + +from nemo_platform_ext.auth.helpers import NMPOIDCConfig, build_effective_scope, discover_nmp_config +from nemo_platform_ext.auth.token_provider import ( + OIDCTokenProvider, + TokenSet, +) +from nemo_platform_ext.auth.workload_exchange import WorkloadTokenExchangeProvider +from nemo_platform_ext.client.tls import client_verify_from_env + +logger = logging.getLogger(__name__) + +# Refresh the access token when fewer than 60 s remain before expiry. +# This gives enough headroom for the refresh HTTP round-trip to complete +# before the token actually expires. +_TOKEN_REFRESH_MARGIN_SECONDS = 60 + +# Guards _TOKEN_PROVIDER_CACHE; acquired only during dict lookup/insert (fast). +_TOKEN_PROVIDER_CACHE_LOCK = threading.Lock() + +# --------------------------------------------------------------------------- +# Data classes +# --------------------------------------------------------------------------- + + +class AccessTokenProvider(Protocol): + def get_access_token(self) -> str: ... + + async def get_access_token_async(self) -> str: ... + + +@dataclass(frozen=True) +class ResolvedBootstrap: + """Result of resolving config + OIDC discovery into client parameters.""" + + base_url: str + workspace: str | None + default_headers: dict[str, str] + token_provider: AccessTokenProvider | None # None for non-OAuth users + client_verify: str | Literal[True] + certificate_authority: str | None = None + + +@dataclass(frozen=True) +class _ProviderCacheKey: + """Composite key for the provider cache. + + Two clients share the same provider iff they read from the same config + file, same context, OIDC settings (endpoint, client, scope), and CA match. + """ + + config_path: Path + context_name: str + token_endpoint: str + client_id: str + refresh_scope: str | None + certificate_authority: str | None + + +# Process-wide cache: (config_path, context, OIDC settings, CA) → shared OIDCTokenProvider. +# This avoids redundant refresh-token grants when the user creates multiple +# NeMoPlatform() instances pointing at the same context. +_TOKEN_PROVIDER_CACHE: dict[_ProviderCacheKey, OIDCTokenProvider] = {} + + +# --------------------------------------------------------------------------- +# OIDC discovery +# --------------------------------------------------------------------------- + +# Fallback returned when NMP configdiscovery fails. +_OIDC_DISCOVERY_FALLBACK = NMPOIDCConfig( + auth_enabled=False, + client_id="", + token_endpoint="", + default_scopes="openid profile email", + scope_prefix=None, +) + + +def _discover_oidc_client_settings(base_url: str, certificate_authority: str | None = None) -> NMPOIDCConfig: + """Fetch OIDC config from the NeMo Platform cluster's discovery endpoint. + + Returns a safe fallback (auth_enabled=False) if the cluster is + unreachable or doesn't have OIDC configured. This lets non-OIDC + clusters work without errors during client construction. + """ + try: + return discover_nmp_config(base_url, certificate_authority=certificate_authority) + except Exception: + logger.debug("Could not discover OIDC settings from %s", base_url, exc_info=True) + return _OIDC_DISCOVERY_FALLBACK + + +def _workload_identity_token_file_from_env() -> Path | None: + """Return the configured workload identity token file, if workload bootstrap is active.""" + token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR) + return Path(token_file) if token_file else None + + +def _create_workload_exchange_provider( + base_url: str, + subject_token_file: Path, + *, + certificate_authority: str | None = None, +) -> WorkloadTokenExchangeProvider: + """Create a workload identity token exchange provider from NeMo auth discovery metadata.""" + oidc_config = _discover_oidc_client_settings(base_url, certificate_authority=certificate_authority) + if not oidc_config.workload_token_exchange_enabled: + raise RuntimeError( + f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} is set but workload token exchange is not enabled by auth discovery" + ) + + token_endpoint = oidc_config.workload_token_endpoint or oidc_config.token_endpoint or "" + client_id = oidc_config.workload_client_id or oidc_config.client_id or "" + if not token_endpoint: + raise RuntimeError( + "Workload token exchange is enabled but auth discovery did not return workload_token_endpoint or token_endpoint" + ) + if not client_id: + raise RuntimeError( + "Workload token exchange is enabled but auth discovery did not return workload_client_id or client_id" + ) + + return WorkloadTokenExchangeProvider( + token_endpoint=token_endpoint, + client_id=client_id, + subject_token_file=subject_token_file, + audience=oidc_config.workload_audience, + scope=oidc_config.workload_scope, + certificate_authority=certificate_authority, + refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, + ) + + +class _LazyWorkloadTokenExchangeProvider: + """Create the workload exchange provider on the first token request.""" + + def __init__( + self, + *, + base_url: str, + subject_token_file: Path, + certificate_authority: str | None = None, + ) -> None: + self._base_url = base_url + self._subject_token_file = subject_token_file + self._certificate_authority = certificate_authority + self._provider: WorkloadTokenExchangeProvider | None = None + self._lock = threading.Lock() + + def _get_provider(self) -> WorkloadTokenExchangeProvider: + provider = self._provider + if provider is not None: + return provider + with self._lock: + provider = self._provider + if provider is None: + provider = _create_workload_exchange_provider( + self._base_url, + self._subject_token_file, + certificate_authority=self._certificate_authority, + ) + self._provider = provider + return provider + + def get_cached_access_token(self) -> str | None: + provider = self._provider + if provider is None: + return None + tokens = provider.tokens + if not tokens.access_token or tokens.is_expired(provider.refresh_margin_seconds): + return None + return tokens.access_token + + def get_access_token(self) -> str: + return self._get_provider().get_access_token() + + async def get_access_token_async(self) -> str: + provider = self._provider + if provider is None: + provider = await asyncio.to_thread(self._get_provider) + return await provider.get_access_token_async() + + +# --------------------------------------------------------------------------- +# httpx event hooks — the core of transparent token injection +# --------------------------------------------------------------------------- + + +def _make_auth_event_hook(provider: AccessTokenProvider): + """Create a **sync** httpx request event hook that injects the Bearer token. + + Called before every SDK HTTP request. ``provider.get_access_token()`` + returns the cached token if still valid, or performs an inline + refresh_token grant if the token is expired/about-to-expire. + """ + + def inject_auth(request: httpx.Request) -> None: + token = provider.get_access_token() + request.headers["Authorization"] = f"Bearer {token}" + + return inject_auth + + +def _make_async_auth_event_hook(provider: AccessTokenProvider): + """Create an **async** httpx request event hook for AsyncNeMoPlatform. + + The actual refresh still runs in a worker thread (via + ``provider.get_access_token_async``) so it doesn't block the event loop. + """ + + async def inject_auth(request: httpx.Request) -> None: + token = await provider.get_access_token_async() + request.headers["Authorization"] = f"Bearer {token}" + + return inject_auth + + +def _headers_with_seeded_auth(headers: Mapping[str, str], provider: AccessTokenProvider) -> dict[str, str]: + seeded_headers = dict(headers) + if isinstance(provider, _LazyWorkloadTokenExchangeProvider): + token = provider.get_cached_access_token() + else: + token = provider.get_access_token() + if token: + seeded_headers["Authorization"] = f"Bearer {token}" + return seeded_headers + + +# --------------------------------------------------------------------------- +# Callbacks wired into OIDCTokenProvider for config-file integration +# --------------------------------------------------------------------------- + + +def _make_config_persister(context_name: str, config_path: Path | None = None): + """Create an ``on_tokens_refreshed`` callback that writes new tokens to the + nmp config file. + + After a successful refresh, the provider calls this so that the CLI and + other SDK processes pick up the rotated tokens without re-authenticating. + """ + from nemo_platform_ext.config.config import Config, ConfigParams + + def persist(tokens: TokenSet) -> None: + params: ConfigParams = {"access_token": tokens.access_token} + if tokens.refresh_token: + params["refresh_token"] = tokens.refresh_token + Config.write(params, context_name=context_name, config_path=config_path) + logger.debug("Persisted refreshed tokens to nmp config (context=%s)", context_name) + + return persist + + +def _make_config_token_loader(context_name: str, config_path: Path): + """Create a ``load_tokens`` callback that re-reads tokens from the config file. + + This is the recovery mechanism for the ``invalid_grant`` scenario: + when another process already rotated the refresh token, our local + copy is stale. The provider calls this to reload whatever that + other process wrote, then retries the refresh with the fresh token. + """ + from nemo_platform_ext.config.config import Config, ConfigParams + from nemo_platform_ext.config.models import OAuthUser + + def load_tokens() -> TokenSet | None: + overrides: ConfigParams = {"current_context": context_name} + try: + config = Config.load(config_path=config_path, overrides=overrides) + resolved = config.resolve() + except Exception: + logger.debug("Failed to reload tokens from nmp config (context=%s)", context_name, exc_info=True) + return None + + if not isinstance(resolved.user, OAuthUser): + return None + + return TokenSet.from_access_token( + resolved.user.token.get_secret_value(), + resolved.user.refresh_token.get_secret_value() if resolved.user.refresh_token else None, + ) + + return load_tokens + + +# --------------------------------------------------------------------------- +# Cross-process file lock for refresh serialization +# --------------------------------------------------------------------------- + + +def _build_refresh_lock_path(config_path: Path, context_name: str) -> Path: + """Derive the lock-file path from the config path and context name. + + Example: ``~/.config/nmp/config.yaml.default.oauth-refresh.lock`` + """ + safe_context = context_name.replace(os.sep, "_") + if os.altsep: + safe_context = safe_context.replace(os.altsep, "_") + return config_path.with_name(f"{config_path.name}.{safe_context}.oauth-refresh.lock") + + +def _make_refresh_lock(config_path: Path, context_name: str): + """Create a ``refresh_lock`` context-manager factory for OIDCTokenProvider. + + Uses ``fcntl.flock`` (POSIX) to serialize refresh transactions across + processes. On platforms without fcntl (Windows), falls back to a no-op + so the provider still works — just without cross-process serialization. + """ + lock_path = _build_refresh_lock_path(config_path, context_name) + + @contextmanager + def refresh_lock(): + try: + import fcntl + except ImportError: + # Windows: no fcntl — skip cross-process locking. + yield + return + + lock_path.parent.mkdir(parents=True, exist_ok=True) + with open(lock_path, "a+", encoding="utf-8") as lock_file: + fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX) + try: + yield + finally: + fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN) + + return refresh_lock + + +def _normalize_config_path(path: Path) -> Path: + """Canonicalize the config path so that cache keys match regardless of + how the caller spelled the path (``~/…`` vs absolute).""" + return path.expanduser().resolve(strict=False) + + +# --------------------------------------------------------------------------- +# Provider cache +# --------------------------------------------------------------------------- + + +def _get_or_create_provider( + key: _ProviderCacheKey, + create_provider: Callable[[], OIDCTokenProvider], +) -> OIDCTokenProvider: + """Return the cached provider for *key*, or create and cache a new one. + + If a cached provider already exists, we call ``reload_tokens()`` so it + picks up any tokens that another process may have written to the config + file since the provider was last used. + """ + with _TOKEN_PROVIDER_CACHE_LOCK: + provider = _TOKEN_PROVIDER_CACHE.get(key) + if provider is None: + provider = create_provider() + _TOKEN_PROVIDER_CACHE[key] = provider + return provider + + # Provider already existed — reload tokens from disk in case another + # process refreshed them since we last checked. + provider.reload_tokens() + return provider + + +# --------------------------------------------------------------------------- +# Config resolution +# --------------------------------------------------------------------------- + + +def _resolve_client_context( + *, + config_path: Path | None, + base_url: str | httpx.URL | None, + context_name: str | None, + access_token: str | None, +) -> tuple[Any, bool, Path]: + """Load the nmp config file, apply overrides, and resolve the active context. + + Returns ``(resolved_context, config_exists, config_path)`` so the caller + knows whether to enable config-backed features (provider caching, + token persistence, file locking). + """ + from nemo_platform_ext.config.config import Config, ConfigParams + + resolved_config_path = config_path or Config.get_default_config_path() + config_exists = resolved_config_path.exists() + if config_exists: + logger.info("Reading nmp config from %s", resolved_config_path) + + # Constructor args override whatever is in the config file. + overrides: ConfigParams | None = None + if context_name is not None or access_token is not None or base_url is not None: + overrides = {} + if base_url is not None: + overrides["base_url"] = str(base_url) + if context_name is not None: + overrides["current_context"] = context_name + if access_token is not None: + overrides["access_token"] = access_token + + config = Config.load(config_path=config_path, overrides=overrides) + if context_name is not None: + available_contexts = [ctx.name for ctx in config.get_config_file().contexts] + if context_name not in available_contexts: + available = ", ".join(available_contexts) if available_contexts else "(none)" + raise ValueError(f"Context '{context_name}' not found. Available contexts: {available}") + + return config.resolve(), config_exists, resolved_config_path + + +# --------------------------------------------------------------------------- +# Bootstrap: config + OIDC discovery → resolved client params +# --------------------------------------------------------------------------- + + +def resolve_bootstrap( + *, + config_path: Path | None, + base_url: str | httpx.URL | None, + context_name: str | None, + access_token: str | None, + extra_headers: Mapping[str, str] | None, +) -> ResolvedBootstrap: + """Resolve the full client bootstrap: config, OIDC discovery, token provider. + + For **non-OAuth users** (no-auth): returns a bootstrap with + ``token_provider=None`` and static headers (e.g. ``Authorization: Bearer ``). + + For **OAuth users**: discovers the cluster's OIDC settings, then either + reuses a cached OIDCTokenProvider or creates a new one. The provider is + wired with: + + - ``load_tokens``: re-reads tokens from the config file (for ``invalid_grant`` recovery) + - ``refresh_lock``: cross-process fcntl lock to serialize refreshes + - ``on_tokens_refreshed``: writes rotated tokens back to the config file + """ + from nemo_platform_ext.config.models import OAuthUser + + resolved, config_exists, resolved_config_path = _resolve_client_context( + config_path=config_path, + base_url=base_url, + context_name=context_name, + access_token=access_token, + ) + + base_url = str(resolved.cluster.base_url) + certificate_authority = resolved.cluster.certificate_authority + client_verify = client_verify_from_env(certificate_authority) + headers: dict[str, str] = dict(extra_headers) if extra_headers else {} + + workload_identity_token_file = _workload_identity_token_file_from_env() + if workload_identity_token_file is not None and access_token is None and not os.environ.get("NMP_ACCESS_TOKEN"): + provider = _LazyWorkloadTokenExchangeProvider( + base_url=base_url, + subject_token_file=workload_identity_token_file, + certificate_authority=certificate_authority, + ) + return ResolvedBootstrap(base_url, resolved.workspace, headers, provider, client_verify, certificate_authority) + + # --- Non-OAuth path (no auth) --- + if not isinstance(resolved.user, OAuthUser): + user_config = resolved.user.get_client_config() if resolved.user else {} + user_headers = user_config.get("default_headers", {}) + if isinstance(user_headers, dict): + headers.update(user_headers) + return ResolvedBootstrap(base_url, resolved.workspace, headers, None, client_verify, certificate_authority) + + # --- OAuth path: set up transparent token refresh --- + try: + oidc_config = discover_nmp_config(base_url, certificate_authority=certificate_authority) + if not oidc_config.auth_enabled and access_token is None: + # Discovery succeeded and confirmed the cluster has no OIDC. The + # stored OAuthUser token can't be refreshed here (no endpoint), so + # fall back to no-auth. This guard only fires on a successful + # discovery response — not on discovery failures, where the stored + # token may still be valid and should be used as-is. + return ResolvedBootstrap(base_url, resolved.workspace, headers, None, client_verify, certificate_authority) + except Exception: + logger.debug("Could not discover OIDC settings from %s", base_url, exc_info=True) + oidc_config = _OIDC_DISCOVERY_FALLBACK + + tokens = TokenSet.from_access_token( + resolved.user.token.get_secret_value(), + resolved.user.refresh_token.get_secret_value() if resolved.user.refresh_token else None, + ) + + token_endpoint = oidc_config.token_endpoint or "" + client_id = oidc_config.client_id or "" + refresh_scope = build_effective_scope(oidc_config.default_scopes, oidc_config.scope_prefix) + + # Only share the provider (and enable persistence/locking) when reading + # from an actual config file. If the caller passed an explicit + # access_token, they own the token lifecycle — don't cache or persist. + share_provider = config_exists and access_token is None + + if share_provider: + normalized_config_path = _normalize_config_path(resolved_config_path) + provider_key = _ProviderCacheKey( + config_path=normalized_config_path, + context_name=resolved.context_name, + token_endpoint=token_endpoint, + client_id=client_id, + refresh_scope=refresh_scope, + certificate_authority=certificate_authority, + ) + on_refreshed = _make_config_persister(resolved.context_name, resolved_config_path) + load_tokens = _make_config_token_loader(resolved.context_name, resolved_config_path) + refresh_lock = _make_refresh_lock(resolved_config_path, resolved.context_name) + + provider = _get_or_create_provider( + provider_key, + lambda: OIDCTokenProvider( + token_endpoint=token_endpoint, + client_id=client_id, + tokens=tokens, + refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, + refresh_scope=refresh_scope, + certificate_authority=certificate_authority, + load_tokens=load_tokens, + refresh_lock=refresh_lock, + on_tokens_refreshed=on_refreshed, + ), + ) + else: + # Ephemeral provider: no persistence, no file locking, no caching. + provider = OIDCTokenProvider( + token_endpoint=token_endpoint, + client_id=client_id, + tokens=tokens, + refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, + refresh_scope=refresh_scope, + certificate_authority=certificate_authority, + ) + + return ResolvedBootstrap(base_url, resolved.workspace, headers, provider, client_verify, certificate_authority) + + +# --------------------------------------------------------------------------- +# Typed client construction +# --------------------------------------------------------------------------- + +# Matches the retry behaviour the generated SDK applied by default, so the CLI +# keeps the same resilience against transient gateway errors. +DEFAULT_RETRY_POLICY = RetryPolicy( + max_retries=2, + retryable_status_codes=(408, 409, 429), + retry_all_server_errors=True, + respect_retry_decision_headers=True, + respect_retry_after_headers=True, +) + + +# Connect phase cap for CLI clients. A blackholed endpoint fails in seconds +# rather than waiting out the full read timeout on every attempt. +DEFAULT_CONNECT_TIMEOUT = 5.0 +DEFAULT_REQUEST_TIMEOUT = 60.0 + + +def resolve_timeout(timeout: float | httpx.Timeout | None) -> httpx.Timeout: + """Normalize a caller timeout into phase-aware httpx form. + + A bare number (or ``None`` for the 60 s default) is the read/write/pool + budget; the connect phase is capped at :data:`DEFAULT_CONNECT_TIMEOUT` + separately. An explicit :class:`httpx.Timeout` is used as given. + """ + if isinstance(timeout, httpx.Timeout): + return timeout + return httpx.Timeout(DEFAULT_REQUEST_TIMEOUT if timeout is None else timeout, connect=DEFAULT_CONNECT_TIMEOUT) + + +def _client_headers(bootstrap: ResolvedBootstrap) -> dict[str, str]: + """Return the default headers for a client built from *bootstrap*. + + A static bearer token (API key user) already lives in ``default_headers``. + Token providers inject ``Authorization`` per request, so nothing is seeded + here; the transport auth handles raw ``_client`` calls too. + """ + return dict(bootstrap.default_headers) + + +def build_nemo_client( + *, + config_path: Path | None = None, + base_url: str | httpx.URL | None = None, + context_name: str | None = None, + access_token: str | None = None, + extra_headers: Mapping[str, str] | None = None, + workspace: str | None = None, + timeout: float | httpx.Timeout | None = None, + retry: RetryPolicy | None = DEFAULT_RETRY_POLICY, +) -> NemoClient: + """Build a sync :class:`NemoClient` from the nmp config and auth bootstrap. + + For OAuth and workload-identity users the resulting client refreshes and + injects the Bearer token before every request; API-key users get static + headers. *workspace* overrides the configured default when given. + """ + bootstrap = resolve_bootstrap( + config_path=config_path, + base_url=base_url, + context_name=context_name, + access_token=access_token, + extra_headers=extra_headers, + ) + resolved_timeout = resolve_timeout(timeout) + http_client = httpx.Client( + headers=_client_headers(bootstrap) or None, + timeout=resolved_timeout, + follow_redirects=True, + verify=bootstrap.client_verify, + auth=TokenProviderAuth(bootstrap.token_provider) if bootstrap.token_provider is not None else None, + ) + return NemoClient( + base_url=bootstrap.base_url, + workspace=workspace if workspace is not None else bootstrap.workspace, + auth=bootstrap.token_provider, + default_headers=_client_headers(bootstrap) or None, + timeout=resolved_timeout, + retry=retry, + http_client=http_client, + ) + + +def build_async_nemo_client( + *, + config_path: Path | None = None, + base_url: str | httpx.URL | None = None, + context_name: str | None = None, + access_token: str | None = None, + extra_headers: Mapping[str, str] | None = None, + workspace: str | None = None, + timeout: float | httpx.Timeout | None = None, + retry: RetryPolicy | None = DEFAULT_RETRY_POLICY, +) -> AsyncNemoClient: + """Async twin of :func:`build_nemo_client`.""" + bootstrap = resolve_bootstrap( + config_path=config_path, + base_url=base_url, + context_name=context_name, + access_token=access_token, + extra_headers=extra_headers, + ) + resolved_timeout = resolve_timeout(timeout) + http_client = httpx.AsyncClient( + headers=_client_headers(bootstrap) or None, + timeout=resolved_timeout, + follow_redirects=True, + verify=bootstrap.client_verify, + auth=TokenProviderAuth(bootstrap.token_provider) if bootstrap.token_provider is not None else None, + ) + return AsyncNemoClient( + base_url=bootstrap.base_url, + workspace=workspace if workspace is not None else bootstrap.workspace, + auth=bootstrap.token_provider, + default_headers=_client_headers(bootstrap) or None, + timeout=resolved_timeout, + retry=retry, + http_client=http_client, + ) + + +def build_direct_nemo_client( + *, + base_url: str, + workspace: str | None = None, + default_headers: Mapping[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + certificate_authority: str | None = None, + retry: RetryPolicy | None = DEFAULT_RETRY_POLICY, +) -> NemoClient: + """Build a sync :class:`NemoClient` without reading the nmp config. + + Direct mode: no config file is consulted and only *default_headers* are + sent. TLS verification still honours the environment override and any + saved cluster certificate authority. + """ + resolved_timeout = resolve_timeout(timeout) + http_client = httpx.Client( + headers=dict(default_headers) if default_headers else None, + timeout=resolved_timeout, + follow_redirects=True, + verify=client_verify_from_env(certificate_authority), + ) + return NemoClient( + base_url=base_url, + workspace=workspace, + default_headers=default_headers, + timeout=resolved_timeout, + retry=retry, + http_client=http_client, + ) + + +def build_direct_async_nemo_client( + *, + base_url: str, + workspace: str | None = None, + default_headers: Mapping[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + certificate_authority: str | None = None, + retry: RetryPolicy | None = DEFAULT_RETRY_POLICY, +) -> AsyncNemoClient: + """Async twin of :func:`build_direct_nemo_client`.""" + resolved_timeout = resolve_timeout(timeout) + http_client = httpx.AsyncClient( + headers=dict(default_headers) if default_headers else None, + timeout=resolved_timeout, + follow_redirects=True, + verify=client_verify_from_env(certificate_authority), + ) + return AsyncNemoClient( + base_url=base_url, + workspace=workspace, + default_headers=default_headers, + timeout=resolved_timeout, + retry=retry, + http_client=http_client, + ) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py b/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py index 8dfb393f9a..bc8696e02b 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py @@ -3,61 +3,16 @@ """Factory for creating NeMoPlatform SDK clients from nmp config. -This module bridges the nmp CLI config file (~/.config/nmp/config.yaml) and the -NeMoPlatform SDK client. When the active user is an OAuthUser, the factory -wires up **transparent token refresh** so that every HTTP request made through the -SDK automatically carries a valid Bearer token — no manual token management needed. - -High-level flow -=============== - - NeMoPlatform() (enhanced.py) - │ - ▼ - build_client_init_kwargs() - │ - ├─ _resolve_bootstrap() - │ ├─ _resolve_client_context() → reads nmp config, resolves context - │ ├─ _discover_oidc_client_settings() → GET /apis/auth/discovery - │ └─ _get_or_create_provider() → reuses or creates OIDCTokenProvider - │ - └─ returns ClientInitConfig with: - • base_url, workspace, default_headers - • httpx client with a request event hook that calls - provider.get_access_token() before every request - -Token refresh is **lazy** (on-demand), not proactive. No background threads are -started. The token is checked on each HTTP request; if it expires within -_TOKEN_REFRESH_MARGIN_SECONDS (60 s), a synchronous refresh_token grant is -performed inline before the request proceeds. - -Concurrency safety -================== - -Three layers prevent race conditions when multiple SDK instances or processes -share the same config file: - -1. **In-process thread lock** — ``OIDCTokenProvider._lock`` serializes concurrent - ``get_access_token()`` calls from different threads. -2. **Provider cache** — ``_TOKEN_PROVIDER_CACHE`` ensures multiple clients created - in the same process with the same (config_path, context) share a single - ``OIDCTokenProvider`` instance, avoiding redundant refreshes. -3. **Cross-process file lock** — An ``fcntl.flock``-based lock file next to the - config prevents multiple processes from racing on refresh and clobbering each - other's tokens. - -See also: ``architecture/docs/auth/sdk-cli-oauth.md`` for a full design doc. +Layers the generated ``NeMoPlatform`` constructors on top of the shared, +SDK-independent bootstrap in :mod:`nemo_platform_ext.client.bootstrap`, which +resolves the nmp config file, discovers OIDC settings, and builds the token +provider that keeps every request authenticated. """ -import asyncio -import logging -import os -import threading from collections.abc import Awaitable, Callable, Mapping -from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path -from typing import Any, Literal, Protocol +from typing import Literal import httpx from nemo_platform import ( @@ -68,36 +23,28 @@ Omit, not_given, ) -from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR -from nemo_platform_ext.auth.helpers import NMPOIDCConfig, build_effective_scope, discover_nmp_config -from nemo_platform_ext.auth.token_provider import ( - OIDCTokenProvider, - TokenSet, +from nemo_platform_ext.client.bootstrap import ( + _TOKEN_PROVIDER_CACHE, + _TOKEN_PROVIDER_CACHE_LOCK, + AccessTokenProvider, + ResolvedBootstrap, + _headers_with_seeded_auth, + _make_async_auth_event_hook, + _make_auth_event_hook, + resolve_bootstrap, ) -from nemo_platform_ext.auth.workload_exchange import WorkloadTokenExchangeProvider -from nemo_platform_ext.client.tls import client_verify_from_env - -logger = logging.getLogger(__name__) - -# Refresh the access token when fewer than 60 s remain before expiry. -# This gives enough headroom for the refresh HTTP round-trip to complete -# before the token actually expires. -_TOKEN_REFRESH_MARGIN_SECONDS = 60 - -# Guards _TOKEN_PROVIDER_CACHE; acquired only during dict lookup/insert (fast). -_TOKEN_PROVIDER_CACHE_LOCK = threading.Lock() - -# --------------------------------------------------------------------------- -# Data classes -# --------------------------------------------------------------------------- - - -class _AccessTokenProvider(Protocol): - def get_access_token(self) -> str: ... - - async def get_access_token_async(self) -> str: ... +__all__ = [ + "_TOKEN_PROVIDER_CACHE", + "_TOKEN_PROVIDER_CACHE_LOCK", + "AccessTokenProvider", + "ClientInitConfig", + "ResolvedBootstrap", + "build_async_client_init_kwargs", + "build_client_init_kwargs", + "create_client", +] _SyncRequestHook = Callable[[httpx.Request], None] _AsyncRequestHook = Callable[[httpx.Request], Awaitable[None]] @@ -107,7 +54,7 @@ async def get_access_token_async(self) -> str: ... @dataclass(frozen=True) class ClientInitConfig: - """Everything the SDK client constructor needs after config resolution. + """Everything a generated SDK client constructor needs after config resolution. For non-OAuth users this just carries base_url/workspace/headers. For OAuth users it also includes a custom httpx client with an event @@ -121,512 +68,28 @@ class ClientInitConfig: client_verify: str | Literal[True] = True -@dataclass(frozen=True) -class _ResolvedBootstrap: - """Intermediate result after resolving config + OIDC discovery.""" - - base_url: str - workspace: str | None - default_headers: dict[str, str | Omit] - token_provider: _AccessTokenProvider | None # None for non-OAuth users - client_verify: str | Literal[True] - certificate_authority: str | None = None - - -@dataclass(frozen=True) -class _ProviderCacheKey: - """Composite key for the provider cache. - - Two clients share the same provider iff they read from the same config - file, same context, OIDC settings (endpoint, client, scope), and CA match. - """ - - config_path: Path - context_name: str - token_endpoint: str - client_id: str - refresh_scope: str | None - certificate_authority: str | None - - -# Process-wide cache: (config_path, context, OIDC settings, CA) → shared OIDCTokenProvider. -# This avoids redundant refresh-token grants when the user creates multiple -# NeMoPlatform() instances pointing at the same context. -_TOKEN_PROVIDER_CACHE: dict[_ProviderCacheKey, OIDCTokenProvider] = {} - - -# --------------------------------------------------------------------------- -# OIDC discovery -# --------------------------------------------------------------------------- - -# Fallback returned when NMP configdiscovery fails. -_OIDC_DISCOVERY_FALLBACK = NMPOIDCConfig( - auth_enabled=False, - client_id="", - token_endpoint="", - default_scopes="openid profile email", - scope_prefix=None, -) - - -def _discover_oidc_client_settings(base_url: str, certificate_authority: str | None = None) -> NMPOIDCConfig: - """Fetch OIDC config from the NeMo Platform cluster's discovery endpoint. - - Returns a safe fallback (auth_enabled=False) if the cluster is - unreachable or doesn't have OIDC configured. This lets non-OIDC - clusters work without errors during client construction. - """ - try: - return discover_nmp_config(base_url, certificate_authority=certificate_authority) - except Exception: - logger.debug("Could not discover OIDC settings from %s", base_url, exc_info=True) - return _OIDC_DISCOVERY_FALLBACK - - -def _workload_identity_token_file_from_env() -> Path | None: - """Return the configured workload identity token file, if workload bootstrap is active.""" - token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR) - return Path(token_file) if token_file else None - - -def _create_workload_exchange_provider( - base_url: str, - subject_token_file: Path, - *, - certificate_authority: str | None = None, -) -> WorkloadTokenExchangeProvider: - """Create a workload identity token exchange provider from NeMo auth discovery metadata.""" - oidc_config = _discover_oidc_client_settings(base_url, certificate_authority=certificate_authority) - if not oidc_config.workload_token_exchange_enabled: - raise RuntimeError( - f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} is set but workload token exchange is not enabled by auth discovery" - ) - - token_endpoint = oidc_config.workload_token_endpoint or oidc_config.token_endpoint or "" - client_id = oidc_config.workload_client_id or oidc_config.client_id or "" - if not token_endpoint: - raise RuntimeError( - "Workload token exchange is enabled but auth discovery did not return workload_token_endpoint or token_endpoint" - ) - if not client_id: - raise RuntimeError( - "Workload token exchange is enabled but auth discovery did not return workload_client_id or client_id" - ) - - return WorkloadTokenExchangeProvider( - token_endpoint=token_endpoint, - client_id=client_id, - subject_token_file=subject_token_file, - audience=oidc_config.workload_audience, - scope=oidc_config.workload_scope, - certificate_authority=certificate_authority, - refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, - ) - - -class _LazyWorkloadTokenExchangeProvider: - """Create the workload exchange provider on the first token request.""" - - def __init__( - self, - *, - base_url: str, - subject_token_file: Path, - certificate_authority: str | None = None, - ) -> None: - self._base_url = base_url - self._subject_token_file = subject_token_file - self._certificate_authority = certificate_authority - self._provider: WorkloadTokenExchangeProvider | None = None - self._lock = threading.Lock() - - def _get_provider(self) -> WorkloadTokenExchangeProvider: - provider = self._provider - if provider is not None: - return provider - with self._lock: - provider = self._provider - if provider is None: - provider = _create_workload_exchange_provider( - self._base_url, - self._subject_token_file, - certificate_authority=self._certificate_authority, - ) - self._provider = provider - return provider - - def get_cached_access_token(self) -> str | None: - provider = self._provider - if provider is None: - return None - tokens = provider.tokens - if not tokens.access_token or tokens.is_expired(provider.refresh_margin_seconds): - return None - return tokens.access_token - - def get_access_token(self) -> str: - return self._get_provider().get_access_token() - - async def get_access_token_async(self) -> str: - provider = self._provider - if provider is None: - provider = await asyncio.to_thread(self._get_provider) - return await provider.get_access_token_async() - - -# --------------------------------------------------------------------------- -# httpx event hooks — the core of transparent token injection -# --------------------------------------------------------------------------- - - -def _make_auth_event_hook(provider: _AccessTokenProvider) -> _SyncRequestHook: - """Create a **sync** httpx request event hook that injects the Bearer token. - - Called before every SDK HTTP request. ``provider.get_access_token()`` - returns the cached token if still valid, or performs an inline - refresh_token grant if the token is expired/about-to-expire. - """ - - def inject_auth(request: httpx.Request) -> None: - token = provider.get_access_token() - request.headers["Authorization"] = f"Bearer {token}" - - return inject_auth - - -def _make_async_auth_event_hook(provider: _AccessTokenProvider) -> _AsyncRequestHook: - """Create an **async** httpx request event hook for AsyncNeMoPlatform. - - The actual refresh still runs in a worker thread (via - ``provider.get_access_token_async``) so it doesn't block the event loop. - """ - - async def inject_auth(request: httpx.Request) -> None: - token = await provider.get_access_token_async() - request.headers["Authorization"] = f"Bearer {token}" - - return inject_auth - - -def _headers_with_seeded_auth( - headers: Mapping[str, str | Omit], - provider: _AccessTokenProvider, -) -> dict[str, str | Omit]: - seeded_headers = dict(headers) - if isinstance(provider, _LazyWorkloadTokenExchangeProvider): - token = provider.get_cached_access_token() - else: - token = provider.get_access_token() - if token: - seeded_headers["Authorization"] = f"Bearer {token}" - return seeded_headers - - -# --------------------------------------------------------------------------- -# Callbacks wired into OIDCTokenProvider for config-file integration -# --------------------------------------------------------------------------- - - -def _make_config_persister(context_name: str, config_path: Path | None = None): - """Create an ``on_tokens_refreshed`` callback that writes new tokens to the - nmp config file. - - After a successful refresh, the provider calls this so that the CLI and - other SDK processes pick up the rotated tokens without re-authenticating. - """ - from nemo_platform_ext.config.config import Config, ConfigParams - - def persist(tokens: TokenSet) -> None: - params: ConfigParams = {"access_token": tokens.access_token} - if tokens.refresh_token: - params["refresh_token"] = tokens.refresh_token - Config.write(params, context_name=context_name, config_path=config_path) - logger.debug("Persisted refreshed tokens to nmp config (context=%s)", context_name) - - return persist - - -def _make_config_token_loader(context_name: str, config_path: Path): - """Create a ``load_tokens`` callback that re-reads tokens from the config file. - - This is the recovery mechanism for the ``invalid_grant`` scenario: - when another process already rotated the refresh token, our local - copy is stale. The provider calls this to reload whatever that - other process wrote, then retries the refresh with the fresh token. - """ - from nemo_platform_ext.config.config import Config, ConfigParams - from nemo_platform_ext.config.models import OAuthUser - - def load_tokens() -> TokenSet | None: - overrides: ConfigParams = {"current_context": context_name} - try: - config = Config.load(config_path=config_path, overrides=overrides) - resolved = config.resolve() - except Exception: - logger.debug("Failed to reload tokens from nmp config (context=%s)", context_name, exc_info=True) - return None - - if not isinstance(resolved.user, OAuthUser): - return None - - return TokenSet.from_access_token( - resolved.user.token.get_secret_value(), - resolved.user.refresh_token.get_secret_value() if resolved.user.refresh_token else None, - ) - - return load_tokens - - -# --------------------------------------------------------------------------- -# Cross-process file lock for refresh serialization -# --------------------------------------------------------------------------- - - -def _build_refresh_lock_path(config_path: Path, context_name: str) -> Path: - """Derive the lock-file path from the config path and context name. - - Example: ``~/.config/nmp/config.yaml.default.oauth-refresh.lock`` - """ - safe_context = context_name.replace(os.sep, "_") - if os.altsep: - safe_context = safe_context.replace(os.altsep, "_") - return config_path.with_name(f"{config_path.name}.{safe_context}.oauth-refresh.lock") - - -def _make_refresh_lock(config_path: Path, context_name: str): - """Create a ``refresh_lock`` context-manager factory for OIDCTokenProvider. - - Uses ``fcntl.flock`` (POSIX) to serialize refresh transactions across - processes. On platforms without fcntl (Windows), falls back to a no-op - so the provider still works — just without cross-process serialization. - """ - lock_path = _build_refresh_lock_path(config_path, context_name) - - @contextmanager - def refresh_lock(): - try: - import fcntl - except ImportError: - # Windows: no fcntl — skip cross-process locking. - yield - return - - lock_path.parent.mkdir(parents=True, exist_ok=True) - with open(lock_path, "a+", encoding="utf-8") as lock_file: - fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX) - try: - yield - finally: - fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN) - - return refresh_lock - - -def _normalize_config_path(path: Path) -> Path: - """Canonicalize the config path so that cache keys match regardless of - how the caller spelled the path (``~/…`` vs absolute).""" - return path.expanduser().resolve(strict=False) - - -# --------------------------------------------------------------------------- -# Provider cache -# --------------------------------------------------------------------------- - - -def _get_or_create_provider( - key: _ProviderCacheKey, - create_provider: Callable[[], OIDCTokenProvider], -) -> OIDCTokenProvider: - """Return the cached provider for *key*, or create and cache a new one. - - If a cached provider already exists, we call ``reload_tokens()`` so it - picks up any tokens that another process may have written to the config - file since the provider was last used. - """ - with _TOKEN_PROVIDER_CACHE_LOCK: - provider = _TOKEN_PROVIDER_CACHE.get(key) - if provider is None: - provider = create_provider() - _TOKEN_PROVIDER_CACHE[key] = provider - return provider - - # Provider already existed — reload tokens from disk in case another - # process refreshed them since we last checked. - provider.reload_tokens() - return provider - - -# --------------------------------------------------------------------------- -# Config resolution -# --------------------------------------------------------------------------- - - -def _resolve_client_context( - *, - config_path: Path | None, - base_url: str | httpx.URL | None, - context_name: str | None, - access_token: str | None, -) -> tuple[Any, bool, Path]: - """Load the nmp config file, apply overrides, and resolve the active context. - - Returns ``(resolved_context, config_exists, config_path)`` so the caller - knows whether to enable config-backed features (provider caching, - token persistence, file locking). - """ - from nemo_platform_ext.config.config import Config, ConfigParams - - resolved_config_path = config_path or Config.get_default_config_path() - config_exists = resolved_config_path.exists() - if config_exists: - logger.info("Reading nmp config from %s", resolved_config_path) - - # Constructor args override whatever is in the config file. - overrides: ConfigParams | None = None - if context_name is not None or access_token is not None or base_url is not None: - overrides = {} - if base_url is not None: - overrides["base_url"] = str(base_url) - if context_name is not None: - overrides["current_context"] = context_name - if access_token is not None: - overrides["access_token"] = access_token - - config = Config.load(config_path=config_path, overrides=overrides) - if context_name is not None: - available_contexts = [ctx.name for ctx in config.get_config_file().contexts] - if context_name not in available_contexts: - available = ", ".join(available_contexts) if available_contexts else "(none)" - raise ValueError(f"Context '{context_name}' not found. Available contexts: {available}") - - return config.resolve(), config_exists, resolved_config_path - - -# --------------------------------------------------------------------------- -# Bootstrap: config + OIDC discovery → resolved client params -# --------------------------------------------------------------------------- - - -def _resolve_bootstrap( - *, - config_path: Path | None, - base_url: str | httpx.URL | None, - context_name: str | None, - access_token: str | None, +def _split_omitted( extra_headers: Mapping[str, str | Omit] | None, -) -> _ResolvedBootstrap: - """Resolve the full client bootstrap: config, OIDC discovery, token provider. - - For **non-OAuth users** (no-auth): returns a bootstrap with - ``token_provider=None`` and static headers (e.g. ``Authorization: Bearer ``). +) -> tuple[dict[str, str], dict[str, Omit]]: + """Separate real header values from ``Omit`` sentinels. - For **OAuth users**: discovers the cluster's OIDC settings, then either - reuses a cached OIDCTokenProvider or creates a new one. The provider is - wired with: - - - ``load_tokens``: re-reads tokens from the config file (for ``invalid_grant`` recovery) - - ``refresh_lock``: cross-process fcntl lock to serialize refreshes - - ``on_tokens_refreshed``: writes rotated tokens back to the config file + ``Omit`` tells the generated SDK to drop one of its own default headers; the + SDK-independent bootstrap only deals in real values, so the sentinels are + carried around it and merged back into the returned config. """ - from nemo_platform_ext.config.models import OAuthUser - - resolved, config_exists, resolved_config_path = _resolve_client_context( - config_path=config_path, - base_url=base_url, - context_name=context_name, - access_token=access_token, - ) - - base_url = str(resolved.cluster.base_url) - certificate_authority = resolved.cluster.certificate_authority - client_verify = client_verify_from_env(certificate_authority) - headers: dict[str, str | Omit] = dict(extra_headers) if extra_headers else {} - - workload_identity_token_file = _workload_identity_token_file_from_env() - if workload_identity_token_file is not None and access_token is None and not os.environ.get("NMP_ACCESS_TOKEN"): - provider = _LazyWorkloadTokenExchangeProvider( - base_url=base_url, - subject_token_file=workload_identity_token_file, - certificate_authority=certificate_authority, - ) - return _ResolvedBootstrap(base_url, resolved.workspace, headers, provider, client_verify, certificate_authority) - - # --- Non-OAuth path (no auth) --- - if not isinstance(resolved.user, OAuthUser): - user_config = resolved.user.get_client_config() if resolved.user else {} - user_headers = user_config.get("default_headers", {}) - if isinstance(user_headers, dict): - headers.update(user_headers) - return _ResolvedBootstrap(base_url, resolved.workspace, headers, None, client_verify, certificate_authority) - - # --- OAuth path: set up transparent token refresh --- - try: - oidc_config = discover_nmp_config(base_url, certificate_authority=certificate_authority) - if not oidc_config.auth_enabled and access_token is None: - # Discovery succeeded and confirmed the cluster has no OIDC. The - # stored OAuthUser token can't be refreshed here (no endpoint), so - # fall back to no-auth. This guard only fires on a successful - # discovery response — not on discovery failures, where the stored - # token may still be valid and should be used as-is. - return _ResolvedBootstrap(base_url, resolved.workspace, headers, None, client_verify, certificate_authority) - except Exception: - logger.debug("Could not discover OIDC settings from %s", base_url, exc_info=True) - oidc_config = _OIDC_DISCOVERY_FALLBACK + values: dict[str, str] = {} + omitted: dict[str, Omit] = {} + for key, value in (extra_headers or {}).items(): + if isinstance(value, Omit): + omitted[key] = value + else: + values[key] = value + return values, omitted - tokens = TokenSet.from_access_token( - resolved.user.token.get_secret_value(), - resolved.user.refresh_token.get_secret_value() if resolved.user.refresh_token else None, - ) - - token_endpoint = oidc_config.token_endpoint or "" - client_id = oidc_config.client_id or "" - refresh_scope = build_effective_scope(oidc_config.default_scopes, oidc_config.scope_prefix) - - # Only share the provider (and enable persistence/locking) when reading - # from an actual config file. If the caller passed an explicit - # access_token, they own the token lifecycle — don't cache or persist. - share_provider = config_exists and access_token is None - - if share_provider: - normalized_config_path = _normalize_config_path(resolved_config_path) - provider_key = _ProviderCacheKey( - config_path=normalized_config_path, - context_name=resolved.context_name, - token_endpoint=token_endpoint, - client_id=client_id, - refresh_scope=refresh_scope, - certificate_authority=certificate_authority, - ) - on_refreshed = _make_config_persister(resolved.context_name, resolved_config_path) - load_tokens = _make_config_token_loader(resolved.context_name, resolved_config_path) - refresh_lock = _make_refresh_lock(resolved_config_path, resolved.context_name) - provider = _get_or_create_provider( - provider_key, - lambda: OIDCTokenProvider( - token_endpoint=token_endpoint, - client_id=client_id, - tokens=tokens, - refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, - refresh_scope=refresh_scope, - certificate_authority=certificate_authority, - load_tokens=load_tokens, - refresh_lock=refresh_lock, - on_tokens_refreshed=on_refreshed, - ), - ) - else: - # Ephemeral provider: no persistence, no file locking, no caching. - provider = OIDCTokenProvider( - token_endpoint=token_endpoint, - client_id=client_id, - tokens=tokens, - refresh_margin_seconds=_TOKEN_REFRESH_MARGIN_SECONDS, - refresh_scope=refresh_scope, - certificate_authority=certificate_authority, - ) - - return _ResolvedBootstrap(base_url, resolved.workspace, headers, provider, client_verify, certificate_authority) +def _with_omitted(headers: Mapping[str, str] | None, omitted: Mapping[str, Omit]) -> dict[str, str | Omit] | None: + merged: dict[str, str | Omit] = {**(headers or {}), **omitted} + return merged or None # --------------------------------------------------------------------------- @@ -649,19 +112,20 @@ def build_client_init_kwargs( has a request event hook that transparently injects and refreshes the Bearer token before every request. """ - bootstrap = _resolve_bootstrap( + header_values, omitted = _split_omitted(extra_headers) + bootstrap = resolve_bootstrap( config_path=config_path, base_url=base_url, context_name=context_name, access_token=access_token, - extra_headers=extra_headers, + extra_headers=header_values, ) if bootstrap.token_provider is None: # Non-OAuth: static headers, no custom http_client needed. return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, - default_headers=bootstrap.default_headers or None, + default_headers=_with_omitted(bootstrap.default_headers, omitted), client_verify=bootstrap.client_verify, ) @@ -683,7 +147,7 @@ def build_client_init_kwargs( return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, - default_headers=headers or None, + default_headers=_with_omitted(headers, omitted), http_client=http_client, client_verify=bootstrap.client_verify, ) @@ -704,18 +168,19 @@ def build_async_client_init_kwargs( whose event hook calls ``provider.get_access_token_async()`` (runs the refresh in a worker thread so it doesn't block the event loop). """ - bootstrap = _resolve_bootstrap( + header_values, omitted = _split_omitted(extra_headers) + bootstrap = resolve_bootstrap( config_path=config_path, base_url=base_url, context_name=context_name, access_token=access_token, - extra_headers=extra_headers, + extra_headers=header_values, ) if bootstrap.token_provider is None: return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, - default_headers=bootstrap.default_headers or None, + default_headers=_with_omitted(bootstrap.default_headers, omitted), client_verify=bootstrap.client_verify, ) @@ -733,7 +198,7 @@ def build_async_client_init_kwargs( return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, - default_headers=headers or None, + default_headers=_with_omitted(headers, omitted), http_client=http_client, client_verify=bootstrap.client_verify, ) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_code_generator.py b/packages/nemo_platform_ext/tests/cli/core/test_code_generator.py index a1babd15d8..1989fd4933 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_code_generator.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_code_generator.py @@ -1,284 +1,208 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Tests for code generation.""" +"""Tests for ``--output-format code`` generation against typed clients.""" import pytest from nemo_platform_ext.cli.core.code_generator import generate_python_code +from nemo_platform_plugin.inference_gateway.client import InferenceGatewayClient +from nemo_platform_plugin.inference_gateway.types import JsonBody +from nemo_platform_plugin.models.client import ModelsClient +from nemo_platform_plugin.models.types import CreateModelDeploymentRequest +from nemo_platform_plugin.secrets.client import SecretsClient +from nemo_platform_plugin.secrets.types import PlatformSecretCreateRequest def test_generate_python_code_simple_list(): - """Test generating code for a simple list operation.""" - code = generate_python_code( - resource_path=["models"], - method="list", - args={}, - base_url=None, - ) + code = generate_python_code(SecretsClient, "list_secrets", {}, result="list") - assert "from nemo_platform import NeMoPlatform" in code - assert "client = NeMoPlatform()" in code - assert "response = client.models.list()" in code - assert "print(response)" in code + assert "from nemo_platform_plugin.secrets.client import SecretsClient" in code + assert "client = SecretsClient.from_config()" in code + assert "response = client.list_secrets()" in code + assert "for item in response.page().items:" in code + assert "print(item)" in code def test_generate_python_code_with_base_url(): - """Test generating code with base URL.""" - code = generate_python_code( - resource_path=["datasets"], - method="list", - args={}, - base_url="http://test.example.com", - ) + code = generate_python_code(SecretsClient, "list_secrets", {}, base_url="http://test.example.com", result="list") - assert 'client = NeMoPlatform(base_url="http://test.example.com")' in code + assert 'client = SecretsClient(base_url="http://test.example.com")' in code -def test_generate_python_code_with_args(): - """Test generating code with arguments.""" - code = generate_python_code( - resource_path=["models"], - method="retrieve", - args={"model_name": "my-model", "namespace": "default"}, - base_url=None, - ) +def test_generate_python_code_with_args_and_entity_result(): + code = generate_python_code(SecretsClient, "get_secret", {"name": "my-secret", "workspace": "default"}) - assert 'model_name="my-model"' in code - assert 'namespace="default"' in code - assert "client.models.retrieve" in code + assert 'response = client.get_secret(name="my-secret", workspace="default")' in code + assert "print(response.data())" in code -def test_generate_python_code_nested_resource(): - """Test generating code for nested resources.""" - code = generate_python_code( - resource_path=["customization", "configs"], - method="list", - args={}, - base_url=None, - ) +def test_generate_python_code_skips_none_args(): + code = generate_python_code(SecretsClient, "get_secret", {"name": "my-secret", "workspace": None}) - assert "client.customization.configs.list()" in code + assert 'response = client.get_secret(name="my-secret")' in code -def test_generate_python_code_with_dict_args(): - """Test generating code with dictionary arguments.""" - code = generate_python_code( - resource_path=["models"], - method="list", - args={ - "filter": {"namespace": "default", "name": "test"}, - "page": 1, - }, - base_url=None, - ) +def test_generate_python_code_renders_request_model_and_imports_it(): + body = PlatformSecretCreateRequest(name="hf-token", value="s3cret", description="HF token") + code = generate_python_code(SecretsClient, "create_secret", {"workspace": "default", "body": body}) - assert "filter=" in code - assert "page=1" in code - assert '"namespace": "default"' in code or "'namespace': 'default'" in code + assert "from nemo_platform_plugin.secrets.types import PlatformSecretCreateRequest" in code + assert 'body=PlatformSecretCreateRequest(name="hf-token", description="HF token", value="***")' in code + assert "s3cret" not in code + assert code.index("from nemo_platform_plugin.secrets.client") < code.index("client = SecretsClient") -def test_generate_python_code_multiline_format(): - """Test that code with many args is formatted multiline.""" +def test_generate_python_code_renders_root_model_body_positionally(): + body = JsonBody({"model": "default/llama", "messages": [{"role": "user", "content": "hi"}]}) code = generate_python_code( - resource_path=["models"], - method="create", - args={ - "name": "my-model", - "namespace": "default", - "description": "A test model", - "files_url": "s3://bucket/path", - }, - base_url=None, + InferenceGatewayClient, + "provider_post", + {"name": "nvidia", "trailing_uri": "v1/chat/completions", "body": body}, ) - # Should be multiline with many args - lines = code.split("\n") - # Check that args are on separate lines - assert any("name=" in line and line.strip().startswith("name=") for line in lines) - assert any("namespace=" in line and line.strip().startswith("namespace=") for line in lines) + assert "from nemo_platform_plugin.inference_gateway.types import JsonBody" in code + assert 'body=JsonBody({"model": "default/llama", "messages": [{"role": "user", "content": "hi"}]})' in code + compile(code, "", "exec") -def test_generate_python_code_with_platform_job_watch(): +def test_generate_python_code_renders_query_params_dict_and_lists(): code = generate_python_code( - resource_path=["customization", "jobs"], - method="create", - args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, - watch_config={"type": "platform_job", "resource_label": "customization job"}, - watch_options={"timeout": 42, "poll_interval": 7}, + SecretsClient, + "list_secrets", + {"query_params": {"page": 2, "page_size": 10, "sort": ["name", "-created_at"]}}, + result="list", ) - assert "import time" not in code - assert "from nemo_platform.jobs.watch import watch_job" not in code - assert "from nemo_platform_plugin.client.adapter import client_from_platform" in code - assert "from nemo_platform_plugin.jobs.client import JobsClient" in code - assert "jobs_client = client_from_platform(client, JobsClient)" in code - assert "response = client.customization.jobs.create" in code - assert 'resource_name = getattr(response, "name", None) or "job-a"' in code - assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code - assert "jobs_client.watch_job(" in code - assert 'workspace="default"' in code - assert "timeout=42" in code - assert "poll_interval=7" in code - assert "print(event)" in code - assert "get_status" not in code - assert "time.sleep" not in code - assert "print(response)" not in code - compile(code, "", "exec") - - -def test_generate_python_code_with_platform_job_watch_has_no_default_timeout(): - code = generate_python_code( - resource_path=["customization", "jobs"], - method="create", - args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, - watch_config={"type": "platform_job", "resource_label": "customization job"}, - watch_options={"poll_interval": 7}, - ) + assert 'query_params={"page": 2, "page_size": 10, "sort": ["name", "-created_at"]}' in code + + +def test_generate_python_code_renders_datetimes_as_iso_strings(): + from datetime import UTC, datetime + + from pydantic import BaseModel + + class Stamped(BaseModel): + started_at: datetime + + body = Stamped(started_at=datetime(2026, 8, 14, tzinfo=UTC)) + code = generate_python_code(SecretsClient, "create_secret", {"body": body}) + + assert 'Stamped(started_at="2026-08-14T00:00:00+00:00")' in code + assert "datetime.datetime" not in code + + +def test_generate_python_code_renders_root_models_by_root_value(): + from pydantic import RootModel - assert "timeout=None" in code - assert "deadline = time.monotonic()" not in code - assert "poll_interval=7" in code - compile(code, "", "exec") + class Payload(RootModel[dict[str, object]]): + pass + body = Payload({"schema_version": "v1", "agent": {"name": "b"}}) + code = generate_python_code(SecretsClient, "create_secret", {"body": body}) -def test_generate_python_code_with_platform_job_wait(): + assert 'body=Payload({"schema_version": "v1", "agent": {"name": "b"}})' in code + + +def test_generate_python_code_multiline_for_many_args(): code = generate_python_code( - resource_path=["jobs"], - method="create", - args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, - wait_config={"type": "platform_job", "resource_label": "job"}, - wait_options={"timeout": 42, "poll_interval": 7}, + SecretsClient, + "get_secret", + {"name": "a" * 50, "workspace": "default", "extra": 1, "more": 2}, ) - assert "import time" not in code - for symbol in ("JobStatusEvent", "JobWatchTimeoutError", "JobsClient", "NeMoPlatform"): - assert symbol in code - assert "from nemo_platform_plugin.jobs.watch_types import JobStatusEvent, JobWatchTimeoutError" in code - assert "jobs_client = client_from_platform(client, JobsClient)" in code - assert "APIConnectionError" not in code - assert "APIStatusError" not in code - assert "APITimeoutError" not in code - assert "NotFoundError" not in code - assert 'raise RuntimeError("Unable to determine created resource name for --wait")' in code - assert "deadline = time.monotonic()" not in code - assert "get_status" not in code - assert "jobs_client.watch_job(" in code - assert "include_logs=False" in code - assert 'resource_label = "job"' in code - assert "isinstance(event, JobStatusEvent)" in code - assert 'f"{resource_label.title()} {resource_name!r} ended with status {event.status!r}"' in code - assert "except JobWatchTimeoutError as exc:" in code - assert "time.sleep" not in code - assert "print(response)" not in code - compile(code, "", "exec") - - -def test_generate_python_code_with_platform_job_wait_requires_timeout(): - with pytest.raises(ValueError, match=r"wait 'platform_job' lifecycle code generation requires timeout"): + assert "response = client.get_secret(\n" in code + assert ' workspace="default",' in code + + +def test_generate_python_code_no_result_block(): + code = generate_python_code(SecretsClient, "delete_secret", {"name": "x"}, result="none") + + assert code.rstrip().endswith('response = client.delete_secret(name="x")') + + +def test_generate_python_code_binary_result(): + code = generate_python_code(SecretsClient, "download", {"name": "x"}, result="binary") + + assert "with response.stream() as chunks:" in code + + +def test_generate_python_code_rejects_wait_and_watch_together(): + with pytest.raises(ValueError, match="Only one of wait_config or watch_config"): generate_python_code( - resource_path=["jobs"], - method="create", - args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, - wait_config={"type": "platform_job", "resource_label": "job"}, - wait_options={"poll_interval": 7}, + ModelsClient, + "create_deployment", + {}, + wait_config={"type": "inference_deployment"}, + wait_options={"timeout": 10}, + watch_config={"type": "inference_deployment"}, + watch_options={"timeout": 10}, ) -def test_generate_python_code_with_platform_job_wait_requires_resource_label(): - with pytest.raises( - ValueError, - match=r"wait 'platform_job' lifecycle code generation requires a non-empty resource_label", - ): +def test_generate_python_code_rejects_unknown_lifecycle(): + with pytest.raises(ValueError, match="Unsupported lifecycle config type"): + generate_python_code(ModelsClient, "create_deployment", {}, wait_config={"type": "bogus"}, wait_options={}) + + +def test_generate_python_code_inference_deployment_wait_requires_timeout(): + with pytest.raises(ValueError, match="requires timeout"): generate_python_code( - resource_path=["jobs"], - method="create", - args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, - wait_config={"type": "platform_job"}, - wait_options={"timeout": 42, "poll_interval": 7}, + ModelsClient, + "create_deployment", + {}, + wait_config={"type": "inference_deployment"}, + wait_options={"poll_interval": 3}, ) -def test_generate_python_code_with_inference_deployment_wait(): +def test_generate_python_code_inference_deployment_wait_block(): + body = CreateModelDeploymentRequest(name="dep-a", config="cfg") code = generate_python_code( - resource_path=["inference", "deployments"], - method="create", - args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, + ModelsClient, + "create_deployment", + {"workspace": "default", "body": body}, + base_url="http://localhost:8080", wait_config={"type": "inference_deployment", "resource_label": "deployment"}, - wait_options={"timeout": 90, "poll_interval": 10}, + wait_options={"timeout": 120, "poll_interval": 5}, ) assert "import time" in code - for symbol in ("APIConnectionError", "APIStatusError", "APITimeoutError", "NeMoPlatform", "NotFoundError"): - assert symbol in code - assert "deadline = time.monotonic() + 90" in code - assert 'resource_name = getattr(response, "name", None) or "deployment-a"' in code - assert 'raise RuntimeError("Unable to determine created resource name for --wait")' in code - assert 'client.inference.deployments.retrieve(resource_name, workspace="default")' in code - assert 'model_provider_id = getattr(deployment, "model_provider_id", None)' in code - assert 'provider_workspace, _, provider_name = model_provider_id.partition("/")' in code - assert "client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace)" in code - assert "except NotFoundError:" in code - assert "except (APIConnectionError, APITimeoutError):" in code - assert "except APIStatusError as exc:" in code - assert "except Exception:" not in code - assert "response = deployment" in code - assert code.rindex("print(response)") > code.index("response = deployment") - assert "time.sleep(min(10, remaining))" in code - compile(code, "", "exec") - - -def test_generate_python_code_with_inference_deployment_watch(): + assert "from nemo_platform_plugin.client.errors import NemoHTTPError, NemoTransportError, NotFoundError" in code + assert "from nemo_platform_plugin.inference_gateway.client import InferenceGatewayClient" in code + assert 'resource_name = getattr(response.data(), "name", None) or "dep-a"' in code + assert "deadline = time.monotonic() + 120" in code + assert 'client.get_deployment(name=resource_name, workspace="default").data()' in code + assert "gateway = InferenceGatewayClient.from_client(client)" in code + assert "gateway.provider_ready(name=provider_name, workspace=provider_workspace)" in code + assert "time.sleep(min(5, remaining))" in code + assert "--wait" in code + # No Stainless artefacts anywhere in the emitted snippet. + assert "NeMoPlatform" not in code + assert "from nemo_platform import" not in code + + +def test_generate_python_code_watch_mode_mentions_watch_flag(): code = generate_python_code( - resource_path=["inference", "deployments"], - method="create", - args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, - watch_config={"type": "inference_deployment", "resource_label": "deployment"}, - watch_options={"timeout": 90, "poll_interval": 10}, + ModelsClient, + "create_deployment", + {"workspace": "default"}, + watch_config={"type": "inference_deployment"}, + watch_options={"timeout": 10}, ) - assert "import time" in code - for symbol in ("APIConnectionError", "APIStatusError", "APITimeoutError", "NeMoPlatform", "NotFoundError"): - assert symbol in code - assert "deadline = time.monotonic() + 90" in code - assert 'resource_name = getattr(response, "name", None) or "deployment-a"' in code - assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code - assert 'client.inference.deployments.retrieve(resource_name, workspace="default")' in code - assert "client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace)" in code - assert "response = deployment" in code - assert code.rindex("print(response)") > code.index("response = deployment") - assert "time.sleep(min(10, remaining))" in code - compile(code, "", "exec") - - -def test_generate_python_code_with_inference_deployment_wait_requires_timeout(): - with pytest.raises(ValueError, match=r"wait 'inference_deployment' lifecycle code generation requires timeout"): - generate_python_code( - resource_path=["inference", "deployments"], - method="create", - args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, - wait_config={"type": "inference_deployment", "resource_label": "deployment"}, - wait_options={"poll_interval": 10}, - ) + assert "--watch" in code + assert "--wait" not in code -def test_generate_python_code_with_platform_job_watch_ignores_label_formatting(): +def test_generated_snippet_is_valid_python(): + body = CreateModelDeploymentRequest(name="dep-a", config="cfg") code = generate_python_code( - resource_path=["customization", "jobs"], - method="create", - args={"workspace": "default", "name": "job-a"}, - watch_config={"type": "platform_job", "resource_label": 'customization "job" {label}'}, - watch_options={"timeout": 42, "poll_interval": 7}, + ModelsClient, + "create_deployment", + {"workspace": "default", "body": body}, + base_url="http://localhost:8080", + wait_config={"type": "inference_deployment"}, + wait_options={"timeout": 120}, ) - compile(code, "", "exec") - assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code - - -def test_generate_python_code_rejects_unknown_lifecycle_type(): - with pytest.raises(ValueError, match="Unsupported lifecycle config type: 'unknown'"): - generate_python_code( - resource_path=["customization", "jobs"], - method="create", - args={"workspace": "default", "name": "job-a"}, - wait_config={"type": "unknown", "resource_label": "customization job"}, - ) + compile(code, "", "exec") diff --git a/packages/nemo_platform_ext/tests/cli/core/test_context.py b/packages/nemo_platform_ext/tests/cli/core/test_context.py index 27066e4c5f..846d17af9c 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_context.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_context.py @@ -173,3 +173,32 @@ def test_get_async_client_uses_config_bootstrap_for_persisted_oauth_context_and_ ) assert client is mock_client_cls.return_value assert client is client2 + + +def test_typed_client_shares_the_platform_clients_transport_and_auth(): + """typed_client derives a service client from the CLI's platform client, sharing its transport.""" + from nemo_platform_plugin.secrets.client import SecretsClient + + ctx = CLIContext(overrides={"base_url": "http://test.example.com", "access_token": "token-123"}) + resolved_context = SimpleNamespace( + cluster=SimpleNamespace(base_url="http://test.example.com", certificate_authority=None), + context_name="dev", + workspace="test-workspace", + user=OAuthUser(name="dev-user", token="token-123"), + ) + + with patch("nemo_platform_ext.config.config.get_context", return_value=resolved_context): + secrets = ctx.typed_client(SecretsClient) + platform = ctx.get_client() + + assert isinstance(secrets, SecretsClient) + assert secrets._http is platform._client + assert secrets.workspace == "test-workspace" + assert secrets.default_headers == {"Authorization": "Bearer token-123"} + + +def test_get_workspace_returns_none_when_no_context_resolves(): + ctx = CLIContext(overrides={}) + + with patch("nemo_platform_ext.config.config.get_context", side_effect=RuntimeError("no config")): + assert ctx.get_workspace() is None diff --git a/packages/nemo_platform_ext/tests/cli/core/test_errors.py b/packages/nemo_platform_ext/tests/cli/core/test_errors.py index 29fab81738..de8b35adde 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_errors.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_errors.py @@ -25,6 +25,7 @@ from nemo_platform_ext.cli.app import app from nemo_platform_ext.cli.core.errors import ( InvalidSearchPatternError, + UnknownInputFieldsError, _format_api_error, handle_exception, ) @@ -542,3 +543,79 @@ def test_internal_server_error_hint(capsys, body_detail, has_list_cmd, expect_li else: assert "server-side issue" in captured.err assert "list" not in captured.err + + +# --------------------------------------------------------------------------- +# Typed-client error shapes +# --------------------------------------------------------------------------- + + +def _typed_http_error(error_class, status_code, body=None): + request = httpx.Request("GET", "http://test/apis/test/v2/things") + response = httpx.Response( + status_code, json=body, request=request, text=None if body is not None else "Error message" + ) + return error_class(response) + + +def test_format_api_error_uses_typed_http_error_detail(): + error = _typed_http_error(plugin_errors.NotFoundError, 404, {"detail": "Workspace 'default' not found"}) + assert _format_api_error(error) == "Workspace 'default' not found" + + +def test_generic_typed_client_error_maps_to_remote_exit_code(capsys): + with pytest.raises(typer.Exit) as exc_info: + handle_exception(plugin_errors.NemoClientError("Generic API error")) + + assert exc_info.value.exit_code == 3 + assert "API error:" in capsys.readouterr().err + + +@pytest.mark.parametrize( + "message", + [ + "Missing workspace argument", + "Missing path parameter 'workspace' for GET /apis/secrets/v2/workspaces/{workspace}/secrets", + "workspace must be provided when the client has no default workspace", + ], +) +def test_missing_workspace_value_error_is_usage_error(capsys, message): + """Both the generated SDK's and the typed client's unresolved-workspace ValueErrors map to exit 2.""" + with pytest.raises(typer.Exit) as exc_info: + handle_exception(ValueError(message)) + + assert exc_info.value.exit_code == 2 + captured = capsys.readouterr() + assert "Missing workspace:" in captured.err + assert "--workspace" in captured.err + + +def test_pydantic_validation_error_is_usage_error(capsys): + """Client-side request-model validation failures are reported as invalid input with exit code 2.""" + from nemo_platform_plugin.secrets.types import PlatformSecretCreateRequest + from pydantic import ValidationError + + try: + PlatformSecretCreateRequest(name="x", value="v") + except ValidationError as error: + with pytest.raises(typer.Exit) as exc_info: + handle_exception(error) + else: # pragma: no cover - guards the test's own assumption + raise AssertionError("expected a validation error") + + assert exc_info.value.exit_code == 2 + captured = capsys.readouterr() + assert "Invalid input:" in captured.err + assert "name" in captured.err + assert "--help" in captured.err + + +def test_unknown_input_fields_error_is_usage_error(capsys): + with pytest.raises(typer.Exit) as exc_info: + handle_exception(UnknownInputFieldsError(["descripton"], "things create", ["description", "name"])) + + assert exc_info.value.exit_code == 2 + captured = capsys.readouterr() + assert "Unknown input fields: descripton" in captured.err + assert "Accepted fields: description, name" in captured.err + assert "things create --help" in captured.err diff --git a/packages/nemo_platform_ext/tests/cli/core/test_help_formatter.py b/packages/nemo_platform_ext/tests/cli/core/test_help_formatter.py index ec8700b8c7..3e26e28d56 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_help_formatter.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_help_formatter.py @@ -580,3 +580,18 @@ def my_func(): my_func() captured = capsys.readouterr() assert captured.err == "" + + +def test_create_typer_app_does_not_add_completion_options(): + """Groups (including plugin-hosted roots) leave shell completion to the root app.""" + import typer.main + from nemo_platform_ext.cli.core.help_formatter import create_typer_app + + app = create_typer_app(name="group", help="A group") + + @app.command("noop") + def _noop() -> None: + pass + + command = typer.main.get_command(app) + assert not any(param.name in {"install_completion", "show_completion"} for param in command.params) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_legacy_code_generator.py b/packages/nemo_platform_ext/tests/cli/core/test_legacy_code_generator.py new file mode 100644 index 0000000000..7c366cc487 --- /dev/null +++ b/packages/nemo_platform_ext/tests/cli/core/test_legacy_code_generator.py @@ -0,0 +1,284 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for code generation.""" + +import pytest +from nemo_platform_ext.cli.core.legacy_code_generator import generate_python_code + + +def test_generate_python_code_simple_list(): + """Test generating code for a simple list operation.""" + code = generate_python_code( + resource_path=["models"], + method="list", + args={}, + base_url=None, + ) + + assert "from nemo_platform import NeMoPlatform" in code + assert "client = NeMoPlatform()" in code + assert "response = client.models.list()" in code + assert "print(response)" in code + + +def test_generate_python_code_with_base_url(): + """Test generating code with base URL.""" + code = generate_python_code( + resource_path=["datasets"], + method="list", + args={}, + base_url="http://test.example.com", + ) + + assert 'client = NeMoPlatform(base_url="http://test.example.com")' in code + + +def test_generate_python_code_with_args(): + """Test generating code with arguments.""" + code = generate_python_code( + resource_path=["models"], + method="retrieve", + args={"model_name": "my-model", "namespace": "default"}, + base_url=None, + ) + + assert 'model_name="my-model"' in code + assert 'namespace="default"' in code + assert "client.models.retrieve" in code + + +def test_generate_python_code_nested_resource(): + """Test generating code for nested resources.""" + code = generate_python_code( + resource_path=["customization", "configs"], + method="list", + args={}, + base_url=None, + ) + + assert "client.customization.configs.list()" in code + + +def test_generate_python_code_with_dict_args(): + """Test generating code with dictionary arguments.""" + code = generate_python_code( + resource_path=["models"], + method="list", + args={ + "filter": {"namespace": "default", "name": "test"}, + "page": 1, + }, + base_url=None, + ) + + assert "filter=" in code + assert "page=1" in code + assert '"namespace": "default"' in code or "'namespace': 'default'" in code + + +def test_generate_python_code_multiline_format(): + """Test that code with many args is formatted multiline.""" + code = generate_python_code( + resource_path=["models"], + method="create", + args={ + "name": "my-model", + "namespace": "default", + "description": "A test model", + "files_url": "s3://bucket/path", + }, + base_url=None, + ) + + # Should be multiline with many args + lines = code.split("\n") + # Check that args are on separate lines + assert any("name=" in line and line.strip().startswith("name=") for line in lines) + assert any("namespace=" in line and line.strip().startswith("namespace=") for line in lines) + + +def test_generate_python_code_with_platform_job_watch(): + code = generate_python_code( + resource_path=["customization", "jobs"], + method="create", + args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, + watch_config={"type": "platform_job", "resource_label": "customization job"}, + watch_options={"timeout": 42, "poll_interval": 7}, + ) + + assert "import time" not in code + assert "from nemo_platform.jobs.watch import watch_job" not in code + assert "from nemo_platform_plugin.client.adapter import client_from_platform" in code + assert "from nemo_platform_plugin.jobs.client import JobsClient" in code + assert "jobs_client = client_from_platform(client, JobsClient)" in code + assert "response = client.customization.jobs.create" in code + assert 'resource_name = getattr(response, "name", None) or "job-a"' in code + assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code + assert "jobs_client.watch_job(" in code + assert 'workspace="default"' in code + assert "timeout=42" in code + assert "poll_interval=7" in code + assert "print(event)" in code + assert "get_status" not in code + assert "time.sleep" not in code + assert "print(response)" not in code + compile(code, "", "exec") + + +def test_generate_python_code_with_platform_job_watch_has_no_default_timeout(): + code = generate_python_code( + resource_path=["customization", "jobs"], + method="create", + args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, + watch_config={"type": "platform_job", "resource_label": "customization job"}, + watch_options={"poll_interval": 7}, + ) + + assert "timeout=None" in code + assert "deadline = time.monotonic()" not in code + assert "poll_interval=7" in code + compile(code, "", "exec") + + +def test_generate_python_code_with_platform_job_wait(): + code = generate_python_code( + resource_path=["jobs"], + method="create", + args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, + wait_config={"type": "platform_job", "resource_label": "job"}, + wait_options={"timeout": 42, "poll_interval": 7}, + ) + + assert "import time" not in code + for symbol in ("JobStatusEvent", "JobWatchTimeoutError", "JobsClient", "NeMoPlatform"): + assert symbol in code + assert "from nemo_platform_plugin.jobs.watch_types import JobStatusEvent, JobWatchTimeoutError" in code + assert "jobs_client = client_from_platform(client, JobsClient)" in code + assert "APIConnectionError" not in code + assert "APIStatusError" not in code + assert "APITimeoutError" not in code + assert "NotFoundError" not in code + assert 'raise RuntimeError("Unable to determine created resource name for --wait")' in code + assert "deadline = time.monotonic()" not in code + assert "get_status" not in code + assert "jobs_client.watch_job(" in code + assert "include_logs=False" in code + assert 'resource_label = "job"' in code + assert "isinstance(event, JobStatusEvent)" in code + assert 'f"{resource_label.title()} {resource_name!r} ended with status {event.status!r}"' in code + assert "except JobWatchTimeoutError as exc:" in code + assert "time.sleep" not in code + assert "print(response)" not in code + compile(code, "", "exec") + + +def test_generate_python_code_with_platform_job_wait_requires_timeout(): + with pytest.raises(ValueError, match=r"wait 'platform_job' lifecycle code generation requires timeout"): + generate_python_code( + resource_path=["jobs"], + method="create", + args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, + wait_config={"type": "platform_job", "resource_label": "job"}, + wait_options={"poll_interval": 7}, + ) + + +def test_generate_python_code_with_platform_job_wait_requires_resource_label(): + with pytest.raises( + ValueError, + match=r"wait 'platform_job' lifecycle code generation requires a non-empty resource_label", + ): + generate_python_code( + resource_path=["jobs"], + method="create", + args={"workspace": "default", "name": "job-a", "spec": {"training_type": "sft"}}, + wait_config={"type": "platform_job"}, + wait_options={"timeout": 42, "poll_interval": 7}, + ) + + +def test_generate_python_code_with_inference_deployment_wait(): + code = generate_python_code( + resource_path=["inference", "deployments"], + method="create", + args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, + wait_config={"type": "inference_deployment", "resource_label": "deployment"}, + wait_options={"timeout": 90, "poll_interval": 10}, + ) + + assert "import time" in code + for symbol in ("APIConnectionError", "APIStatusError", "APITimeoutError", "NeMoPlatform", "NotFoundError"): + assert symbol in code + assert "deadline = time.monotonic() + 90" in code + assert 'resource_name = getattr(response, "name", None) or "deployment-a"' in code + assert 'raise RuntimeError("Unable to determine created resource name for --wait")' in code + assert 'client.inference.deployments.retrieve(resource_name, workspace="default")' in code + assert 'model_provider_id = getattr(deployment, "model_provider_id", None)' in code + assert 'provider_workspace, _, provider_name = model_provider_id.partition("/")' in code + assert "client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace)" in code + assert "except NotFoundError:" in code + assert "except (APIConnectionError, APITimeoutError):" in code + assert "except APIStatusError as exc:" in code + assert "except Exception:" not in code + assert "response = deployment" in code + assert code.rindex("print(response)") > code.index("response = deployment") + assert "time.sleep(min(10, remaining))" in code + compile(code, "", "exec") + + +def test_generate_python_code_with_inference_deployment_watch(): + code = generate_python_code( + resource_path=["inference", "deployments"], + method="create", + args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, + watch_config={"type": "inference_deployment", "resource_label": "deployment"}, + watch_options={"timeout": 90, "poll_interval": 10}, + ) + + assert "import time" in code + for symbol in ("APIConnectionError", "APIStatusError", "APITimeoutError", "NeMoPlatform", "NotFoundError"): + assert symbol in code + assert "deadline = time.monotonic() + 90" in code + assert 'resource_name = getattr(response, "name", None) or "deployment-a"' in code + assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code + assert 'client.inference.deployments.retrieve(resource_name, workspace="default")' in code + assert "client.inference.gateway.provider.ready(provider_name, workspace=provider_workspace)" in code + assert "response = deployment" in code + assert code.rindex("print(response)") > code.index("response = deployment") + assert "time.sleep(min(10, remaining))" in code + compile(code, "", "exec") + + +def test_generate_python_code_with_inference_deployment_wait_requires_timeout(): + with pytest.raises(ValueError, match=r"wait 'inference_deployment' lifecycle code generation requires timeout"): + generate_python_code( + resource_path=["inference", "deployments"], + method="create", + args={"workspace": "default", "name": "deployment-a", "config": "deployment-config"}, + wait_config={"type": "inference_deployment", "resource_label": "deployment"}, + wait_options={"poll_interval": 10}, + ) + + +def test_generate_python_code_with_platform_job_watch_ignores_label_formatting(): + code = generate_python_code( + resource_path=["customization", "jobs"], + method="create", + args={"workspace": "default", "name": "job-a"}, + watch_config={"type": "platform_job", "resource_label": 'customization "job" {label}'}, + watch_options={"timeout": 42, "poll_interval": 7}, + ) + + compile(code, "", "exec") + assert 'raise RuntimeError("Unable to determine created resource name for --watch")' in code + + +def test_generate_python_code_rejects_unknown_lifecycle_type(): + with pytest.raises(ValueError, match="Unsupported lifecycle config type: 'unknown'"): + generate_python_code( + resource_path=["customization", "jobs"], + method="create", + args={"workspace": "default", "name": "job-a"}, + wait_config={"type": "unknown", "resource_label": "customization job"}, + ) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py index 4839bb4258..9f7421586c 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py @@ -4,42 +4,92 @@ """Tests for pagination utilities.""" import logging +from typing import Any from unittest.mock import Mock +from urllib.parse import parse_qs +import httpx import pytest from nemo_platform_ext.cli.core.pagination import ( AllCursorPagesResponse, AllPagesResponse, + CursorPageResponse, + OffsetPageResponse, PaginationType, _fetch_all_pages_cursor, _fetch_all_pages_page_number, + collect_cursor_pages, + collect_offset_pages, + collect_pages, fetch_all_pages, + warn_if_more_pages, ) +from nemo_platform_plugin.client.client import NemoClient +from nemo_platform_plugin.client.endpoint import get +from nemo_platform_plugin.client.types import CursorPagination, Paginated +from pydantic import BaseModel # ============================================================================= -# Helper to create mock SDK pagination response +# Helpers: real typed paginated responses served by an in-memory transport # ============================================================================= -def create_mock_page_response(data: list, page: int, total_pages: int, page_size: int = 10): - """Create a mock SDK pagination response that supports iter_pages().""" - mock_response = Mock() - mock_response.data = data - mock_response.pagination = Mock() - mock_response.pagination.page = page - mock_response.pagination.total_pages = total_pages - mock_response.pagination.page_size = page_size - mock_response.pagination.current_page_size = len(data) - mock_response.pagination.total_results = total_pages * page_size # approximate - return mock_response +class Item(BaseModel): + name: str -def create_mock_cursor_response(data: list, next_page: str | None = None): - """Create a mock SDK cursor pagination response that supports iter_pages().""" - mock_response = Mock() - mock_response.data = data - mock_response.next_page = next_page - return mock_response +@get("/apis/test/v2/items") +def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: ... + + +@get("/apis/test/v2/logs") +def list_logs(*, query_params: dict[str, Any] | None = None) -> Paginated[Item, CursorPagination]: ... + + +def _offset_client( + pages: dict[int, list[str]], *, page_size: int = 2, envelope: dict[str, Any] | None = None +) -> NemoClient: + """Serve ``pages`` (page number -> names) with offset pagination metadata and optional envelope fields.""" + total_results = sum(len(v) for v in pages.values()) + calls: list[int] = [] + + def handler(request: httpx.Request) -> httpx.Response: + page = int(parse_qs(request.url.query.decode()).get("page", ["1"])[0]) + calls.append(page) + data = pages.get(page, []) + return httpx.Response( + 200, + json={ + "data": [{"name": n} for n in data], + **(envelope or {}), + "pagination": { + "page": page, + "page_size": page_size, + "current_page_size": len(data), + "total_pages": len(pages), + "total_results": total_results, + }, + }, + ) + + client = NemoClient(base_url="http://test", http_client=httpx.Client(transport=httpx.MockTransport(handler))) + client.calls = calls # type: ignore[attr-defined] + return client + + +def _cursor_client(pages: dict[str | None, tuple[list[str], str | None]]) -> NemoClient: + """Serve ``pages`` (cursor -> (names, next_cursor)) with cursor pagination metadata.""" + total = sum(len(v[0]) for v in pages.values()) + + def handler(request: httpx.Request) -> httpx.Response: + cursor = parse_qs(request.url.query.decode()).get("page_cursor", [None])[0] + data, next_page = pages[cursor] + return httpx.Response( + 200, + json={"data": [{"name": n} for n in data], "total": total, "next_page": next_page, "prev_page": None}, + ) + + return NemoClient(base_url="http://test", http_client=httpx.Client(transport=httpx.MockTransport(handler))) # ============================================================================= @@ -134,10 +184,206 @@ def test_pagination_type_is_string(): # ============================================================================= -# fetch_all_pages Tests (using SDK's iter_pages) +# collect_*_pages Tests (typed NemoPaginatedResponse) +# ============================================================================= + + +def test_collect_offset_pages_single_page_keeps_server_metadata(): + client = _offset_client({1: ["a", "b"], 2: ["c"]}) + + result = collect_offset_pages(client.send(list_items()), all_pages=False) + + assert isinstance(result, OffsetPageResponse) + assert [item.name for item in result.data] == ["a", "b"] + assert result.pagination.page == 1 + assert result.pagination.total_pages == 2 + assert result.pagination.total_results == 3 + assert client.calls == [1] + dumped = result.model_dump() + assert dumped["data"] == [{"name": "a"}, {"name": "b"}] + assert dumped["pagination"]["total_pages"] == 2 + + +def test_collect_offset_pages_all_pages_merges_every_page(): + client = _offset_client({1: ["a", "b"], 2: ["c", "d"], 3: ["e"]}) + + result = collect_offset_pages(client.send(list_items()), all_pages=True, show_progress=False) + + assert isinstance(result, AllPagesResponse) + assert [item.name for item in result.data] == ["a", "b", "c", "d", "e"] + assert result.pagination.total_results == 5 + assert result.pagination.total_pages == 1 + assert result.pagination.page_size == 2 + assert client.calls == [1, 2, 3] + + +ENVELOPE = {"filter": {"role": "Viewer"}, "sort": "-created_at", "grouped_by": ["session_id"]} + + +def test_collect_offset_pages_single_page_keeps_envelope_fields_in_wire_order(): + """The server's sort/filter/grouped_by echo must survive into JSON output.""" + client = _offset_client({1: ["a"]}, envelope=ENVELOPE) + + dumped = collect_offset_pages(client.send(list_items()), all_pages=False).model_dump() + + assert list(dumped) == ["data", "filter", "sort", "grouped_by", "pagination"] + assert dumped["filter"] == {"role": "Viewer"} + assert dumped["sort"] == "-created_at" + assert dumped["grouped_by"] == ["session_id"] + + +def test_collect_offset_pages_all_pages_keeps_first_page_envelope(): + client = _offset_client({1: ["a"], 2: ["b"]}, envelope=ENVELOPE) + + dumped = collect_offset_pages(client.send(list_items()), all_pages=True, show_progress=False).model_dump() + + assert dumped["data"] == [{"name": "a"}, {"name": "b"}] + assert dumped["filter"] == {"role": "Viewer"} + assert dumped["sort"] == "-created_at" + assert dumped["pagination"]["total_results"] == 2 + + +def test_collect_offset_pages_without_envelope_fields_emits_only_data_and_pagination(): + client = _offset_client({1: ["a"]}) + + dumped = collect_offset_pages(client.send(list_items()), all_pages=False).model_dump() + + assert list(dumped) == ["data", "pagination"] + + +def test_collect_offset_pages_all_pages_empty(): + client = _offset_client({1: []}) + + result = collect_offset_pages(client.send(list_items()), all_pages=True, show_progress=False) + + assert isinstance(result, AllPagesResponse) + assert result.data == [] + assert result.pagination.total_results == 0 + + +def test_collect_cursor_pages_single_page_keeps_cursors(): + client = _cursor_client({None: (["a", "b"], "cur-2"), "cur-2": (["c"], None)}) + + result = collect_cursor_pages(client.send(list_logs()), all_pages=False) + + assert isinstance(result, CursorPageResponse) + assert [item.name for item in result.data] == ["a", "b"] + assert result.next_page == "cur-2" + assert result.total == 3 + assert result.model_dump()["next_page"] == "cur-2" + + +def test_collect_cursor_pages_all_pages_follows_cursors(): + client = _cursor_client({None: (["a", "b"], "cur-2"), "cur-2": (["c"], "cur-3"), "cur-3": (["d"], None)}) + + result = collect_cursor_pages(client.send(list_logs()), all_pages=True, limit=2, show_progress=False) + + assert isinstance(result, AllCursorPagesResponse) + assert [item.name for item in result.data] == ["a", "b", "c", "d"] + assert result.next_page is None + assert result._limit == 2 + + +def test_collect_pages_dispatches_on_pagination_type(): + offset = collect_pages( + _offset_client({1: ["a"]}).send(list_items()), + all_pages=False, + pagination_type=PaginationType.PAGE_NUMBER, + ) + cursor = collect_pages( + _cursor_client({None: (["a"], None)}).send(list_logs()), + all_pages=False, + pagination_type=PaginationType.CURSOR, + ) + + assert isinstance(offset, OffsetPageResponse) + assert isinstance(cursor, CursorPageResponse) + + +def test_collect_pages_propagates_http_errors(): + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, json={"detail": "boom"}) + + client = NemoClient(base_url="http://test", http_client=httpx.Client(transport=httpx.MockTransport(handler))) + + with pytest.raises(Exception): + collect_offset_pages(client.send(list_items()), all_pages=False) + + +def test_model_dump_handles_plain_dict_items(): + result = OffsetPageResponse([{"name": "a"}, {"nested": [{"x": 1}]}], {"page": 1, "total_pages": 1}) + assert result.model_dump()["data"] == [{"name": "a"}, {"nested": [{"x": 1}]}] + + +# ============================================================================= +# warn_if_more_pages Tests # ============================================================================= +def test_warn_if_more_pages_offset(monkeypatch): + warnings: list[str] = [] + monkeypatch.setattr("nemo_platform_ext.cli.core.pagination.add_warning", warnings.append) + client = _offset_client({1: ["a"], 2: ["b"]}) + + warn_if_more_pages(collect_offset_pages(client.send(list_items()), all_pages=False), PaginationType.PAGE_NUMBER) + assert len(warnings) == 1 + + warnings.clear() + warn_if_more_pages( + collect_offset_pages(client.send(list_items()), all_pages=True, show_progress=False), + PaginationType.PAGE_NUMBER, + ) + assert warnings == [] + + +def test_warn_if_more_pages_cursor(monkeypatch): + warnings: list[str] = [] + monkeypatch.setattr("nemo_platform_ext.cli.core.pagination.add_warning", warnings.append) + client = _cursor_client({None: (["a"], "cur-2"), "cur-2": (["b"], None)}) + + warn_if_more_pages(collect_cursor_pages(client.send(list_logs()), all_pages=False), PaginationType.CURSOR) + assert len(warnings) == 1 + + warnings.clear() + warn_if_more_pages( + collect_cursor_pages(client.send(list_logs()), all_pages=True, show_progress=False), PaginationType.CURSOR + ) + assert warnings == [] + + +def test_warn_if_more_pages_not_paginated(monkeypatch): + warnings: list[str] = [] + monkeypatch.setattr("nemo_platform_ext.cli.core.pagination.add_warning", warnings.append) + warn_if_more_pages(object(), PaginationType.NOT_PAGINATED) + assert warnings == [] + + +# ============================================================================= +# Legacy fetch_all_pages (generated commands, removed with them) +# ============================================================================= + + +def create_mock_page_response(data: list, page: int, total_pages: int, page_size: int = 10): + """Create a mock SDK pagination response that supports iter_pages().""" + mock_response = Mock() + mock_response.data = data + mock_response.pagination = Mock() + mock_response.pagination.page = page + mock_response.pagination.total_pages = total_pages + mock_response.pagination.page_size = page_size + mock_response.pagination.current_page_size = len(data) + mock_response.pagination.total_results = total_pages * page_size # approximate + return mock_response + + +def create_mock_cursor_response(data: list, next_page: str | None = None): + """Create a mock SDK cursor pagination response that supports iter_pages().""" + mock_response = Mock() + mock_response.data = data + mock_response.next_page = next_page + return mock_response + + def test_fetch_all_pages_page_number_single_page(): """Test fetch_all_pages with page-number pagination (single page).""" # Create mock response with iter_pages that yields just itself diff --git a/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py b/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py index 5796e8df27..b136f1b846 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py @@ -11,13 +11,16 @@ import pytest from click import UsageError +from nemo_platform_ext.cli.core.errors import UnknownInputFieldsError from nemo_platform_ext.cli.core.stdin_utils import ( + build_request_body, is_stdin_available, merge_stdin_with_options, read_data_from_stdin, read_secret_from_file, resolve_secret_value, ) +from pydantic import BaseModel, ConfigDict, Field, RootModel, ValidationError class TestIsStdinAvailable: @@ -206,3 +209,61 @@ def test_optional_empty_data_raises(self) -> None: """When --value is empty/whitespace, raises.""" with pytest.raises(UsageError, match="Secret value cannot be empty"): resolve_secret_value(None, " ", required=False) + + +# --------------------------------------------------------------------------- +# build_request_body +# --------------------------------------------------------------------------- + + +class _Body(BaseModel): + name: str + description: str | None = None + schema_: dict | None = Field(default=None, alias="schema") + + +class _PassThroughBody(BaseModel): + model_config = ConfigDict(extra="allow") + model: str + + +class _RootBody(RootModel[dict]): + pass + + +def test_build_request_body_marks_only_provided_keys_as_set() -> None: + body = build_request_body(_Body, {"name": "x", "workspace": "ws"}, exclude={"workspace"}) + + assert body.model_dump(exclude_unset=True) == {"name": "x"} + + +def test_build_request_body_accepts_aliases() -> None: + body = build_request_body(_Body, {"name": "x", "schema": {"a": 1}}) + + assert body.schema_ == {"a": 1} + + +def test_build_request_body_rejects_unknown_keys_with_accepted_list() -> None: + with pytest.raises(UnknownInputFieldsError) as excinfo: + build_request_body(_Body, {"name": "x", "descripton": "typo", "extra": 1}, command_name="things create") + + assert excinfo.value.unknown_fields == ["descripton", "extra"] + assert excinfo.value.command_name == "things create" + assert excinfo.value.known_fields == ["description", "name", "schema_"] + + +def test_build_request_body_validates_types() -> None: + with pytest.raises(ValidationError): + build_request_body(_Body, {"name": ["not", "a", "string"]}) + + +def test_build_request_body_passes_extra_keys_through_for_extra_allow_models() -> None: + body = build_request_body(_PassThroughBody, {"model": "m", "seed": 3}) + + assert body.model_dump() == {"model": "m", "seed": 3} + + +def test_build_request_body_leaves_unstructured_root_models_unconstrained() -> None: + body = build_request_body(_RootBody, {"anything": 1}) + + assert body.root == {"anything": 1} diff --git a/packages/nemo_platform_ext/tests/cli/core/test_waiters.py b/packages/nemo_platform_ext/tests/cli/core/test_waiters.py index 5cbb96b5f1..fbcbc4b335 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_waiters.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_waiters.py @@ -3,15 +3,16 @@ from __future__ import annotations +import sys from collections.abc import Iterator from datetime import datetime, timezone -from types import SimpleNamespace from unittest.mock import MagicMock, patch import httpx import pytest -from nemo_platform import APIConnectionError, APIStatusError, AuthenticationError from nemo_platform_ext.cli.core import waiters +from nemo_platform_plugin.client.client import NemoClient +from nemo_platform_plugin.client.errors import InternalServerError from nemo_platform_plugin.jobs.schemas import PlatformJobStatus, PlatformJobStatusResponse WAITERS_MODULE = "nemo_platform_ext.cli.core.waiters" @@ -59,6 +60,48 @@ def data(self) -> PlatformJobStatusResponse: return self._status +def _deployment_json(status: str, *, status_message: str = "", model_provider_id: str | None = None) -> dict: + return { + "id": "dep-1", + "name": "deployment-a", + "workspace": "default", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", + "entity_version": 1, + "config": "cfg", + "config_version": 1, + "status": status, + "status_message": status_message, + "status_history": [], + "model_provider_id": model_provider_id, + } + + +def _scripted_client(steps: list[httpx.Response | Exception]) -> NemoClient: + """A real NemoClient whose transport answers each request with the next scripted step. + + A step that is an exception is raised from the transport, which the client + surfaces as ``NemoTransportError``. Retries are disabled so every request + consumes exactly one step. + """ + remaining = list(steps) + + def handler(request: httpx.Request) -> httpx.Response: + if not remaining: + raise AssertionError(f"unexpected request {request.method} {request.url}") + step = remaining.pop(0) + if isinstance(step, Exception): + raise step + step.request = request + return step + + return NemoClient( + base_url="http://test", + workspace="default", + http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + def _status_response(status: str | PlatformJobStatus) -> _StatusResponse: return _StatusResponse( PlatformJobStatusResponse( @@ -92,7 +135,7 @@ def _no_telemetry() -> Iterator[None]: @pytest.fixture def frozen_time() -> Iterator[MagicMock]: - with patch(f"{WAITERS_MODULE}.time.time", return_value=0) as time: + with patch(f"{WAITERS_MODULE}.time.monotonic", return_value=0) as time: yield time @@ -131,7 +174,7 @@ def test_platform_job_wait_live_display_recomputes_elapsed() -> None: with ( patch(f"{WAITERS_MODULE}.datetime") as datetime_mock, - patch(f"{WAITERS_MODULE}.time.time", side_effect=[101.0, 109.0]), + patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[101.0, 109.0]), ): datetime_mock.now.return_value.strftime.return_value = "12:34:56" @@ -161,14 +204,9 @@ def test_wait_for_platform_job_uses_dynamic_live_display_for_unchanged_status_po def test_wait_for_inference_deployment_uses_remaining_timeout_for_gateway(gateway_wait: MagicMock) -> None: - client = MagicMock() - client.inference.deployments.retrieve.return_value = SimpleNamespace( - status="READY", - status_message="", - status_history=[], - ) + client = _scripted_client([httpx.Response(200, json=_deployment_json("READY"))]) - with patch(f"{WAITERS_MODULE}.time.time", side_effect=[100.0, 104.0, 104.0, 104.0]): + with patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[100.0, 104.0, 104.0, 104.0]): assert waiters.wait_for_inference_deployment( client, "deployment-a", @@ -184,15 +222,11 @@ def test_wait_for_inference_deployment_uses_remaining_timeout_for_gateway(gatewa def test_wait_for_inference_deployment_uses_model_provider_id_for_gateway(gateway_wait: MagicMock) -> None: - client = MagicMock() - client.inference.deployments.retrieve.return_value = SimpleNamespace( - status="READY", - status_message="", - status_history=[], - model_provider_id="provider-workspace/generated-provider", + client = _scripted_client( + [httpx.Response(200, json=_deployment_json("READY", model_provider_id="provider-workspace/generated-provider"))] ) - with patch(f"{WAITERS_MODULE}.time.time", side_effect=[100.0, 104.0, 104.0, 104.0]): + with patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[100.0, 104.0, 104.0, 104.0]): assert waiters.wait_for_inference_deployment( client, "deployment-a", @@ -205,14 +239,9 @@ def test_wait_for_inference_deployment_uses_model_provider_id_for_gateway(gatewa def test_wait_for_inference_deployment_quiet_mode_uses_quiet_gateway(gateway_wait: MagicMock) -> None: - client = MagicMock() - client.inference.deployments.retrieve.return_value = SimpleNamespace( - status="READY", - status_message="", - status_history=[], - ) + client = _scripted_client([httpx.Response(200, json=_deployment_json("READY"))]) - with patch(f"{WAITERS_MODULE}.time.time", side_effect=[100.0, 104.0, 104.0, 104.0]): + with patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[100.0, 104.0, 104.0, 104.0]): assert waiters.wait_for_inference_deployment( client, "deployment-a", @@ -228,15 +257,12 @@ def test_wait_for_inference_deployment_quiet_mode_uses_quiet_gateway(gateway_wai def test_wait_for_inference_deployment_retries_transient_status_error( frozen_time: MagicMock, waiter_pause: MagicMock, gateway_wait: MagicMock ) -> None: - client = MagicMock() - client.inference.deployments.retrieve.side_effect = [ - APIConnectionError(request=httpx.Request("GET", "http://test")), - SimpleNamespace( - status="READY", - status_message="", - status_history=[], - ), - ] + client = _scripted_client( + [ + httpx.ConnectError("connection refused"), + httpx.Response(200, json=_deployment_json("READY")), + ] + ) assert waiters.wait_for_inference_deployment( client, @@ -252,12 +278,7 @@ def test_wait_for_inference_deployment_retries_transient_status_error( def test_wait_for_inference_deployment_returns_false_on_error_status(frozen_time: MagicMock) -> None: - client = MagicMock() - client.inference.deployments.retrieve.return_value = SimpleNamespace( - status="ERROR", - status_message="boom", - status_history=[], - ) + client = _scripted_client([httpx.Response(200, json=_deployment_json("ERROR", status_message="boom"))]) assert ( waiters.wait_for_inference_deployment( @@ -273,15 +294,10 @@ def test_wait_for_inference_deployment_returns_false_on_error_status(frozen_time def test_wait_for_inference_deployment_does_not_sleep_past_timeout(waiter_pause: MagicMock) -> None: - client = MagicMock() - client.inference.deployments.retrieve.return_value = SimpleNamespace( - status="PENDING", - status_message="", - status_history=[], - ) + client = _scripted_client([httpx.Response(200, json=_deployment_json("PENDING"))] * 3) with ( - patch(f"{WAITERS_MODULE}.time.time", side_effect=[0.0, 0.0, 0.0, 4.0, 5.0, 5.0]), + patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[0.0, 0.0, 0.0, 4.0, 5.0, 5.0]), ): assert ( waiters.wait_for_inference_deployment( @@ -302,9 +318,17 @@ def test_wait_for_platform_job_does_not_sleep_past_timeout() -> None: jobs = MagicMock() jobs.get_job_status.return_value = _status_response("active") + # ``time.monotonic`` is one global attribute, so script it per caller: the + # waiter's own elapsed reads stay at 0 while the watch loop sees the clock + # advance to 4s and then 5s of its 5s deadline. + watch_clock = iter([0.0, 0.0, 4.0, 5.0]) + + def monotonic() -> float: + caller = sys._getframe(1).f_globals["__name__"] + return next(watch_clock) if caller == WATCH_MODULE else 0.0 + with ( - patch(f"{WAITERS_MODULE}.time.time", return_value=0.0), - patch(f"{WATCH_MODULE}.time.monotonic", side_effect=[0.0, 0.0, 4.0, 5.0]), + patch("time.monotonic", monotonic), patch(f"{WATCH_MODULE}.time.sleep") as watch_sleep, ): assert waiters.wait_for_platform_job(jobs, "job-a", workspace="default", timeout=5, poll_interval=10) is False @@ -317,7 +341,7 @@ def test_wait_for_platform_job_retries_transient_status_error(frozen_time: Magic request = httpx.Request("GET", "http://test") response = httpx.Response(503, request=request) jobs.get_job_status.side_effect = [ - APIStatusError("service unavailable", response=response, body=None), + InternalServerError(response), _status_response("completed"), ] @@ -329,12 +353,9 @@ def test_wait_for_platform_job_retries_transient_status_error(frozen_time: Magic def test_wait_for_gateway_does_not_sleep_past_timeout(waiter_pause: MagicMock) -> None: - client = MagicMock() - client.inference.gateway.provider.ready.side_effect = APIConnectionError( - request=httpx.Request("GET", "http://test") - ) + client = _scripted_client([httpx.ConnectError("connection refused")] * 3) - with patch(f"{WAITERS_MODULE}.time.time", side_effect=[0.0, 0.0, 0.0, 4.0, 5.0, 5.0]): + with patch(f"{WAITERS_MODULE}.time.monotonic", side_effect=[0.0, 0.0, 0.0, 4.0, 5.0, 5.0]): assert ( waiters.wait_for_gateway( client, @@ -350,22 +371,28 @@ def test_wait_for_gateway_does_not_sleep_past_timeout(waiter_pause: MagicMock) - def test_wait_for_gateway_returns_false_on_non_transient_status_error(frozen_time: MagicMock) -> None: - client = MagicMock() - request = httpx.Request("GET", "http://test") - response = httpx.Response(401, request=request) - client.inference.gateway.provider.ready.side_effect = AuthenticationError( - "unauthorized", - response=response, - body=None, - ) + client = _scripted_client([httpx.Response(401, json={"detail": "unauthorized"})]) assert waiters.wait_for_gateway(client, "provider-a", workspace="default") is False frozen_time.assert_called() +def test_wait_for_gateway_returns_true_when_provider_ready(frozen_time: MagicMock) -> None: + client = _scripted_client( + [ + httpx.Response(404, json={"detail": "Model provider not found for default/provider-a"}), + httpx.Response(503, json={"detail": "warming up"}), + httpx.Response(200, json={"workspace": "default", "name": "provider-a"}), + ] + ) + + with patch(f"{WAITERS_MODULE}._sleep_until_next_poll", return_value=True): + assert waiters.wait_for_gateway(client, "provider-a", workspace="default") is True + frozen_time.assert_called() + + def test_wait_for_gateway_reraises_unexpected_errors(frozen_time: MagicMock) -> None: - client = MagicMock() - client.inference.gateway.provider.ready.side_effect = RuntimeError("boom") + client = _scripted_client([RuntimeError("boom")]) with pytest.raises(RuntimeError, match="boom"): waiters.wait_for_gateway(client, "provider-a", workspace="default") diff --git a/packages/nemo_platform_ext/tests/cli/telemetry/test_job_events.py b/packages/nemo_platform_ext/tests/cli/telemetry/test_job_events.py index f8a80acdea..8133abfa66 100644 --- a/packages/nemo_platform_ext/tests/cli/telemetry/test_job_events.py +++ b/packages/nemo_platform_ext/tests/cli/telemetry/test_job_events.py @@ -3,6 +3,7 @@ from __future__ import annotations +import sys from collections.abc import Iterator from datetime import datetime, timezone from unittest.mock import MagicMock, patch @@ -230,16 +231,10 @@ def test_duration_uses_job_created_at_when_available() -> None: "completed", created_at=datetime.fromtimestamp(90.0, tz=timezone.utc), ) - # Live snapshot and duration both call time.time(); extra snapshots must not - # exhaust the mock or emit_event is skipped (swallowed in the waiter). - time_calls = {"n": 0} - - def fake_time() -> float: - time_calls["n"] += 1 - return 100.0 if time_calls["n"] <= 2 else 130.0 - + # Elapsed display uses the monotonic clock; the duration compares the job's + # created_at against wall-clock time, which is what the event reports. with ( - patch(f"{WAITERS_MODULE}.time.time", side_effect=fake_time), + patch(f"{WAITERS_MODULE}.time.time", return_value=130.0), patch(EMIT_TARGET) as emit_event, ): assert waiters.wait_for_platform_job(jobs, "job-a", workspace="default") is True @@ -274,9 +269,14 @@ def test_timeout_emits_nothing_and_does_not_crash() -> None: jobs = MagicMock() jobs.get_job_status.return_value = _status_response("active") + watch_clock = iter([0.0, 0.0, 4.0, 5.0]) + + def monotonic() -> float: + caller = sys._getframe(1).f_globals["__name__"] + return next(watch_clock) if caller == WATCH_MODULE else 0.0 + with ( - patch(f"{WAITERS_MODULE}.time.time", return_value=0.0), - patch(f"{WATCH_MODULE}.time.monotonic", side_effect=[0.0, 0.0, 4.0, 5.0]), + patch("time.monotonic", monotonic), patch(f"{WATCH_MODULE}.time.sleep"), patch(EMIT_TARGET) as emit_event, ): diff --git a/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py b/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py new file mode 100644 index 0000000000..606f4a7066 --- /dev/null +++ b/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py @@ -0,0 +1,253 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""The four typed-client builders in ``client.bootstrap``: transport shape, auth wiring, retry, TLS.""" + +from __future__ import annotations + +import json +import time +from base64 import urlsafe_b64encode +from pathlib import Path +from unittest.mock import patch + +import httpx +import pytest +import yaml +from nemo_platform_ext.auth.helpers import NMPOIDCConfig +from nemo_platform_ext.client.bootstrap import ( + DEFAULT_CONNECT_TIMEOUT, + DEFAULT_RETRY_POLICY, + build_async_nemo_client, + build_direct_async_nemo_client, + build_direct_nemo_client, + build_nemo_client, + resolve_timeout, +) +from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.auth import TokenProviderAuth +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.endpoint import get +from nemo_platform_plugin.client.types import RetryPolicy +from pydantic import BaseModel + + +class Probe(BaseModel): + ok: bool + + +@get("/apis/test/v2/probe") +def probe() -> Probe: ... + + +def _wire(client: NemoClient) -> list[httpx.Request]: + """Swap the transport for a recorder that answers every request, keeping the builder's auth hook.""" + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={"ok": True}) + + client._http = httpx.Client( + transport=httpx.MockTransport(handler), auth=client._http.auth, headers=client._http.headers + ) + return seen + + +_OIDC = NMPOIDCConfig(auth_enabled=True, client_id="nmp-client-id", token_endpoint="https://idp/token") + + +def _jwt(exp: float) -> str: + def b64(data: bytes) -> str: + return urlsafe_b64encode(data).rstrip(b"=").decode() + + return ".".join([b64(b'{"alg":"RS256"}'), b64(json.dumps({"exp": exp, "sub": "u"}).encode()), b64(b"sig")]) + + +def _write_config(tmp_path: Path, *, user: dict, certificate_authority: str | None = None) -> Path: + cluster: dict = {"name": "default", "base_url": "http://localhost:8080"} + if certificate_authority: + cluster["certificate_authority"] = certificate_authority + config = { + "current_context": "default", + "clusters": [cluster], + "users": [{"name": "default", **user}], + "contexts": [{"name": "default", "cluster": "default", "user": "default", "workspace": "ws"}], + } + path = tmp_path / "config.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _oauth_config(tmp_path: Path, **kwargs) -> Path: + token = _jwt(time.time() + 3600) + return _write_config(tmp_path, user={"type": "oauth", "token": token, "refresh_token": "r"}, **kwargs) + + +def _api_key_config(tmp_path: Path) -> Path: + return _write_config(tmp_path, user={"type": "api-key", "api_key": "nvapi-secret"}) + + +# --------------------------------------------------------------------------- +# timeout shape +# --------------------------------------------------------------------------- + + +def test_resolve_timeout_keeps_a_short_connect_phase_for_bare_numbers() -> None: + resolved = resolve_timeout(60.0) + + assert resolved.connect == DEFAULT_CONNECT_TIMEOUT + assert resolved.read == 60.0 + assert resolved.write == 60.0 + assert resolved.pool == 60.0 + + +def test_resolve_timeout_default_matches_the_generated_sdk_shape() -> None: + resolved = resolve_timeout(None) + + assert resolved == httpx.Timeout(60.0, connect=5.0) + + +def test_resolve_timeout_respects_an_explicit_httpx_timeout() -> None: + explicit = httpx.Timeout(10.0, connect=30.0) + + assert resolve_timeout(explicit) is explicit + + +@pytest.mark.parametrize("timeout", [None, 60.0, 15]) +def test_direct_builders_apply_the_connect_cap_to_the_transport_and_per_request(timeout: float | None) -> None: + """CLIContext passes a bare float; both the httpx client and the per-request timeout must get connect=5.""" + client = build_direct_nemo_client(base_url="http://localhost:8080", timeout=timeout) + async_client = build_direct_async_nemo_client(base_url="http://localhost:8080", timeout=timeout) + + for built in (client, async_client): + expected = httpx.Timeout(60.0 if timeout is None else timeout, connect=DEFAULT_CONNECT_TIMEOUT) + assert built._http.timeout == expected + assert built._timeout == expected + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_config_builders_apply_the_connect_cap(_discover, tmp_path: Path) -> None: + config = _oauth_config(tmp_path) + + client = build_nemo_client(config_path=config, timeout=45.0) + async_client = build_async_nemo_client(config_path=config, timeout=45.0) + + for built in (client, async_client): + assert built._http.timeout == httpx.Timeout(45.0, connect=DEFAULT_CONNECT_TIMEOUT) + assert built._timeout == httpx.Timeout(45.0, connect=DEFAULT_CONNECT_TIMEOUT) + + +# --------------------------------------------------------------------------- +# retry, TLS, headers +# --------------------------------------------------------------------------- + + +def test_direct_builder_defaults_to_the_generated_sdk_retry_policy() -> None: + client = build_direct_nemo_client(base_url="http://localhost:8080") + + assert client.retry == DEFAULT_RETRY_POLICY + assert DEFAULT_RETRY_POLICY.max_retries == 2 + assert DEFAULT_RETRY_POLICY.retryable_status_codes == (408, 409, 429) + + +def test_direct_builder_honours_a_retry_override() -> None: + policy = RetryPolicy(max_retries=0) + + assert build_direct_nemo_client(base_url="http://localhost:8080", retry=policy).retry is policy + assert build_direct_nemo_client(base_url="http://localhost:8080", retry=None).retry is None + + +def test_direct_builder_sends_default_headers_on_the_wire() -> None: + client = build_direct_nemo_client( + base_url="http://localhost:8080", + workspace="ws", + default_headers={"Authorization": "Bearer api-key", "X-Extra": "1"}, + ) + seen = _wire(client) + + client.send(probe()) + + assert seen[0].headers["Authorization"] == "Bearer api-key" + assert seen[0].headers["X-Extra"] == "1" + assert client.workspace == "ws" + assert client._auth is None + + +def test_direct_builder_uses_the_certificate_authority_for_verification(tmp_path: Path) -> None: + with patch("nemo_platform_ext.client.bootstrap.httpx.Client") as client_cls: + build_direct_nemo_client(base_url="https://nmp.example", certificate_authority="/etc/nmp/ca.pem") + + assert client_cls.call_args.kwargs["verify"] == "/etc/nmp/ca.pem" + + +def test_direct_builder_verifies_by_default() -> None: + with patch("nemo_platform_ext.client.bootstrap.httpx.Client") as client_cls: + build_direct_nemo_client(base_url="https://nmp.example") + + assert client_cls.call_args.kwargs["verify"] is True + + +def test_direct_builder_prefers_the_env_ca_bundle(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + env_ca = tmp_path / "env.pem" + env_ca.write_text("cert") + monkeypatch.setenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, str(env_ca)) + + with patch("nemo_platform_ext.client.bootstrap.httpx.Client") as client_cls: + build_direct_nemo_client(base_url="https://nmp.example", certificate_authority="/other/ca.pem") + + assert client_cls.call_args.kwargs["verify"] == str(env_ca) + + +# --------------------------------------------------------------------------- +# auth wiring +# --------------------------------------------------------------------------- + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_oauth_builder_installs_the_provider_on_both_layers(_discover, tmp_path: Path) -> None: + client = build_nemo_client(config_path=_oauth_config(tmp_path)) + + assert isinstance(client, NemoClient) + assert client._auth is not None + assert isinstance(client._http.auth, TokenProviderAuth) + assert client.workspace == "ws" + assert client.base_url.rstrip("/") == "http://localhost:8080" + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_oauth_builder_sends_the_stored_token_on_the_wire(_discover, tmp_path: Path) -> None: + token = _jwt(time.time() + 3600) + config = _write_config(tmp_path, user={"type": "oauth", "token": token, "refresh_token": "r"}) + client = build_nemo_client(config_path=config) + seen = _wire(client) + + client.send(probe()) + + assert seen[0].headers["Authorization"] == f"Bearer {token}" + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_async_oauth_builder_installs_the_provider_on_both_layers(_discover, tmp_path: Path) -> None: + client = build_async_nemo_client(config_path=_oauth_config(tmp_path)) + + assert isinstance(client, AsyncNemoClient) + assert client._auth is not None + assert isinstance(client._http.auth, TokenProviderAuth) + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_api_key_builder_sends_the_key_as_a_bearer_header(_discover, tmp_path: Path) -> None: + client = build_nemo_client(config_path=_api_key_config(tmp_path)) + seen = _wire(client) + + client.send(probe()) + + assert seen[0].headers["Authorization"] == "Bearer nvapi-secret" + + +@patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_OIDC) +def test_config_builder_workspace_override_wins(_discover, tmp_path: Path) -> None: + client = build_nemo_client(config_path=_api_key_config(tmp_path), workspace="other") + + assert client.workspace == "other" diff --git a/packages/nemo_platform_ext/tests/client/test_client.py b/packages/nemo_platform_ext/tests/client/test_client.py index 0464537128..458381f682 100644 --- a/packages/nemo_platform_ext/tests/client/test_client.py +++ b/packages/nemo_platform_ext/tests/client/test_client.py @@ -94,7 +94,7 @@ def _write_config( class TestCreateClientOAuth: - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_creates_client_from_stored_oauth_tokens(self, _mock_discover, tmp_path): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( @@ -108,7 +108,7 @@ def test_creates_client_from_stored_oauth_tokens(self, _mock_discover, tmp_path) assert str(client.base_url).rstrip("/") == "http://localhost:8080" assert client.workspace == "test-workspace" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_event_hook_injects_fresh_token(self, _mock_discover, tmp_path): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( @@ -126,7 +126,7 @@ def test_event_hook_injects_fresh_token(self, _mock_discover, tmp_path): httpx_client._event_hooks["request"][0](request) assert request.headers["Authorization"] == f"Bearer {token}" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_oauth_uses_sdk_default_httpx_client(self, _mock_discover, tmp_path): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( @@ -142,7 +142,7 @@ def test_oauth_uses_sdk_default_httpx_client(self, _mock_discover, tmp_path): client.close() @patch("nemo_platform_ext.client.factory.DefaultHttpxClient") - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_oauth_uses_nemo_scoped_ca_bundle(self, _mock_discover, mock_default_httpx_client, tmp_path, monkeypatch): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( @@ -163,7 +163,7 @@ def test_oauth_uses_nemo_scoped_ca_bundle(self, _mock_discover, mock_default_htt assert mock_default_httpx_client.call_args.kwargs["verify"] == "/tmp/nemo-ca.pem" @patch("nemo_platform_ext.client.factory.DefaultHttpxClient") - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_oauth_uses_context_certificate_authority( self, _mock_discover, mock_default_httpx_client, tmp_path, monkeypatch ): @@ -189,7 +189,7 @@ def test_oauth_uses_context_certificate_authority( assert _mock_discover.call_args.kwargs["certificate_authority"] == context_ca @patch("nemo_platform_ext.client.factory.DefaultHttpxClient") - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_env_ca_bundle_overrides_context_certificate_authority( self, _mock_discover, mock_default_httpx_client, tmp_path, monkeypatch ): @@ -213,7 +213,7 @@ def test_env_ca_bundle_overrides_context_certificate_authority( assert mock_default_httpx_client.call_args.kwargs["verify"] == "/tmp/env-ca.pem" assert _mock_discover.call_args.kwargs["certificate_authority"] == "/tmp/context-ca.pem" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) @patch("nemo_platform_ext.auth.token_provider.httpx.post") def test_persist_refreshed_tokens_writes_to_config(self, mock_post, _mock_discover, tmp_path): expired_token = _make_jwt({"exp": int(time.time()) - 100, "sub": "user1"}) @@ -244,7 +244,7 @@ def test_persist_refreshed_tokens_writes_to_config(self, mock_post, _mock_discov assert saved_user["token"] == new_token assert saved_user["refresh_token"] == "new_refresh" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_explicit_access_token_overrides_config_auth(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") @@ -265,11 +265,11 @@ class TestCreateClientOAuthUserAuthDisabledCluster: """ @patch( - "nemo_platform.client.factory.discover_nmp_config", + "nemo_platform.client.bootstrap.discover_nmp_config", return_value=NMPOIDCConfig(auth_enabled=False, client_id="", token_endpoint=""), ) @patch( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=NMPOIDCConfig(auth_enabled=False, client_id="", token_endpoint=""), ) @patch("nemo_platform_ext.auth.token_provider.httpx.post") @@ -290,11 +290,11 @@ def test_expired_token_on_auth_disabled_cluster_does_not_attempt_refresh( assert "Authorization" not in client._custom_headers @patch( - "nemo_platform.client.factory.discover_nmp_config", + "nemo_platform.client.bootstrap.discover_nmp_config", return_value=NMPOIDCConfig(auth_enabled=False, client_id="", token_endpoint=""), ) @patch( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=NMPOIDCConfig(auth_enabled=False, client_id="", token_endpoint=""), ) def test_valid_token_on_auth_disabled_cluster_skips_token_provider( @@ -313,7 +313,7 @@ def test_valid_token_on_auth_disabled_cluster_skips_token_provider( assert client._client._event_hooks["request"] == [] @patch( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", side_effect=Exception("network error"), ) def test_discovery_failure_preserves_stored_token(self, _mock_discover, tmp_path): @@ -330,7 +330,7 @@ def test_discovery_failure_preserves_stored_token(self, _mock_discover, tmp_path class TestCreateClientWorkloadIdentity: - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) @patch("nemo_platform_ext.auth.workload_exchange.token_exchange_grant") def test_exchanges_workload_identity_token_file(self, mock_exchange, _mock_discover, tmp_path, monkeypatch): subject_token_file = tmp_path / "workload-token" @@ -362,7 +362,7 @@ def test_exchanges_workload_identity_token_file(self, mock_exchange, _mock_disco assert mock_exchange.call_args.kwargs["scope"] == "openid email groups" @patch("nemo_platform_ext.client.factory.DefaultHttpxClient") - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) @patch("nemo_platform_ext.auth.workload_exchange.token_exchange_grant") def test_workload_identity_discovery_uses_context_certificate_authority( self, mock_exchange, _mock_discover, mock_default_httpx_client, tmp_path, monkeypatch @@ -395,7 +395,7 @@ def default_httpx_client(*args, **kwargs): assert mock_default_httpx_client.call_args.kwargs["verify"] == context_ca @pytest.mark.asyncio - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) @patch("nemo_platform_ext.auth.workload_exchange.token_exchange_grant") async def test_async_exchanges_workload_identity_token_file_at_request_time( self, mock_exchange, _mock_discover, tmp_path, monkeypatch @@ -428,7 +428,7 @@ async def test_async_exchanges_workload_identity_token_file_at_request_time( assert mock_exchange.call_args.kwargs["audience"] == "nemo-platform" assert mock_exchange.call_args.kwargs["scope"] == "openid email groups" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_WORKLOAD_NMP_CONFIG) def test_env_access_token_takes_precedence_over_workload_identity_file(self, _mock_discover, tmp_path, monkeypatch): subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-one\n", encoding="utf-8") @@ -447,7 +447,7 @@ def test_env_access_token_takes_precedence_over_workload_identity_file(self, _mo class TestCreateClientApiKey: - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_creates_client_with_api_key(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") @@ -458,7 +458,7 @@ def test_creates_client_with_api_key(self, _mock_discover, tmp_path): assert "Authorization" in client._custom_headers assert client._custom_headers["Authorization"] == "Bearer nvapi-test-key-123" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_creates_client_with_email_api_key(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="admin@example.com") @@ -524,8 +524,8 @@ def test_explicit_timeout_is_forwarded(self, mock_client_ctor, tmp_path): class TestCreateClientProviderReuse: - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) - @patch("nemo_platform_ext.client.factory.OIDCTokenProvider") + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.OIDCTokenProvider") def test_reuses_oauth_provider_for_same_context(self, mock_provider_cls, _mock_discover, tmp_path): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( @@ -547,9 +547,9 @@ def test_reuses_oauth_provider_for_same_context(self, mock_provider_cls, _mock_d assert callable(provider_kwargs["load_tokens"]) assert callable(provider_kwargs["refresh_lock"]) - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) @patch("nemo_platform_ext.client.factory.DefaultHttpxClient") - @patch("nemo_platform_ext.client.factory.OIDCTokenProvider") + @patch("nemo_platform_ext.client.bootstrap.OIDCTokenProvider") def test_context_certificate_authority_participates_in_provider_cache_key( self, mock_provider_cls, mock_default_httpx_client, _mock_discover, tmp_path, monkeypatch ): @@ -593,7 +593,7 @@ def test_context_certificate_authority_participates_in_provider_cache_key( class TestCreateClientOverrides: - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_base_url_override_uses_explicit_url_with_context_auth(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") @@ -603,7 +603,7 @@ def test_base_url_override_uses_explicit_url_with_context_auth(self, _mock_disco assert client.workspace == "test-workspace" assert client._custom_headers["Authorization"] == "Bearer nvapi-test-key-123" - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_context_override_uses_selected_context(self, _mock_discover, tmp_path): config = { "current_context": "default", @@ -646,7 +646,7 @@ def test_context_override_fails_for_missing_context(self, tmp_path): with pytest.raises(ValueError, match="Context 'missing-context' not found"): create_client(config_path=config_path, context_name="missing-context") - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_access_token_override_uses_bearer_token(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "override-user"}) @@ -665,7 +665,7 @@ def test_explicit_missing_config_file_fails_fast(self, tmp_path): with pytest.raises(FileNotFoundError, match=f"Config file not found at {missing_config_path}"): create_client(config_path=missing_config_path) - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) def test_expired_oauth_token_without_refresh_token_fails(self, _mock_discover, tmp_path): expired_token = _make_jwt({"exp": int(time.time()) - 100, "sub": "user1"}) config_path = _write_config( @@ -677,7 +677,7 @@ def test_expired_oauth_token_without_refresh_token_fails(self, _mock_discover, t with pytest.raises(RuntimeError, match="no refresh token is available"): create_client(config_path=config_path) - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) @patch("nemo_platform_ext.auth.token_provider.httpx.post") def test_refresh_grant_failure_surfaces_clear_error(self, mock_post, _mock_discover, tmp_path): expired_token = _make_jwt({"exp": int(time.time()) - 100, "sub": "user1"}) @@ -1029,7 +1029,7 @@ async def test_async_constructor_passes_context_name_to_bootstrap(self, mock_bui class TestAsyncNeMoPlatformInit: @pytest.mark.asyncio - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) async def test_async_client_uses_config_for_api_key(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") @@ -1100,8 +1100,8 @@ async def test_async_client_raises_attribute_error_for_sync_only_plugin(self): await client.close() @pytest.mark.asyncio - @patch("nemo_platform.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) - @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.bootstrap.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) async def test_async_client_oauth_hook_injects_fresh_token(self, _mock_ext_discover, _mock_sdk_discover, tmp_path): token = _make_jwt({"exp": int(time.time()) + 3600, "sub": "user1"}) config_path = _write_config( diff --git a/packages/nmp_common/tests/sdk_factory/test_sdk.py b/packages/nmp_common/tests/sdk_factory/test_sdk.py index 73c9335910..8aa13c5026 100644 --- a/packages/nmp_common/tests/sdk_factory/test_sdk.py +++ b/packages/nmp_common/tests/sdk_factory/test_sdk.py @@ -194,7 +194,7 @@ def token_exchange_grant(**kwargs): monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) monkeypatch.setattr( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", lambda _base_url, **_kwargs: _workload_oidc_config(), ) monkeypatch.setattr("nemo_platform_ext.auth.workload_exchange.token_exchange_grant", token_exchange_grant) @@ -383,7 +383,7 @@ def token_exchange_grant(**kwargs): monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) monkeypatch.setattr( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", lambda _base_url, **_kwargs: _workload_oidc_config(), ) monkeypatch.setattr("nemo_platform_ext.auth.workload_exchange.token_exchange_grant", token_exchange_grant) From 9c1f1660fd4d9a804d1ffcaf72f83b1c75523f3a Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Fri, 11 Sep 2026 12:06:07 -0400 Subject: [PATCH 3/8] test: satisfy the ty gate and patch OIDC discovery where the bootstrap calls it The new test endpoint stubs raise NotImplementedError like their siblings instead of an ellipsis body, which ty rejects as an implicit None return. The auth-idp CLI refresh contract test patches discover_nmp_config on client.bootstrap, where the CLI now resolves OIDC settings. Signed-off-by: Max Dubrinsky --- .../nemo_platform_ext/tests/cli/core/test_pagination.py | 6 ++++-- .../tests/client/test_bootstrap_builders.py | 3 ++- .../tests/client/test_auth_per_attempt.py | 9 ++++++--- tests/auth_idp/contracts/test_cli_refresh.py | 4 ++-- 4 files changed, 14 insertions(+), 8 deletions(-) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py index 9f7421586c..023fc16efe 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py @@ -39,11 +39,13 @@ class Item(BaseModel): @get("/apis/test/v2/items") -def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: ... +def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: + raise NotImplementedError @get("/apis/test/v2/logs") -def list_logs(*, query_params: dict[str, Any] | None = None) -> Paginated[Item, CursorPagination]: ... +def list_logs(*, query_params: dict[str, Any] | None = None) -> Paginated[Item, CursorPagination]: + raise NotImplementedError def _offset_client( diff --git a/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py b/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py index 606f4a7066..34c920c248 100644 --- a/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py +++ b/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py @@ -37,7 +37,8 @@ class Probe(BaseModel): @get("/apis/test/v2/probe") -def probe() -> Probe: ... +def probe() -> Probe: + raise NotImplementedError def _wire(client: NemoClient) -> list[httpx.Request]: diff --git a/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py b/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py index bea22614a1..e7a95809c0 100644 --- a/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py +++ b/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py @@ -29,15 +29,18 @@ class Item(BaseModel): @get("/apis/test/v2/items") -def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: ... +def list_items(*, query_params: dict[str, Any] | None = None) -> Paginated[Item]: + raise NotImplementedError @get("/apis/test/v2/items/{name}") -def get_item(*, name: str) -> Item: ... +def get_item(*, name: str) -> Item: + raise NotImplementedError @get("/apis/test/v2/download") -def download() -> BinaryContent: ... +def download() -> BinaryContent: + raise NotImplementedError class RotatingProvider: diff --git a/tests/auth_idp/contracts/test_cli_refresh.py b/tests/auth_idp/contracts/test_cli_refresh.py index 0c5125917d..109fbb63e8 100644 --- a/tests/auth_idp/contracts/test_cli_refresh.py +++ b/tests/auth_idp/contracts/test_cli_refresh.py @@ -136,9 +136,9 @@ def test_cli_api_command_auto_refreshes_expired_device_flow_token( device_authorization_endpoint=runtime_device_authorization_endpoint, token_endpoint=runtime_token_endpoint, ) - monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda *_args, **_kwargs: runtime_oidc) + monkeypatch.setattr("nemo_platform.client.bootstrap.discover_nmp_config", lambda *_args, **_kwargs: runtime_oidc) monkeypatch.setattr( - "nemo_platform_ext.client.factory.discover_nmp_config", + "nemo_platform_ext.client.bootstrap.discover_nmp_config", lambda *_args, **_kwargs: runtime_oidc, ) From 055d0e997f5aeec0a94d7e69703c2dceefb86be0 Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Fri, 11 Sep 2026 12:15:00 -0400 Subject: [PATCH 4/8] fix(cli): keep the sort key in merged list output AllPagesResponse and OffsetPageResponse emit the server envelope fields, but callers that build them without an envelope (fetch_all_pages behind the generated list commands, and nemo jobs list --all-pages) lost the sort key that list output has always carried. sort is now always present, null when the server did not echo one, and a real envelope keeps its own field order. Signed-off-by: Max Dubrinsky --- .../nemo_platform_ext/cli/core/pagination.py | 17 +++++++++++++++-- .../tests/cli/core/test_pagination.py | 14 ++++++++++++-- 2 files changed, 27 insertions(+), 4 deletions(-) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py index a1b4663f8f..d34114cf0b 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/pagination.py @@ -51,6 +51,19 @@ class PaginationType(str, Enum): NOT_PAGINATED = "not_paginated" # List operation without pagination support +def _envelope_with_sort(envelope: dict[str, Any] | None) -> dict[str, Any]: + """Return the envelope fields with ``sort`` always present. + + List output has always carried a ``sort`` key (``null`` when the server did + not echo one); callers without an envelope keep that shape, and a server + envelope keeps its own field order. + """ + fields = dict(envelope or {}) + if "sort" not in fields: + fields = {"sort": None, **fields} + return fields + + class AllPagesResponse: """ A response object that mimics a single-page response but contains all items. @@ -68,7 +81,7 @@ def __init__( envelope: dict[str, Any] | None = None, ): self.data = data - self.envelope = dict(envelope or {}) + self.envelope = _envelope_with_sort(envelope) self.sort = self.envelope.get("sort") # Create pagination info for all items @@ -171,7 +184,7 @@ class OffsetPageResponse: def __init__(self, items: list[Any], metadata: dict[str, Any], envelope: dict[str, Any] | None = None) -> None: self.data = items - self.envelope = dict(envelope or {}) + self.envelope = _envelope_with_sort(envelope) self.sort = self.envelope.get("sort") self.pagination = SimpleNamespace(**metadata) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py index 023fc16efe..2c8bdb3470 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_pagination.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_pagination.py @@ -245,12 +245,22 @@ def test_collect_offset_pages_all_pages_keeps_first_page_envelope(): assert dumped["pagination"]["total_results"] == 2 -def test_collect_offset_pages_without_envelope_fields_emits_only_data_and_pagination(): +def test_collect_offset_pages_without_envelope_fields_keeps_a_null_sort(): + """List output always carries ``sort``; a server that echoes nothing yields null, as it always has.""" client = _offset_client({1: ["a"]}) dumped = collect_offset_pages(client.send(list_items()), all_pages=False).model_dump() - assert list(dumped) == ["data", "pagination"] + assert list(dumped) == ["data", "sort", "pagination"] + assert dumped["sort"] is None + + +def test_all_pages_response_without_envelope_keeps_legacy_shape(): + """Callers that build the merged response directly (fetch_all_pages, jobs) keep data/sort/pagination.""" + dumped = AllPagesResponse(data=[{"id": 1}], total_items=1, total_pages=1).model_dump() + + assert list(dumped) == ["data", "sort", "pagination"] + assert dumped["sort"] is None def test_collect_offset_pages_all_pages_empty(): From c0e1d2617ad34cac7d02efc99bc2021e45a72e10 Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Fri, 11 Sep 2026 13:24:01 -0400 Subject: [PATCH 5/8] fix(cli): address review comments on typed-client plumbing - Narrow the platform handle from object to PlatformClient in client_from_platform and platform_default_headers (adapter.py) and the nemo.sdk resource-factory owner in NemoPluginSDKResources (sdk.py), so the typed-client boundary stops erasing the type where plugins are built. - build_request_body reports each field's accepted input alias in the unknown-input hint (schema, not the internal schema_) so the advertised key is one pydantic actually accepts. - with_options now clears a clone's cached nemo.sdk plugin resources, which were built against the original client's transport, so typed clients accessed on the clone bind to its headers/retry/timeout. - Drop the always-false _PLATFORM_JOB_LIFECYCLE watch exclusions from the legacy code generator's lifecycle helpers. Signed-off-by: Max Dubrinsky --- .../cli/core/legacy_code_generator.py | 16 ++---- .../nemo_platform_ext/cli/core/stdin_utils.py | 12 +++- .../tests/cli/core/test_stdin_utils.py | 2 +- .../nemo_platform_plugin/client/adapter.py | 10 ++-- .../src/nemo_platform_plugin/client/client.py | 8 +++ .../src/nemo_platform_plugin/sdk.py | 7 ++- .../tests/client/test_resource_clone.py | 57 +++++++++++++++++++ 7 files changed, 92 insertions(+), 20 deletions(-) create mode 100644 packages/nemo_platform_plugin/tests/client/test_resource_clone.py diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py index 4c0e778be1..52b5877e97 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py @@ -98,9 +98,9 @@ def generate_python_code( lifecycle_mode = "watch" if watch_config else "wait" if wait_config else None - if _lifecycle_uses_deadline(lifecycle_type, lifecycle_mode): + if _lifecycle_uses_deadline(lifecycle_type): lines.append("import time") - if _lifecycle_uses_status_error_handling(lifecycle_type, lifecycle_mode): + if _lifecycle_uses_status_error_handling(lifecycle_type): lines.append( "from nemo_platform import APIConnectionError, APIStatusError, APITimeoutError, NeMoPlatform, NotFoundError" ) @@ -182,16 +182,12 @@ def _format_python_literal(value: Any) -> str: return repr(value) -def _lifecycle_uses_deadline(lifecycle_type: object, mode: str | None) -> bool: - return lifecycle_type in _LIFECYCLE_TYPES_WITH_DEADLINES and not ( - lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" - ) +def _lifecycle_uses_deadline(lifecycle_type: object) -> bool: + return lifecycle_type in _LIFECYCLE_TYPES_WITH_DEADLINES -def _lifecycle_uses_status_error_handling(lifecycle_type: object, mode: str | None) -> bool: - return lifecycle_type in _LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING and not ( - lifecycle_type == _PLATFORM_JOB_LIFECYCLE and mode == "watch" - ) +def _lifecycle_uses_status_error_handling(lifecycle_type: object) -> bool: + return lifecycle_type in _LIFECYCLE_TYPES_WITH_STATUS_ERROR_HANDLING def _require_timeout(timeout: Any, lifecycle_type: object, mode: str | None) -> Any: diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py index fab0049cd5..d5b114b29e 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/stdin_utils.py @@ -192,6 +192,16 @@ def _accepted_field_names(model_cls: type[BaseModel]) -> list[str]: return names +def _accepted_input_names(model_cls: type[BaseModel]) -> list[str]: + """Return each field's accepted input name, preferring its alias. + + The internal Pydantic field name may differ from what a caller types + (e.g. ``schema_`` with ``alias="schema"``), so the user-facing hint must + advertise the alias, not the field name. + """ + return [field.alias or field.validation_alias or name for name, field in model_cls.model_fields.items()] + + def build_request_body( model_cls: type[RequestModelT], payload: Mapping[str, Any], @@ -216,7 +226,7 @@ def build_request_body( if field_model is not None: unknown = sorted(set(body) - set(_accepted_field_names(field_model))) if unknown: - raise UnknownInputFieldsError(unknown, command_name, sorted(field_model.model_fields)) + raise UnknownInputFieldsError(unknown, command_name, sorted(_accepted_input_names(field_model))) return model_cls.model_validate(body) diff --git a/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py b/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py index b136f1b846..15abcfc976 100644 --- a/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py +++ b/packages/nemo_platform_ext/tests/cli/core/test_stdin_utils.py @@ -249,7 +249,7 @@ def test_build_request_body_rejects_unknown_keys_with_accepted_list() -> None: assert excinfo.value.unknown_fields == ["descripton", "extra"] assert excinfo.value.command_name == "things create" - assert excinfo.value.known_fields == ["description", "name", "schema_"] + assert excinfo.value.known_fields == ["description", "name", "schema"] def test_build_request_body_validates_types() -> None: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py index 463fa17a5f..569e236917 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py @@ -12,7 +12,7 @@ from nemo_platform_plugin.client.adapter import client_from_platform - def make_sync_resource(platform: object) -> NemoClient: + def make_sync_resource(platform: PlatformClient) -> NemoClient: return client_from_platform(platform, NemoClient) """ @@ -58,7 +58,7 @@ class _PlatformClient(Protocol): def _prepare_url(self, url: str) -> httpx.URL: ... -def platform_default_headers(platform: object) -> dict[str, str]: +def platform_default_headers(platform: PlatformClient) -> dict[str, str]: """Return a copy of the default headers *platform* sends on every request. Reads ``default_headers`` off a :class:`NemoClient` / :class:`AsyncNemoClient` @@ -72,13 +72,13 @@ def platform_default_headers(platform: object) -> dict[str, str]: @overload -def client_from_platform(platform: object, client_cls: type[SyncT]) -> SyncT: ... +def client_from_platform(platform: PlatformClient, client_cls: type[SyncT]) -> SyncT: ... @overload -def client_from_platform(platform: object, client_cls: type[AsyncT]) -> AsyncT: ... +def client_from_platform(platform: PlatformClient, client_cls: type[AsyncT]) -> AsyncT: ... def client_from_platform( - platform: object, + platform: PlatformClient, client_cls: type[NemoClient] | type[AsyncNemoClient], ) -> NemoClient | AsyncNemoClient: """Create a typed client sharing a platform client's transport. 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 978ef0b90f..abb374c678 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 @@ -526,6 +526,13 @@ def with_options( client.with_options(timeout=300).update_fileset(...) """ clone = copy.copy(self) + # Cached plugin resources were built against the original transport and + # must be rebuilt against the clone's options. + cached = self.__dict__.get("_cached_resources") + if cached: + for name in cached: + clone.__dict__.pop(name, None) + clone.__dict__["_cached_resources"] = set(cached) if headers: clone._default_headers = {**self._default_headers, **headers} if retry is not None: @@ -663,6 +670,7 @@ def __getattr__(self, name: str) -> Any: instance = factory(self) self.__dict__[name] = instance + self.__dict__.setdefault("_cached_resources", set()).add(name) return instance def _resolve_query_params(self, request: PreparedRequest) -> dict[str, str | int | bool] | None: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py index 848abbffdb..ec185d1bb2 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py @@ -7,9 +7,10 @@ from collections.abc import Callable from dataclasses import dataclass -from typing import Any, Generic, TypeVar +from typing import Generic, TypeVar from nemo_platform import AsyncNeMoPlatform, NeMoPlatform +from nemo_platform_plugin.client.adapter import PlatformClient SyncResourceT = TypeVar("SyncResourceT") AsyncResourceT = TypeVar("AsyncResourceT") @@ -26,8 +27,8 @@ class NemoPluginSDKResources(Generic[SyncResourceT, AsyncResourceT]): entry-point surface. """ - sync_resource: Callable[[Any], SyncResourceT] | None = None - async_resource: Callable[[Any], AsyncResourceT] | None = None + sync_resource: Callable[[PlatformClient], SyncResourceT] | None = None + async_resource: Callable[[PlatformClient], AsyncResourceT] | None = None def __post_init__(self) -> None: if self.sync_resource is None and self.async_resource is None: diff --git a/packages/nemo_platform_plugin/tests/client/test_resource_clone.py b/packages/nemo_platform_plugin/tests/client/test_resource_clone.py new file mode 100644 index 0000000000..06526d0b3a --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_resource_clone.py @@ -0,0 +1,57 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""``with_options`` must not carry a clone's cached plugin-resource namespaces across.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import httpx +from nemo_platform_plugin.client.client import NemoClient + +BASE = "http://test:8000" + + +def _make_client(factory: MagicMock, monkeypatch) -> NemoClient: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + client = NemoClient(base_url=BASE, workspace="default", http_client=http_client) + resources = MagicMock(sync_resource=factory, async_resource=None) + discovered = MagicMock(get=MagicMock(return_value=resources)) + monkeypatch.setattr("nemo_platform_plugin.discovery.discover_sdk", lambda: discovered) + return client + + +def test_with_options_rebuilds_cached_plugin_resources(monkeypatch) -> None: + factory = MagicMock(side_effect=lambda _owner: object()) + client = _make_client(factory, monkeypatch) + + first = client.example + assert factory.call_count == 1 + + clone = client.with_options(headers={"X-Extra": "1"}) + + # The clone must not reuse the resource built for the original client. + second = clone.example + assert factory.call_count == 2 + assert second is not first + assert factory.call_args.args[0] is clone + + # The original keeps its own cached resource. + assert client.example is first + assert factory.call_count == 2 + + +def test_with_options_isolates_the_cached_resource_set(monkeypatch) -> None: + factory = MagicMock(side_effect=lambda _owner: object()) + client = _make_client(factory, monkeypatch) + + client.example + clone_a = client.with_options(timeout=5.0) + clone_b = client.with_options(headers={"X-Other": "1"}) + + # The original's cache set is not shared with its clones. + clone_a.other + assert "other" in clone_a.__dict__["_cached_resources"] + assert "other" not in client.__dict__.get("_cached_resources", set()) + assert "other" not in clone_b.__dict__.get("_cached_resources", set()) From 877f379b03df00ed93b00de64434321a8d7b0b2f Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Fri, 11 Sep 2026 14:00:35 -0400 Subject: [PATCH 6/8] fix(plugin): drop the unused nemo_platform SDK re-export in sdk.py The main merge kept the typed resource factories on PlatformClient but took main's __all__ which no longer re-exports NeMoPlatform/AsyncNeMoPlatform, so the import became dead and tripped the ruff pre-commit hook (lint-python-style). Signed-off-by: Max Dubrinsky --- packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py | 1 - 1 file changed, 1 deletion(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py index 3a3b2c22ef..5d20238703 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py @@ -9,7 +9,6 @@ from dataclasses import dataclass from typing import Generic, TypeVar -from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.client.adapter import PlatformClient SyncResourceT = TypeVar("SyncResourceT") From a7fd46cfd3232109b70e07d9f580afff3e936605 Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Fri, 11 Sep 2026 14:07:45 -0400 Subject: [PATCH 7/8] fix(plugin): restore lazy nemo_platform re-export in sdk.py Plugins (nemo-evaluator among others) import the generated NeMoPlatform / AsyncNeMoPlatform classes from nemo_platform_plugin.sdk. My earlier drop of that import broke them with an ImportError. Re-export lazily through a module __getattr__ so plugin discovery still imports this module without requiring the generated SDK at load time. Signed-off-by: Max Dubrinsky --- .../src/nemo_platform_plugin/sdk.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py index 5d20238703..af88e2f37f 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk.py @@ -7,7 +7,7 @@ from collections.abc import Callable from dataclasses import dataclass -from typing import Generic, TypeVar +from typing import Any, Generic, TypeVar from nemo_platform_plugin.client.adapter import PlatformClient @@ -35,5 +35,18 @@ def __post_init__(self) -> None: __all__ = [ + "AsyncNeMoPlatform", # noqa: F822 (resolved lazily by module __getattr__) + "NeMoPlatform", # noqa: F822 (resolved lazily by module __getattr__) "NemoPluginSDKResources", ] + + +def __getattr__(name: str) -> Any: + # Plugins import the generated SDK classes from here. Resolve them on first + # use so plugin discovery, which imports this module for the resource + # container, does not require the generated SDK to be installed. + if name in ("AsyncNeMoPlatform", "NeMoPlatform"): + import nemo_platform + + return getattr(nemo_platform, name) + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") From 633afbfe6011dd656c7e9523986e0166201a040e Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Thu, 10 Sep 2026 17:39:19 -0400 Subject: [PATCH 8/8] refactor(cli): host guardrail, intake and experiments in their owning packages The three functional groups move out of the generated command tree and become nemo.cli entry points on the packages that own the services: guardrail in nemo-guardrails-plugin (GuardrailCLI), intake and experiments in nmp-intake (IntakeCLI, ExperimentsCLI). They appear only when that package is installed, the same way every other plugin group does, and are written on GuardrailClient and IntakeClient via state.typed_client(). Command names, flags, defaults and columns are unchanged; the group help line becomes the standard "Plugin commands for ." row. The intake typed client gains the read and experiment endpoints the CLI needs. The nmp-intake bundle inherits nemo.* entry points so the nemo-platform wheel actually carries the groups (the generated entry-point table is regenerated with make vendor), and a new test asserts every bundled package's nemo.* entry points are exposed by the wrapper or SDK pyproject so a missing inherit fails in CI rather than in the shipped wheel. The generator skips the three resources. GuardrailConfig.data is Optional again after #1966 typed it as a bare RailsConfig: the guardrails entity stores it as Optional and the API returns "data": null for configs created without a body, which the CLI integration tests here create. The middleware already handled None. Rebased over #1965, which added its own experiment endpoints and read models: main's ExperimentCreateRequest/UpdateRequest/Response (dict-typed pareto and column_layout) and RetrieveTraceQueryParams are used; the endpoint stubs keep the create exist_ok / list / delete variants the CLI needs, and span groups go back through collect_offset_pages now that list_span_groups returns Paginated[SpanGroup] on main. Signed-off-by: Max Dubrinsky --- docs/cli/reference.mdx | 6 +- packages/nemo_platform/pyproject.toml | 7 +- .../scripts/docs_generator.py | 3 + .../cli/commands/api/__init__.py | 24 - .../cli/commands/api/intake/__init__.py | 25 - .../cli/commands/api/intake/annotations.py | 273 ---- .../commands/api/intake/evaluator_results.py | 276 ---- .../commands/api/intake/ingest/__init__.py | 21 - .../cli/commands/api/intake/ingest/atif.py | 127 -- .../api/intake/ingest/chat_completions.py | 150 -- .../cli/commands/api/intake/ingest/spans.py | 95 -- .../cli/commands/api/intake/sessions.py | 50 - .../cli/commands/api/intake/spans/__init__.py | 205 --- .../api/intake/spans/evaluator_results.py | 78 - .../cli/commands/api/intake/spans/groups.py | 164 -- .../cli/commands/api/intake/traces.py | 260 --- .../tests/cli/commands/test_agent.py | 9 +- .../nemo_platform_ext/tests/cli/test_app.py | 11 +- .../tests/cli/test_docs_generator.py | 3 + .../nemo_platform_plugin/guardrail/types.py | 2 +- .../src/nemo_platform_plugin/intake/client.py | 9 + .../nemo_platform_plugin/intake/endpoints.py | 113 +- .../src/nemo_platform_plugin/intake/types.py | 256 ++- .../tests/guardrail/test_types.py | 19 +- .../test_read_and_experiment_endpoints.py | 337 ++++ plugins/nemo-guardrails/pyproject.toml | 3 + .../src/nemo_guardrails_plugin/cli.py | 27 + .../cli_commands}/configs.py | 151 +- .../cli_commands/guardrail.py | 35 +- plugins/nemo-guardrails/tests/cli/conftest.py | 99 ++ plugins/nemo-guardrails/tests/cli/test_cli.py | 574 +++++++ .../tests/cli/test_cli_integration.py | 145 ++ services/intake/pyproject.toml | 8 + services/intake/src/nmp/intake/cli.py | 39 + .../src/nmp/intake/cli_commands/common.py | 18 + .../nmp/intake/cli_commands}/experiments.py | 238 ++- .../src/nmp/intake/cli_commands/intake.py | 1417 +++++++++++++++++ services/intake/tests/cli/conftest.py | 104 ++ .../intake/tests/cli/test_experiments_cli.py | 265 +++ .../cli/test_experiments_cli_integration.py | 132 ++ services/intake/tests/cli/test_intake_cli.py | 917 +++++++++++ .../tests/test_clickhouse_architecture.py | 5 +- .../sdk/cli_generator/cli_config.yaml | 44 +- .../sdk/vendor/test_wrapper_entry_points.py | 83 + uv.lock | 12 + 45 files changed, 4808 insertions(+), 2031 deletions(-) delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/__init__.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/annotations.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/evaluator_results.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/ingest/__init__.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/ingest/atif.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/ingest/chat_completions.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/ingest/spans.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/sessions.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/spans/__init__.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/spans/evaluator_results.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/spans/groups.py delete mode 100644 packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/traces.py create mode 100644 packages/nemo_platform_plugin/tests/intake/test_read_and_experiment_endpoints.py create mode 100644 plugins/nemo-guardrails/src/nemo_guardrails_plugin/cli.py rename {packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/guardrail => plugins/nemo-guardrails/src/nemo_guardrails_plugin/cli_commands}/configs.py (67%) rename packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/guardrail/__init__.py => plugins/nemo-guardrails/src/nemo_guardrails_plugin/cli_commands/guardrail.py (90%) create mode 100644 plugins/nemo-guardrails/tests/cli/conftest.py create mode 100644 plugins/nemo-guardrails/tests/cli/test_cli.py create mode 100644 plugins/nemo-guardrails/tests/cli/test_cli_integration.py create mode 100644 services/intake/src/nmp/intake/cli.py create mode 100644 services/intake/src/nmp/intake/cli_commands/common.py rename {packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api => services/intake/src/nmp/intake/cli_commands}/experiments.py (65%) create mode 100644 services/intake/src/nmp/intake/cli_commands/intake.py create mode 100644 services/intake/tests/cli/conftest.py create mode 100644 services/intake/tests/cli/test_experiments_cli.py create mode 100644 services/intake/tests/cli/test_experiments_cli_integration.py create mode 100644 services/intake/tests/cli/test_intake_cli.py create mode 100644 tools/nemo-platform-sdk-tools/tests/sdk/vendor/test_wrapper_entry_points.py diff --git a/docs/cli/reference.mdx b/docs/cli/reference.mdx index e3fa8273fd..69875236c5 100644 --- a/docs/cli/reference.mdx +++ b/docs/cli/reference.mdx @@ -8043,7 +8043,7 @@ Some spec fields cannot be represented as CLI flags: `generate.file_extensions`, ### nemo guardrail -Manage guardrails. +Plugin commands for guardrail. **Usage:** @@ -10492,7 +10492,7 @@ Some spec fields cannot be represented as CLI flags: `config.data.max_sequences_ ### nemo experiments -Manage experiments. +Plugin commands for experiments. **Usage:** @@ -10710,7 +10710,7 @@ nemo experiments update [OPTIONS] PATH_NAME ### nemo intake -Intake operations. +Plugin commands for intake. **Usage:** diff --git a/packages/nemo_platform/pyproject.toml b/packages/nemo_platform/pyproject.toml index 03b8ad68cd..bf6f8c913b 100644 --- a/packages/nemo_platform/pyproject.toml +++ b/packages/nemo_platform/pyproject.toml @@ -201,6 +201,8 @@ intake-service = [ "pydantic>=2.9.2, <3.0.0", "pydantic-settings>=2.6.1, <3.0.0", "nmp-common", + "nemo-platform-plugin", + "nemo-platform-ext", "clickhouse-connect>=0.7,<1.0", "docker>=7.1.0", "opentelemetry-proto>=1.27.0", @@ -670,6 +672,9 @@ data-designer = "nemo_data_designer_plugin.cli.main:DataDesignerCLI" evaluator = "nemo_evaluator.cli:EvaluatorPluginCLI" insights = "nemo_insights_plugin.cli:InsightsCLI" customization = "nemo_customizer.cli:CustomizationCLI" +guardrail = "nemo_guardrails_plugin.cli:GuardrailCLI" +intake = "nmp.intake.cli:IntakeCLI" +experiments = "nmp.intake.cli:ExperimentsCLI" # Generated from [tool.bundle-package]; do not edit this table by hand. [project.entry-points."nemo.cli.agents"] @@ -890,7 +895,7 @@ nmp-inference-gateway = { source = "../../services/core/inference-gateway/src/nm nmp-guardrails = { source = "../../services/guardrails/src/nmp/guardrails", module = "nmp/guardrails", deps_group = "guardrails-service" } nmp-platform-seed = { source = "../../services/platform-seed/src/nmp/platform_seed", module = "nmp/platform_seed", deps_group = "platform-seed-service" } nmp-hello-world = { source = "../../services/hello-world/src/nmp/hello_world", module = "nmp/hello_world", deps_group = "hello-world-service" } -nmp-intake = { source = "../../services/intake/src/nmp/intake", module = "nmp/intake", deps_group = "intake-service" } +nmp-intake = { source = "../../services/intake/src/nmp/intake", module = "nmp/intake", deps_group = "intake-service", inherit = { "entry-points" = ["nemo.*"] } } # Customization task packages: compile glue and schemas only. Their container # entrypoint scripts are deliberately not re-exported, and the GPU stacks they # drive (torch, unsloth, nemo-rl) live in the training images, not this wheel. diff --git a/packages/nemo_platform_ext/scripts/docs_generator.py b/packages/nemo_platform_ext/scripts/docs_generator.py index d88b7583e1..a80ca2d4e3 100644 --- a/packages/nemo_platform_ext/scripts/docs_generator.py +++ b/packages/nemo_platform_ext/scripts/docs_generator.py @@ -708,7 +708,10 @@ def _escape_mdx_line(line: str) -> str: "customization", "data-designer", "evaluator", + "experiments", + "guardrail", "insights", + "intake", "safe-synthesizer", ) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/__init__.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/__init__.py index 6fad1fd8c3..2c9677db40 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/__init__.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/__init__.py @@ -15,14 +15,6 @@ kind="group", hidden=True, ), - TopLevelEntry( - import_path=f"{__package__}.experiments:app", - name="experiments", - help="Manage experiments.", - panel="Functional plugins", - kind="group", - hidden=False, - ), TopLevelEntry( import_path=f"{__package__}.files:app", name="files", @@ -31,14 +23,6 @@ kind="group", hidden=False, ), - TopLevelEntry( - import_path=f"{__package__}.guardrail:app", - name="guardrail", - help="Manage guardrails.", - panel="Functional plugins", - kind="group", - hidden=False, - ), TopLevelEntry( import_path=f"{__package__}.inference:app", name="inference", @@ -47,14 +31,6 @@ kind="group", hidden=False, ), - TopLevelEntry( - import_path=f"{__package__}.intake:app", - name="intake", - help="Intake operations.", - panel="Functional plugins", - kind="group", - hidden=False, - ), TopLevelEntry( import_path=f"{__package__}.models:app", name="models", diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/__init__.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/__init__.py deleted file mode 100644 index 04767cd639..0000000000 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/__init__.py +++ /dev/null @@ -1,25 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# NOTE: This file is auto-generated -from __future__ import annotations - -from importlib import import_module as _importlib_import_module - -from nemo_platform_ext.cli.core.help_formatter import create_typer_app - -_cli_child_annotations = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.annotations") -_cli_child_evaluator_results = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.evaluator_results") -_cli_child_ingest = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.ingest") -_cli_child_sessions = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.sessions") -_cli_child_spans = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.spans") -_cli_child_traces = _importlib_import_module("nemo_platform_ext.cli.commands.api.intake.traces") - -app = create_typer_app(name="intake", help="Intake operations") - -app.add_typer(_cli_child_annotations.app, name="annotations") -app.add_typer(_cli_child_evaluator_results.app, name="evaluator-results") -app.add_typer(_cli_child_ingest.app, name="ingest") -app.add_typer(_cli_child_sessions.app, name="sessions") -app.add_typer(_cli_child_spans.app, name="spans") -app.add_typer(_cli_child_traces.app, name="traces") diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/annotations.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/annotations.py deleted file mode 100644 index 1abd943cb7..0000000000 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/intake/annotations.py +++ /dev/null @@ -1,273 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# NOTE: This file is auto-generated -from __future__ import annotations - -from typing import Annotated, Literal - -import typer - -from nemo_platform_ext.cli.core.api import build_kwargs, merge_filter_dict -from nemo_platform_ext.cli.core.code_generator import handle_code_generation -from nemo_platform_ext.cli.core.context import CLIContext -from nemo_platform_ext.cli.core.errors import handle_errors -from nemo_platform_ext.cli.core.formatters import ( - Column, - check_output_columns_with_format, - format_output, - validate_stream_output_format, -) -from nemo_platform_ext.cli.core.help_formatter import collect_warnings, create_typer_app -from nemo_platform_ext.cli.core.pagination import PaginationType, fetch_all_pages, warn_if_more_pages -from nemo_platform_ext.cli.core.stdin_utils import read_data_input_with_flags, read_payload, validate_required_fields -from nemo_platform_ext.cli.core.types import ( - EntityOutputFormatOption, - ListOutputFormatOption, - NoTruncateOption, - OutputColumnsOption, - StreamOutputOption, -) - -app = create_typer_app(name="annotations", help="Manage annotations") - - -@app.command("create") -@collect_warnings -@handle_errors -def create_annotations( - ctx: typer.Context, - name: Annotated[str | None, typer.Argument()] = None, - workspace: Annotated[str | None, typer.Option("--workspace")] = None, - kind: Annotated[ - Literal["feedback", "note", "metadata", "label"] | None, typer.Option("--kind", help="(required)") - ] = None, - session_id: Annotated[str | None, typer.Option("--session-id", help="(required)")] = None, - value: Annotated[str | None, typer.Option("--value")] = None, - span_id: Annotated[str | None, typer.Option("--span-id")] = None, - text: Annotated[str | None, typer.Option("--text")] = None, - metadata: Annotated[str | None, typer.Option("--metadata", help="JSON string")] = None, - value_type: Annotated[Literal["text", "numeric"] | None, typer.Option("--value-type")] = None, - exist_ok: Annotated[bool | None, typer.Option("--exist-ok")] = None, - input_file: Annotated[ - str | None, - typer.Option("--input-file", help="Path to JSON file (use '-' for stdin)", rich_help_panel="Input Options"), - ] = None, - input_data: Annotated[ - str | None, - typer.Option("--input-data", help="Input data for the request (JSON or YAML)", rich_help_panel="Input Options"), - ] = None, - output_format: EntityOutputFormatOption = None, -) -> None: - """Create annotations. - - [bold red]Required fields:[/] kind, session_id - - [green]Examples:[/] - nemo intake annotations create --input-file config.json - nemo intake annotations create --input-data '{"kind": "value", "session_id": "value"}' - echo '{"json": "data"}' | nemo intake annotations create --input-file - - nemo intake annotations create --