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 e9491b60e1..47327dcd59 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,87 +16,21 @@ 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, + served_model_name_for_entity, + 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 - served_model_name: str | None = 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, - served_model_name: str | None = 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, - served_model_name=served_model_name, - ) - - -def served_model_name_for_entity(provider: ModelProvider, model_entity: ModelEntity) -> str | None: - """Return the provider model id mapped to this Model Entity.""" - entity_ref = f"{model_entity.workspace}/{model_entity.name}" - for mapping in getattr(provider, "served_models", None) or (): - if mapping.model_entity_id == entity_ref: - return mapping.served_model_name - return None - - -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_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 dc0d7d8e06..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,13 +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) or "Missing path parameter 'workspace'" 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." @@ -340,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..52b5877e97 --- /dev/null +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/core/legacy_code_generator.py @@ -0,0 +1,381 @@ +# 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): + lines.append("import time") + if _lifecycle_uses_status_error_handling(lifecycle_type): + 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) -> bool: + return lifecycle_type in _LIFECYCLE_TYPES_WITH_DEADLINES + + +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: + 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..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 @@ -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 @@ -41,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. @@ -49,9 +72,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 = _envelope_with_sort(envelope) + self.sort = self.envelope.get("sort") # Create pagination info for all items self.pagination = type( @@ -80,7 +111,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 +155,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 = _envelope_with_sort(envelope) + 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..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 @@ -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,67 @@ 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 _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], + *, + 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(_accepted_input_names(field_model))) + 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..2c8bdb3470 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,94 @@ """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]: + raise NotImplementedError + + +@get("/apis/test/v2/logs") +def list_logs(*, query_params: dict[str, Any] | None = None) -> Paginated[Item, CursorPagination]: + raise NotImplementedError + + +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 +186,216 @@ 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_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", "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(): + 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..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 @@ -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..34c920c248 --- /dev/null +++ b/packages/nemo_platform_ext/tests/client/test_bootstrap_builders.py @@ -0,0 +1,254 @@ +# 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: + raise NotImplementedError + + +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/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py index 5babe6d9a3..acaa3e30b4 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,20 @@ # 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 generated ``NeMoPlatform`` SDK with the new typed client, -allowing plugins registered via ``NemoPluginSDKResources`` to use the -new endpoint/client infrastructure internally. - -Usage:: - - from nemo_platform_plugin.client.adapter import client_from_platform - - def make_sync_resource(platform: NeMoPlatform) -> NemoClient: - return client_from_platform(platform, NemoClient) +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. """ 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,7 +22,49 @@ def make_sync_resource(platform: NeMoPlatform) -> NemoClient: AsyncT = TypeVar("AsyncT", bound=AsyncNemoClient) -def _platform_default_headers(platform: NeMoPlatform | AsyncNeMoPlatform) -> dict[str, str] | None: +@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: 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` + 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) + + +def _platform_default_headers(platform: _PlatformClient) -> dict[str, str] | None: # 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. @@ -40,51 +76,70 @@ def _platform_default_headers(platform: NeMoPlatform | AsyncNeMoPlatform) -> dic @overload -def client_from_platform(platform: NeMoPlatform, client_cls: type[SyncT]) -> SyncT: ... +def client_from_platform(platform: PlatformClient, client_cls: type[SyncT]) -> SyncT: ... @overload -def client_from_platform(platform: AsyncNeMoPlatform, client_cls: type[AsyncT]) -> AsyncT: ... +def client_from_platform(platform: PlatformClient, client_cls: type[AsyncT]) -> AsyncT: ... def client_from_platform( - platform: NeMoPlatform | AsyncNeMoPlatform, + platform: PlatformClient, 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. """ - headers = _platform_default_headers(platform) + 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) + headers = _platform_default_headers(platform_client) 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, timeout=timeout, retry=retry, - http_client=platform._client, + http_client=platform_client._client, owns_http_client=False, url_resolver=url_resolver, ) @@ -92,12 +147,12 @@ def client_from_platform( 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, timeout=timeout, retry=retry, - http_client=platform._client, + http_client=platform_client._client, owns_http_client=False, 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 2ab45e9f4c..3156bc2cd0 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 @@ -487,6 +492,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 @@ -514,6 +528,13 @@ def with_options( """ clone = copy.copy(self) clone._owns_http = False + # 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: @@ -654,6 +675,31 @@ 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 + self.__dict__.setdefault("_cached_resources", set()).add(name) + 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: @@ -824,14 +870,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) @@ -865,6 +905,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, @@ -877,7 +926,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) @@ -911,7 +964,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: @@ -1100,9 +1157,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) @@ -1136,6 +1190,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, @@ -1148,7 +1209,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) @@ -1182,7 +1247,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/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 848b7c85ae..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,9 +7,9 @@ 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 +from nemo_platform_plugin.client.adapter import PlatformClient SyncResourceT = TypeVar("SyncResourceT") AsyncResourceT = TypeVar("AsyncResourceT") @@ -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[[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: @@ -32,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}") 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 f5aaf303c2..6fe6a20885 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 @@ -198,6 +199,64 @@ def test_client_from_platform_carries_disabled_timeout() -> None: 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())) == {} + + def test_client_from_platform_carries_authorization_header() -> None: """A statically configured bearer reaches the typed client. 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..e7a95809c0 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_auth_per_attempt.py @@ -0,0 +1,244 @@ +# 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]: + raise NotImplementedError + + +@get("/apis/test/v2/items/{name}") +def get_item(*, name: str) -> Item: + raise NotImplementedError + + +@get("/apis/test/v2/download") +def download() -> BinaryContent: + raise NotImplementedError + + +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/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()) 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/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 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) 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, )