diff --git a/packages/data_designer_nemo/src/data_designer_nemo/model_provider.py b/packages/data_designer_nemo/src/data_designer_nemo/model_provider.py index c0d6947f9d..10abb21ea4 100644 --- a/packages/data_designer_nemo/src/data_designer_nemo/model_provider.py +++ b/packages/data_designer_nemo/src/data_designer_nemo/model_provider.py @@ -10,15 +10,11 @@ from data_designer.engine.model_provider import ModelProviderRegistry, resolve_model_provider_registry from data_designer_nemo.errors import NDDInternalError, NDDInvalidConfigError from data_designer_nemo.sdk_translation import sync_to_async_sdk -from nemo_platform import ( - APIConnectionError, - APITimeoutError, - AsyncNeMoPlatform, - NeMoPlatform, - NotFoundError, - PermissionDeniedError, -) -from nemo_platform.types.inference import ModelProvider as NMPModelProvider +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import NemoTransportError, NotFoundError, PermissionDeniedError +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient +from nemo_platform_plugin.models.types import ModelProvider as NMPModelProvider logger = logging.getLogger(__name__) @@ -183,9 +179,10 @@ async def _add_provider( if nmp_provider is None: return + models = client_from_platform(self.sdk, AsyncModelsClient) ndd_provider = NDDModelProvider( name=user_supplied_provider_name, - endpoint=self.sdk.models.get_provider_route_openai_url(nmp_provider), + endpoint=models.get_provider_route_openai_url(nmp_provider), extra_headers={k: v for k, v in self.sdk.default_headers.items() if isinstance(v, str)}, ) providers = (ndd_provider, nmp_provider) @@ -202,7 +199,7 @@ async def _get_nmp_provider( f"Cannot access provider {user_supplied_provider_name!r}. Check that it exists and you have access to it." ) self.inaccessible_providers.add(user_supplied_provider_name) - except (APIConnectionError, APITimeoutError) as e: + except NemoTransportError as e: logger.debug( "Error connecting while retrieving model provider", extra={"provider_name": provider_name, "workspace": workspace}, @@ -265,14 +262,11 @@ async def make_model_provider_registry( def get_nmp_provider(sdk: NeMoPlatform, workspace: str, provider_name: str) -> NMPModelProvider: - return sdk.inference.providers.retrieve( - workspace=workspace, - name=provider_name, - ) + models = client_from_platform(sdk, ModelsClient) + return models.get_provider(workspace=workspace, name=provider_name).data() async def get_nmp_provider_async(sdk: AsyncNeMoPlatform, workspace: str, provider_name: str) -> NMPModelProvider: - return await sdk.inference.providers.retrieve( - workspace=workspace, - name=provider_name, - ) + models = client_from_platform(sdk, AsyncModelsClient) + response = await models.get_provider(workspace=workspace, name=provider_name) + return response.data() diff --git a/packages/data_designer_nemo/tests/unit/test_models_client_migration.py b/packages/data_designer_nemo/tests/unit/test_models_client_migration.py new file mode 100644 index 0000000000..2de4e58726 --- /dev/null +++ b/packages/data_designer_nemo/tests/unit/test_models_client_migration.py @@ -0,0 +1,41 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from data_designer_nemo.model_provider import get_nmp_provider, get_nmp_provider_async +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient + + +def test_get_nmp_provider_uses_models_client() -> None: + sdk = MagicMock() + provider = MagicMock() + response = MagicMock() + response.data.return_value = provider + models = MagicMock() + models.get_provider.return_value = response + + with patch("data_designer_nemo.model_provider.client_from_platform", return_value=models) as make_client: + result = get_nmp_provider(sdk, "workspace", "provider") + + make_client.assert_called_once_with(sdk, ModelsClient) + models.get_provider.assert_called_once_with(workspace="workspace", name="provider") + assert result is provider + + +@pytest.mark.asyncio +async def test_get_nmp_provider_async_uses_models_client() -> None: + sdk = MagicMock() + provider = MagicMock() + response = MagicMock() + response.data.return_value = provider + models = MagicMock() + models.get_provider = AsyncMock(return_value=response) + + with patch("data_designer_nemo.model_provider.client_from_platform", return_value=models) as make_client: + result = await get_nmp_provider_async(sdk, "workspace", "provider") + + make_client.assert_called_once_with(sdk, AsyncModelsClient) + models.get_provider.assert_awaited_once_with(workspace="workspace", name="provider") + assert result is provider diff --git a/packages/models/src/models/resources.py b/packages/models/src/models/resources.py index b3faf5e11a..2de127c026 100644 --- a/packages/models/src/models/resources.py +++ b/packages/models/src/models/resources.py @@ -1,41 +1,63 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Extended ModelsResource with high-level helper methods.""" +"""High-level Models helpers built on the typed :class:`ModelsClient`. + +Historically ``ModelsResource`` / ``AsyncModelsResource`` extended the +Stainless-generated ``ModelsResource`` base to inherit its CRUD surface and +layered convenience helpers on top. As part of AIRCORE-876 the Stainless- +resource *inheritance* is removed: these classes now hold a ``NeMoPlatform`` +SDK and drive the typed ``nemo_platform_plugin.models.client.ModelsClient`` +(built from that SDK via ``client_from_platform``) for their own genuine public +surface -- the inference-gateway route builders, OpenAI client factories, and +deployment/provider status polling. + +CRUD (``retrieve`` / ``create`` / ``list`` / adapter sub-resource) is +intentionally *not* re-implemented here: reproducing the Stainless resource +method/param shapes would be a compatibility proxy. Callers that need CRUD +should use the typed client directly, e.g.:: + + from nemo_platform_plugin.client.adapter import client_from_platform + from nemo_platform_plugin.models.client import ModelsClient + + models = client_from_platform(sdk, ModelsClient) + entity = models.get_model(name="llama", workspace="default").data() + +.. warning:: + **Vendoring gate.** This package is vendored into the SDK as + ``nemo_platform.models`` (``sdk.models`` resolves to this ``ModelsResource`` + via ``packages/models`` ``[tool.vendor-package]``). Because the Stainless + CRUD inheritance is dropped above, running ``make vendor`` will remove + ``sdk.models.retrieve`` / ``create`` / ``list`` / ``adapters.*`` from the + vendored SDK and break every consumer still on that surface (automodel + compiler, provider/deployment reconcilers, models_controller, adapter + sidecar, model_spec task, ``nmp_customization_common``, evaluator resolver, + inference-gateway model cache, generated CLI). Those call sites must first be + migrated to the typed client and repointed from ``nemo_platform`` exceptions + to ``nemo_platform_plugin.client.errors`` (``NotFoundError`` / ``ConflictError`` + are distinct classes). Do not run ``make vendor`` for this package until that + consumer migration lands -- it is a separate follow-up under the AIRCORE-827 + migration umbrella. + +The inference-gateway *readiness* probe (:meth:`wait_for_gateway`) still calls +through the ``NeMoPlatform`` SDK because it targets the separate +inference-gateway service, which has not yet been migrated to a typed client. +""" + +from __future__ import annotations -import asyncio import time from datetime import datetime -from nemo_platform import NotFoundError -from nemo_platform.resources.models import AsyncModelsResource as BaseAsyncModelsResource -from nemo_platform.resources.models import ModelsResource as BaseModelsResource +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform, NotFoundError from nemo_platform.types.inference import ModelDeployment, ModelProvider from nemo_platform.types.models import ModelEntity +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient -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: - return None - if isinstance(entry_timestamp, str): - try: - entry_timestamp = datetime.fromisoformat(entry_timestamp.replace("Z", "+00:00")) - except (ValueError, TypeError): - return None - if not hasattr(entry_timestamp, "timestamp") or not hasattr(created_at, "timestamp"): - return None - try: - return int(entry_timestamp.timestamp() - created_at.timestamp()) - except (TypeError, OSError): - return None - - -class ModelsResource(BaseModelsResource): - """Extended ModelsResource with high-level helper methods. - - All existing methods (create, retrieve, list, etc.) work unchanged. - Adds convenience methods for OpenAI integration and deployment management. +class ModelsResource: + """Sync Models helpers backed by a typed :class:`ModelsClient`. Example: >>> sdk = NeMoPlatform(base_url="http://nmp-host", workspace="default") @@ -43,152 +65,56 @@ class ModelsResource(BaseModelsResource): >>> sdk.models.wait_for_status("my-deployment", "READY") """ + def __init__(self, client: NeMoPlatform) -> None: + self._client = client + self._typed: ModelsClient | None = None + + @property + def models(self) -> ModelsClient: + """The typed Models client sharing this SDK's transport (built lazily).""" + if self._typed is None: + self._typed = client_from_platform(self._client, ModelsClient) + return self._typed + def _get_base_url_str(self) -> str: """Get the base URL as a string with trailing slash removed.""" return str(self._client.base_url).rstrip("/") - def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: - """ - Generate the base URL for the OpenAI proxy route. - - This route uses the `model` field in the request body for routing, - formatted as `workspace/model_entity_name`. + # -- OpenAI inference-gateway route builders (delegate to the typed client) -- - Args: - workspace: The workspace identifier - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> base_url = sdk.models.get_openai_route_base_url() - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/openai/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - base_url = self._get_base_url_str() - return f"{base_url}/apis/inference-gateway/v2/workspaces/{workspace}/openai/-/v1" + def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: + """Base URL for the OpenAI proxy route (routes on the request body ``model``).""" + return self.models.get_openai_route_base_url(workspace=workspace) def get_client_default_headers(self) -> dict[str, str]: - """Get string-only default headers for third-party client libraries. + """String-only default headers for third-party client libraries (OpenAI SDK, LiteLLM). - Use this helper when constructing external clients (for example OpenAI - SDK or LiteLLM) so auth and identity headers from the SDK are forwarded. - This is required for successful inference requests when platform auth/ - authorization is enabled. + Forwards the SDK's auth/identity headers, required for inference when + platform authorization is enabled. """ return {key: value for key, value in self._client.default_headers.items() if isinstance(value, str)} def get_openai_client(self, *, workspace: str | None = None): - """ - Get a sync OpenAI client configured for NeMo Platform's inference gateway. - - This method returns an OpenAI client with the base_url set to the - OpenAI proxy route for the specified workspace. The client can be - used directly with the standard OpenAI SDK interface. - - Args: - workspace: The workspace identifier - - Returns: - An OpenAI client configured for the inference gateway - - Example: - >>> client = sdk.models.get_openai_client() - >>> response = client.chat.completions.create( - ... model="default/meta_llama-3.2-1b-instruct", - ... messages=[{"role": "user", "content": "Hello!"}] - ... ) - """ + """A sync OpenAI client configured for NeMo Platform's inference gateway.""" import openai base_url = self.get_openai_route_base_url(workspace=workspace) - # Preserve auth and identity headers from the parent SDK client. default_headers = self.get_client_default_headers() return openai.OpenAI(base_url=base_url, api_key="not-needed", default_headers=default_headers) def get_provider_route_openai_url(self, provider: ModelProvider) -> str: - """ - Generate an OpenAI SDK-compatible URL for the provider proxy route. - - Handles the conditional /v1 suffix based on the provider's host_url: - - If host_url ends with /v1, no suffix is added - - Otherwise, /v1 is appended - - Args: - provider: The ModelProvider object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> provider = sdk.inference.providers.retrieve("my-provider", workspace="default") - >>> base_url = sdk.models.get_provider_route_openai_url(provider) - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - base_url = self._get_base_url_str() - route = f"{base_url}/apis/inference-gateway/v2/workspaces/{provider.workspace}/provider/{provider.name}/-" - - host_url_normalized = provider.host_url.rstrip("/") - if not host_url_normalized.endswith("/v1"): - route = f"{route}/v1" - - return route + """OpenAI SDK-compatible URL for a provider proxy route (conditional ``/v1``).""" + return self.models.get_provider_route_openai_url(provider) def get_provider_route_openai_url_for_deployment(self, deployment: ModelDeployment) -> str: - """ - Generate an OpenAI SDK-compatible URL for a deployment's model provider. - - This is a convenience method that fetches the ModelProvider associated - with the deployment and returns the provider route URL. - - Args: - deployment: The ModelDeployment object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Raises: - ValueError: If the deployment has no associated model_provider_id - - Example: - >>> deployment = sdk.inference.deployments.retrieve("my-deployment", workspace="default") - >>> base_url = sdk.models.get_provider_route_openai_url_for_deployment(deployment) - >>> openai_client = OpenAI(base_url=base_url) - """ - if not deployment.model_provider_id: - raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") - - # model_provider_id is in "workspace/name" format - workspace, name = deployment.model_provider_id.split("/", 1) - provider = self._client.inference.providers.retrieve(name, workspace=workspace) - return self.get_provider_route_openai_url(provider) + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + return self.models.get_provider_route_openai_url_for_deployment(deployment) def get_model_entity_route_openai_url(self, model_entity: ModelEntity) -> str: - """ - Generate an OpenAI SDK-compatible URL for the model entity proxy route. - - Always appends /v1 suffix since the client doesn't interact directly - with the provider's host_url. + """OpenAI SDK-compatible URL for a model-entity proxy route (always ``/v1``).""" + return self.models.get_model_entity_route_openai_url(model_entity) - Args: - model_entity: The ModelEntity object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> entity = sdk.models.retrieve("my-model", workspace="default") - >>> base_url = sdk.models.get_model_entity_route_openai_url(entity) - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/model/my-model/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - base_url = self._get_base_url_str() - return ( - f"{base_url}/apis/inference-gateway/v2/workspaces/{model_entity.workspace}/model/{model_entity.name}/-/v1" - ) + # -- Deployment / provider status polling -- def wait_for_status( self, @@ -199,85 +125,31 @@ def wait_for_status( timeout: int = 1200, check_gateway: bool = True, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. - - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", verify the - gateway can route to the provider before returning (default: True). + """Wait for a ModelDeployment to reach ``desired_status``. - Returns: - True if desired status reached, False if timeout + When ``desired_status`` is ``"READY"`` and ``check_gateway`` is set, also + waits for the inference gateway to be able to route to the provider. """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - deployment_status = self.wait_for_deployment_status( + if not self.models.wait_for_deployment_status( deployment_name, desired_status, workspace=workspace, timeout=timeout - ) - if not deployment_status: + ): return False - - # Verify gateway can route to the provider if desired_status == "READY" and check_gateway: - gateway_ready = self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) - if not gateway_ready: - return False - + return self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) return True - def wait_for_gateway( + def wait_for_deployment_status( self, - provider_name: str, + deployment_name: str, + desired_status: str, *, workspace: str | None = None, - timeout: int = 60, + timeout: int = 1200, ) -> bool: - """ - Wait for the inference gateway to be able to route to a provider. - - Polls the gateway's /ready endpoint until it returns success, indicating - the gateway has refreshed its cache and is aware of the provider. - - Args: - provider_name: Name of the model provider - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - - Returns: - True if gateway is ready, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - - print("Waiting for gateway to be ready...") - - while time.time() - start_time < timeout: - try: - self._client.inference.gateway.provider.ready( - provider_name, - workspace=workspace, - ) - timestamp = datetime.now().strftime("%H:%M:%S") - print(f" [{timestamp}] Gateway is ready!\n") - return True - except NotFoundError: - # Gateway doesn't know about the provider yet, keep waiting - time.sleep(1) - except Exception: - # Connection error or other issue, keep waiting - time.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"Gateway timeout after {elapsed}s\n") - return False + """Wait for a ModelDeployment to reach ``desired_status`` (or 404 for ``DELETED``).""" + return self.models.wait_for_deployment_status( + deployment_name, desired_status, workspace=workspace, timeout=timeout + ) def wait_for_provider( self, @@ -288,300 +160,93 @@ def wait_for_provider( timeout: int = 60, check_gateway: bool = True, ) -> bool: - """ - Wait for a provider to reach the desired status. - - This is useful for external providers (like NVIDIA Build or OpenAI) where - you need to wait for the provider to be ready before making inference calls. - - Args: - provider_name: Name of the provider - desired_status: Target status (default: "READY") - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", also verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - last_status = "" - - print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") - - while time.time() - start_time < timeout: - try: - provider = self._client.inference.providers.retrieve( - provider_name, - workspace=workspace, - ) - current_status = provider.status - - if current_status != last_status: - timestamp = datetime.now().strftime("%H:%M:%S") - elapsed = int(time.time() - start_time) - print(f" [{timestamp}] ({elapsed}s) Status: {current_status}") - last_status = current_status - - if current_status == desired_status: - if desired_status == "READY" and check_gateway: - return self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) - print() - return True - - if current_status == "ERROR": - print(f"\nProvider entered ERROR state: {provider.status_message}\n") - return False - - except NotFoundError: - print(f"\nProvider '{provider_name}' not found\n") - return False - - time.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"\nProvider timeout after {elapsed}s. Last status: {last_status}\n") - return False + """Wait for a ModelProvider to reach ``desired_status`` (optionally gateway-ready).""" + if not self.models.wait_for_provider_status( + provider_name, desired_status, workspace=workspace, timeout=timeout + ): + return False + if desired_status == "READY" and check_gateway: + return self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) + return True - def wait_for_deployment_status( + def wait_for_gateway( self, - deployment_name: str, - desired_status: str, + provider_name: str, *, workspace: str | None = None, - timeout: int = 1200, + timeout: int = 60, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. + """Wait for the inference gateway to be able to route to a provider. - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - - Returns: - True if desired status reached, False if timeout + Targets the separate inference-gateway service, which is not yet migrated + to a typed client, so this still calls through the ``NeMoPlatform`` SDK. """ if workspace is None: workspace = self._client._get_workspace_path_param() start_time = time.time() - last_status = "" - last_message = "" - last_history_len = 0 - - print(f"Waiting for status: {desired_status}...\n") - + print("Waiting for gateway to be ready...") while time.time() - start_time < timeout: try: - deployment = self._client.inference.deployments.retrieve(deployment_name, workspace=workspace) - history = getattr(deployment, "status_history", None) - created_at = getattr(deployment, "created_at", None) - # API guarantees last history entry is current state; fall back to top-level fields if no history - if history and len(history) > 0: - last_entry = history[-1] - current_status = getattr(last_entry, "status", deployment.status) - status_message = getattr(last_entry, "status_message", "") or "" - else: - current_status = deployment.status - status_message = deployment.status_message or "" - last_status = current_status - last_message = status_message - - # Only print status from history; elapsed shown is seconds since deployment creation - if history and len(history) > last_history_len: - for i in range(last_history_len, len(history)): - entry = history[i] - ts = getattr(entry, "timestamp", None) - ts_str = ts.strftime("%H:%M:%S") if hasattr(ts, "strftime") else str(ts) if ts else "" - st = getattr(entry, "status", "") - msg = getattr(entry, "status_message", "") or "" - secs = _seconds_since_creation(ts, created_at) - part = f" [{ts_str}] " - if secs is not None: - part += f"(+{secs}s) " - part += f"Status: {st}" - if msg: - part += f" - {msg}" - print(part) - last_history_len = len(history) - - # Check if we've reached the desired status - # For DELETED status, we need to wait for the actual 404 (garbage collection) - if current_status == desired_status and desired_status != "DELETED": - print(f"Deployment reached {desired_status} status!\n") - return True - - # Handle error states - if current_status == "ERROR": - print(f"Deployment entered ERROR state: {status_message}\n") - return False - + self._client.inference.gateway.provider.ready(provider_name, workspace=workspace) + print(f" [{datetime.now().strftime('%H:%M:%S')}] Gateway is ready!\n") + return True except NotFoundError: - # For DELETED status, not found means success - if desired_status == "DELETED": - print(f"Deployment {desired_status}!\n") - return True - # For other statuses, not found is an error - print("Deployment not found\n") - return False - - time.sleep(3) - - # Timeout reached (wait_elapsed is time since we started polling) - wait_elapsed = int(time.time() - start_time) - detail = f"Last status: {last_status}" - if last_message: - detail += f" - {last_message}" - print(f"Timeout after {wait_elapsed}s. {detail}\n") + time.sleep(1) + except Exception: + time.sleep(1) + print(f"Gateway timeout after {int(time.time() - start_time)}s\n") return False -class AsyncModelsResource(BaseAsyncModelsResource): - """Extended AsyncModelsResource with high-level helper methods. +class AsyncModelsResource: + """Async twin of :class:`ModelsResource`. - All existing async methods (create, retrieve, list, etc.) work unchanged. - Adds convenience methods for OpenAI integration and deployment management. + Route builders are synchronous (no I/O) and safe to call from async code; + methods that perform I/O are async. + """ - URL builder methods are synchronous (no I/O) and safe to call from async code. - Methods that perform I/O are properly async. + def __init__(self, client: AsyncNeMoPlatform) -> None: + self._client = client + self._typed: AsyncModelsClient | None = None - Example: - >>> sdk = AsyncNeMoPlatform(base_url="http://nmp-host", workspace="default") - >>> sdk.models.get_openai_route_base_url() - >>> await sdk.models.wait_for_status("my-deployment", "READY") - """ + @property + def models(self) -> AsyncModelsClient: + """The typed async Models client sharing this SDK's transport (built lazily).""" + if self._typed is None: + self._typed = client_from_platform(self._client, AsyncModelsClient) + return self._typed def _get_base_url_str(self) -> str: """Get the base URL as a string with trailing slash removed.""" return str(self._client.base_url).rstrip("/") def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: - """ - Generate the base URL for the OpenAI proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - workspace: The workspace identifier - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - base_url = self._get_base_url_str() - return f"{base_url}/apis/inference-gateway/v2/workspaces/{workspace}/openai/-/v1" + """Base URL for the OpenAI proxy route. Synchronous (no I/O).""" + return self.models.get_openai_route_base_url(workspace=workspace) def get_client_default_headers(self) -> dict[str, str]: - """Get string-only default headers for third-party client libraries. - - Use this helper when constructing external clients (for example OpenAI - SDK or LiteLLM) so auth and identity headers from the SDK are forwarded. - This is required for successful inference requests when platform auth/ - authorization is enabled. - """ + """String-only default headers for third-party client libraries.""" return {key: value for key, value in self._client.default_headers.items() if isinstance(value, str)} def get_async_openai_client(self, *, workspace: str | None = None): - """ - Get an async OpenAI client configured for NeMo Platform's inference gateway. - - This method returns an AsyncOpenAI client with the base_url set to the - OpenAI proxy route for the specified workspace. - - Args: - workspace: The workspace identifier - - Returns: - An AsyncOpenAI client configured for the inference gateway - - Example: - >>> client = sdk.models.get_async_openai_client() - >>> response = await client.chat.completions.create( - ... model="default/meta_llama-3.2-1b-instruct", - ... messages=[{"role": "user", "content": "Hello!"}] - ... ) - """ + """An async OpenAI client configured for NeMo Platform's inference gateway.""" import openai base_url = self.get_openai_route_base_url(workspace=workspace) - # Preserve auth and identity headers from the parent SDK client. default_headers = self.get_client_default_headers() return openai.AsyncOpenAI(base_url=base_url, api_key="not-needed", default_headers=default_headers) def get_provider_route_openai_url(self, provider: ModelProvider) -> str: - """ - Generate an OpenAI SDK-compatible URL for the provider proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - provider: The ModelProvider object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - base_url = self._get_base_url_str() - route = f"{base_url}/apis/inference-gateway/v2/workspaces/{provider.workspace}/provider/{provider.name}/-" - - host_url_normalized = provider.host_url.rstrip("/") - if not host_url_normalized.endswith("/v1"): - route = f"{route}/v1" - - return route + """OpenAI SDK-compatible URL for a provider proxy route. Synchronous (no I/O).""" + return self.models.get_provider_route_openai_url(provider) async def get_provider_route_openai_url_for_deployment(self, deployment: ModelDeployment) -> str: - """ - Generate an OpenAI SDK-compatible URL for a deployment's model provider. - - This is an async method that fetches the ModelProvider associated - with the deployment and returns the provider route URL. - - Args: - deployment: The ModelDeployment object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Raises: - ValueError: If the deployment has no associated model_provider_id - - Example: - >>> deployment = await sdk.inference.deployments.retrieve("my-deployment", workspace="default") - >>> base_url = await sdk.models.get_provider_route_openai_url_for_deployment(deployment) - >>> openai_client = AsyncOpenAI(base_url=base_url) - """ - if not deployment.model_provider_id: - raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") - - # model_provider_id is in "workspace/name" format - workspace, name = deployment.model_provider_id.split("/", 1) - provider = await self._client.inference.providers.retrieve(name, workspace=workspace) - return self.get_provider_route_openai_url(provider) + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + return await self.models.get_provider_route_openai_url_for_deployment(deployment) def get_model_entity_route_openai_url(self, model_entity: ModelEntity) -> str: - """ - Generate an OpenAI SDK-compatible URL for the model entity proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - model_entity: The ModelEntity object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - base_url = self._get_base_url_str() - return ( - f"{base_url}/apis/inference-gateway/v2/workspaces/{model_entity.workspace}/model/{model_entity.name}/-/v1" - ) + """OpenAI SDK-compatible URL for a model-entity proxy route. Synchronous (no I/O).""" + return self.models.get_model_entity_route_openai_url(model_entity) async def wait_for_status( self, @@ -592,85 +257,27 @@ async def wait_for_status( timeout: int = 1200, check_gateway: bool = True, ) -> bool: - """ - Wait for a ModelDeployment and ModelProvider to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. - - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - deployment_status = await self.wait_for_deployment_status( + """Wait for a ModelDeployment to reach ``desired_status`` (optionally gateway-ready).""" + if not await self.models.wait_for_deployment_status( deployment_name, desired_status, workspace=workspace, timeout=timeout - ) - if not deployment_status: + ): return False - - # Verify gateway can route to the provider if desired_status == "READY" and check_gateway: - gateway_ready = await self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) - if not gateway_ready: - return False - + return await self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) return True - async def wait_for_gateway( + async def wait_for_deployment_status( self, - provider_name: str, + deployment_name: str, + desired_status: str, *, workspace: str | None = None, - timeout: int = 60, + timeout: int = 1200, ) -> bool: - """ - Wait for the inference gateway to be able to route to a provider. - - Polls the gateway's /ready endpoint until it returns success, indicating - the gateway has refreshed its cache and is aware of the provider. - - Args: - provider_name: Name of the model provider - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - - Returns: - True if gateway is ready, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - - print("Waiting for gateway to be ready...") - - while time.time() - start_time < timeout: - try: - await self._client.inference.gateway.provider.ready( - provider_name, - workspace=workspace, - ) - timestamp = datetime.now().strftime("%H:%M:%S") - print(f" [{timestamp}] Gateway is ready!\n") - return True - except NotFoundError: - # Gateway doesn't know about the provider yet, keep waiting - await asyncio.sleep(1) - except Exception: - # Connection error or other issue, keep waiting - await asyncio.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"Gateway timeout after {elapsed}s\n") - return False + """Wait for a ModelDeployment to reach ``desired_status`` (or 404 for ``DELETED``).""" + return await self.models.wait_for_deployment_status( + deployment_name, desired_status, workspace=workspace, timeout=timeout + ) async def wait_for_provider( self, @@ -681,156 +288,41 @@ async def wait_for_provider( timeout: int = 60, check_gateway: bool = True, ) -> bool: - """ - Wait for a provider to reach the desired status (async version). - - This is useful for external providers (like NVIDIA Build or OpenAI) where - you need to wait for the provider to be ready before making inference calls. - - Args: - provider_name: Name of the provider - desired_status: Target status (default: "READY") - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", also verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - last_status = "" - - print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") - - while time.time() - start_time < timeout: - try: - provider = await self._client.inference.providers.retrieve( - provider_name, - workspace=workspace, - ) - current_status = provider.status - - if current_status != last_status: - timestamp = datetime.now().strftime("%H:%M:%S") - elapsed = int(time.time() - start_time) - print(f" [{timestamp}] ({elapsed}s) Status: {current_status}") - last_status = current_status - - if current_status == desired_status: - if desired_status == "READY" and check_gateway: - return await self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) - print() - return True - - if current_status == "ERROR": - print(f"\nProvider entered ERROR state: {provider.status_message}\n") - return False - - except NotFoundError: - print(f"\nProvider '{provider_name}' not found\n") - return False - - await asyncio.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"\nProvider timeout after {elapsed}s. Last status: {last_status}\n") - return False + """Wait for a ModelProvider to reach ``desired_status`` (optionally gateway-ready).""" + if not await self.models.wait_for_provider_status( + provider_name, desired_status, workspace=workspace, timeout=timeout + ): + return False + if desired_status == "READY" and check_gateway: + return await self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) + return True - async def wait_for_deployment_status( + async def wait_for_gateway( self, - deployment_name: str, - desired_status: str, + provider_name: str, *, workspace: str | None = None, - timeout: int = 1200, + timeout: int = 60, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status (async version). - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. + """Wait for the inference gateway to be able to route to a provider. - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - - Returns: - True if desired status reached, False if timeout + Targets the separate inference-gateway service (not yet migrated to a + typed client), so this still calls through the ``AsyncNeMoPlatform`` SDK. """ + import asyncio + if workspace is None: workspace = self._client._get_workspace_path_param() start_time = time.time() - last_status = "" - last_message = "" - last_history_len = 0 - - print(f"Waiting for status: {desired_status}...\n") - + print("Waiting for gateway to be ready...") while time.time() - start_time < timeout: try: - deployment = await self._client.inference.deployments.retrieve(deployment_name, workspace=workspace) - history = getattr(deployment, "status_history", None) - created_at = getattr(deployment, "created_at", None) - # API guarantees last history entry is current state; fall back to top-level fields if no history - if history and len(history) > 0: - last_entry = history[-1] - current_status = getattr(last_entry, "status", deployment.status) - status_message = getattr(last_entry, "status_message", "") or "" - else: - current_status = deployment.status - status_message = deployment.status_message or "" - last_status = current_status - last_message = status_message - - # Only print status from history; elapsed shown is seconds since deployment creation - if history and len(history) > last_history_len: - for i in range(last_history_len, len(history)): - entry = history[i] - ts = getattr(entry, "timestamp", None) - ts_str = ts.strftime("%H:%M:%S") if hasattr(ts, "strftime") else str(ts) if ts else "" - st = getattr(entry, "status", "") - msg = getattr(entry, "status_message", "") or "" - secs = _seconds_since_creation(ts, created_at) - part = f" [{ts_str}] " - if secs is not None: - part += f"(+{secs}s) " - part += f"Status: {st}" - if msg: - part += f" - {msg}" - print(part) - last_history_len = len(history) - - # Check if we've reached the desired status - # For DELETED status, we need to wait for the actual 404 (garbage collection) - if current_status == desired_status and desired_status != "DELETED": - print(f"Deployment reached {desired_status} status!\n") - return True - - # Handle error states - if current_status == "ERROR": - print(f"Deployment entered ERROR state: {status_message}\n") - return False - + await self._client.inference.gateway.provider.ready(provider_name, workspace=workspace) + print(f" [{datetime.now().strftime('%H:%M:%S')}] Gateway is ready!\n") + return True except NotFoundError: - # For DELETED status, not found means success - if desired_status == "DELETED": - print(f"Deployment {desired_status}!\n") - return True - # For other statuses, not found is an error - print("Deployment not found\n") - return False - - await asyncio.sleep(3) - - # Timeout reached (wait_elapsed is time since we started polling) - wait_elapsed = int(time.time() - start_time) - detail = f"Last status: {last_status}" - if last_message: - detail += f" - {last_message}" - print(f"Timeout after {wait_elapsed}s. {detail}\n") + await asyncio.sleep(1) + except Exception: + await asyncio.sleep(1) + print(f"Gateway timeout after {int(time.time() - start_time)}s\n") return False diff --git a/packages/models/tests/test_client.py b/packages/models/tests/test_client.py index 9dcb341b77..28d62d001f 100644 --- a/packages/models/tests/test_client.py +++ b/packages/models/tests/test_client.py @@ -1,556 +1,291 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -from unittest.mock import AsyncMock, MagicMock, patch +"""Tests for the extended ModelsResource / AsyncModelsResource helpers. +These exercise the *source* package (``models.resources``) which drives the +typed ``ModelsClient`` built from the SDK. The route builders are pure (no I/O); +the deployment/provider helpers are driven through a mocked httpx transport +shared by the SDK and the typed client (via ``client_from_platform``). +""" + +from unittest.mock import MagicMock, patch + +import httpx import pytest +from models.resources import AsyncModelsResource, ModelsResource from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform.types.inference import ModelDeployment, ModelProvider from nemo_platform.types.models import ModelEntity -# ============================================================================ -# Fixtures -# ============================================================================ - - -@pytest.fixture -def sdk(): - """Create a real NeMoPlatform SDK instance for testing.""" - return NeMoPlatform(base_url="https://nmp.example.com/") - - -@pytest.fixture -def sdk_with_workspace(): - """Create SDK with client-level workspace set.""" - return NeMoPlatform(base_url="https://nmp.example.com/", workspace="client-ws") - - -@pytest.fixture -def sdk_no_trailing_slash(): - """Create SDK with base_url without trailing slash.""" - return NeMoPlatform(base_url="https://nmp.example.com") - - -@pytest.fixture -def async_sdk(): - """Create a real AsyncNeMoPlatform SDK instance for testing.""" - return AsyncNeMoPlatform(base_url="https://nmp.example.com/") - - -@pytest.fixture -def async_sdk_with_workspace(): - """Create async SDK with client-level workspace set.""" - return AsyncNeMoPlatform(base_url="https://nmp.example.com/", workspace="client-ws") - - -# ============================================================================ -# ModelsResource Tests -# ============================================================================ - - -# Tests for _get_base_url_str - - -def test_get_base_url_str_removes_trailing_slash(sdk): - """Test that trailing slash is removed from base URL.""" - result = sdk.models._get_base_url_str() - - assert result == "https://nmp.example.com" - assert not result.endswith("/") - - -def test_get_base_url_str_handles_no_trailing_slash(sdk_no_trailing_slash): - """Test that URLs without trailing slash are unchanged.""" - result = sdk_no_trailing_slash.models._get_base_url_str() - - assert result == "https://nmp.example.com" - - -# Tests for get_openai_route_base_url +def _resource(base_url: str = "https://nmp.example.com/", workspace: str | None = None, **kwargs) -> ModelsResource: + return ModelsResource(NeMoPlatform(base_url=base_url, workspace=workspace, **kwargs)) -def test_get_openai_route_base_url_explicit_workspace(sdk): - """Test URL generation with explicit workspace.""" - result = sdk.models.get_openai_route_base_url(workspace="default") - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" +def _async_resource( + base_url: str = "https://nmp.example.com/", workspace: str | None = None, **kwargs +) -> AsyncModelsResource: + return AsyncModelsResource(AsyncNeMoPlatform(base_url=base_url, workspace=workspace, **kwargs)) -def test_get_openai_route_base_url_custom_workspace(sdk): - """Test URL generation with custom workspace name.""" - result = sdk.models.get_openai_route_base_url(workspace="my-workspace") +# --------------------------------------------------------------------------- +# Base URL / route builders +# --------------------------------------------------------------------------- - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/my-workspace/openai/-/v1" +def test_get_base_url_str_removes_trailing_slash() -> None: + assert _resource("https://nmp.example.com/")._get_base_url_str() == "https://nmp.example.com" + assert _resource("https://nmp.example.com")._get_base_url_str() == "https://nmp.example.com" -def test_get_openai_route_base_url_with_trailing_slash(sdk): - """Test that trailing slash in base_url is handled correctly.""" - result = sdk.models.get_openai_route_base_url(workspace="default") - assert "//v2" not in result - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - - -def test_get_openai_route_base_url_without_trailing_slash(sdk_no_trailing_slash): - """Test URL generation when base_url has no trailing slash.""" - result = sdk_no_trailing_slash.models.get_openai_route_base_url(workspace="default") - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" +def test_get_openai_route_base_url_explicit_workspace() -> None: + r = _resource() + assert ( + r.get_openai_route_base_url(workspace="default") + == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" + ) + assert "//v2" not in r.get_openai_route_base_url(workspace="default") -# Tests for get_provider_route_openai_url +def test_get_openai_route_base_url_without_trailing_slash() -> None: + r = _resource("https://nmp.example.com") + assert ( + r.get_openai_route_base_url(workspace="default") + == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" + ) -def test_get_provider_route_openai_url_appends_v1(sdk): - """Test that /v1 is appended when provider host_url doesn't end with /v1.""" +def test_provider_route_openai_url_appends_v1() -> None: provider = MagicMock(spec=ModelProvider) provider.workspace = "default" provider.name = "openai-provider" provider.host_url = "https://api.openai.com" - - result = sdk.models.get_provider_route_openai_url(provider) - assert ( - result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/openai-provider/-/v1" + _resource().get_provider_route_openai_url(provider) + == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/openai-provider/-/v1" ) -def test_get_provider_route_openai_url_no_v1_when_host_ends_with_v1(sdk): - """Test that /v1 is NOT appended when provider host_url already ends with /v1.""" - provider = MagicMock(spec=ModelProvider) - provider.workspace = "default" - provider.name = "nim-provider" - provider.host_url = "https://nim.example.com/v1" - - result = sdk.models.get_provider_route_openai_url(provider) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/nim-provider/-" - - -def test_get_provider_route_openai_url_no_v1_when_host_ends_with_v1_slash(sdk): - """Test that /v1 is NOT appended when provider host_url ends with /v1/.""" - provider = MagicMock(spec=ModelProvider) - provider.workspace = "default" - provider.name = "nim-provider" - provider.host_url = "https://nim.example.com/v1/" - - result = sdk.models.get_provider_route_openai_url(provider) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/nim-provider/-" - - -def test_get_provider_route_openai_url_custom_workspace(sdk): - """Test provider URL generation with custom workspace.""" - provider = MagicMock(spec=ModelProvider) - provider.workspace = "production" - provider.name = "my-provider" - provider.host_url = "https://api.example.com" - - result = sdk.models.get_provider_route_openai_url(provider) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/production/provider/my-provider/-/v1" - - -# Tests for get_model_entity_route_openai_url - - -def test_get_model_entity_route_openai_url_default_workspace(sdk): - """Test URL generation for model entity.""" - model_entity = MagicMock(spec=ModelEntity) - model_entity.workspace = "default" - model_entity.name = "llama3-70b-instruct" - - result = sdk.models.get_model_entity_route_openai_url(model_entity) - - assert ( - result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/model/llama3-70b-instruct/-/v1" - ) +def test_provider_route_openai_url_no_v1_when_host_ends_with_v1() -> None: + for host in ("https://nim.example.com/v1", "https://nim.example.com/v1/"): + provider = MagicMock(spec=ModelProvider) + provider.workspace = "default" + provider.name = "nim-provider" + provider.host_url = host + assert ( + _resource().get_provider_route_openai_url(provider) + == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/nim-provider/-" + ) -def test_get_model_entity_route_openai_url_custom_workspace(sdk): - """Test URL generation for model entity with custom workspace.""" +def test_model_entity_route_openai_url_always_v1() -> None: model_entity = MagicMock(spec=ModelEntity) model_entity.workspace = "ml-team" model_entity.name = "custom-model" - - result = sdk.models.get_model_entity_route_openai_url(model_entity) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/ml-team/model/custom-model/-/v1" - - -def test_get_model_entity_route_openai_url_always_appends_v1(sdk): - """Test that /v1 is always appended for model entity routes.""" - model_entity = MagicMock(spec=ModelEntity) - model_entity.workspace = "default" - model_entity.name = "test-model" - - result = sdk.models.get_model_entity_route_openai_url(model_entity) - - assert result.endswith("/v1") - - -# Tests for get_provider_route_openai_url_for_deployment - - -def test_get_provider_route_openai_url_for_deployment_fetches_provider(sdk): - """Test that provider is fetched and URL is generated correctly.""" - mock_provider = MagicMock(spec=ModelProvider) - mock_provider.workspace = "default" - mock_provider.name = "my-provider" - mock_provider.host_url = "https://api.example.com" - - deployment = MagicMock(spec=ModelDeployment) - deployment.name = "my-deployment" - deployment.model_provider_id = "default/my-provider" - - # Mock the provider retrieval - with patch.object(sdk.inference.providers, "retrieve", return_value=mock_provider) as mock_retrieve: - result = sdk.models.get_provider_route_openai_url_for_deployment(deployment) - - mock_retrieve.assert_called_once_with("my-provider", workspace="default") - assert ( - result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1" - ) - - -def test_get_provider_route_openai_url_for_deployment_respects_v1_suffix(sdk): - """Test that /v1 is not appended when provider host_url ends with /v1.""" - mock_provider = MagicMock(spec=ModelProvider) - mock_provider.workspace = "production" - mock_provider.name = "nim-provider" - mock_provider.host_url = "https://nim.example.com/v1" - - deployment = MagicMock(spec=ModelDeployment) - deployment.name = "nim-deployment" - deployment.model_provider_id = "production/nim-provider" - - with patch.object(sdk.inference.providers, "retrieve", return_value=mock_provider): - result = sdk.models.get_provider_route_openai_url_for_deployment(deployment) - - assert ( - result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/production/provider/nim-provider/-" - ) + assert ( + _resource().get_model_entity_route_openai_url(model_entity) + == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/ml-team/model/custom-model/-/v1" + ) -def test_get_provider_route_openai_url_for_deployment_raises_when_no_provider_id(sdk): - """Test that ValueError is raised when deployment has no model_provider_id.""" - deployment = MagicMock(spec=ModelDeployment) - deployment.name = "orphan-deployment" - deployment.model_provider_id = None +# --------------------------------------------------------------------------- +# Workspace resolution (client-level fallback) +# --------------------------------------------------------------------------- - with pytest.raises(ValueError) as exc_info: - sdk.models.get_provider_route_openai_url_for_deployment(deployment) - assert "orphan-deployment" in str(exc_info.value) - assert "no associated model_provider_id" in str(exc_info.value) +def test_openai_route_base_url_uses_client_workspace() -> None: + assert "/workspaces/client-ws/" in _resource(workspace="client-ws").get_openai_route_base_url() -def test_get_provider_route_openai_url_for_deployment_raises_when_empty_provider_id(sdk): - """Test that ValueError is raised when deployment has empty model_provider_id.""" - deployment = MagicMock(spec=ModelDeployment) - deployment.name = "empty-provider-deployment" - deployment.model_provider_id = "" +def test_openai_route_base_url_explicit_overrides_client() -> None: + r = _resource(workspace="client-ws") + result = r.get_openai_route_base_url(workspace="override") + assert "/workspaces/override/" in result + assert "/workspaces/client-ws/" not in result - with pytest.raises(ValueError) as exc_info: - sdk.models.get_provider_route_openai_url_for_deployment(deployment) - assert "empty-provider-deployment" in str(exc_info.value) +def test_openai_route_base_url_raises_without_workspace() -> None: + with pytest.raises(ValueError, match="Missing workspace"): + _resource().get_openai_route_base_url() -# Tests for get_openai_client +# --------------------------------------------------------------------------- +# OpenAI client factories +# --------------------------------------------------------------------------- -def test_get_openai_client_returns_configured_client(sdk): - """Test that get_openai_client returns an OpenAI client with correct base_url.""" +def test_get_openai_client_returns_configured_client() -> None: + r = _resource() with patch("openai.OpenAI") as mock_openai_cls: mock_openai_cls.return_value = MagicMock() - - sdk.models.get_openai_client(workspace="default") - expected_headers = sdk.models.get_client_default_headers() - + r.get_openai_client(workspace="default") mock_openai_cls.assert_called_once_with( base_url="https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1", api_key="not-needed", - default_headers=expected_headers, + default_headers=r.get_client_default_headers(), ) -def test_get_openai_client_custom_workspace(sdk): - """Test get_openai_client with custom workspace.""" - with patch("openai.OpenAI") as mock_openai_cls: - mock_openai_cls.return_value = MagicMock() - - sdk.models.get_openai_client(workspace="production") - expected_headers = sdk.models.get_client_default_headers() - - mock_openai_cls.assert_called_once_with( - base_url="https://nmp.example.com/apis/inference-gateway/v2/workspaces/production/openai/-/v1", - api_key="not-needed", - default_headers=expected_headers, +def test_get_openai_client_includes_auth_headers() -> None: + r = ModelsResource( + NeMoPlatform( + base_url="https://nmp.example.com/", + default_headers={"Authorization": "Bearer token-123", "X-NMP-Principal-Id": "user@example.com"}, ) - - -def test_get_openai_client_includes_auth_headers(): - """Test that auth headers from the SDK are propagated to OpenAI client.""" - sdk_with_auth = NeMoPlatform( - base_url="https://nmp.example.com/", - default_headers={"Authorization": "Bearer token-123", "X-NMP-Principal-Id": "user@example.com"}, ) - with patch("openai.OpenAI") as mock_openai_cls: mock_openai_cls.return_value = MagicMock() - - sdk_with_auth.models.get_openai_client(workspace="default") - default_headers = mock_openai_cls.call_args.kwargs["default_headers"] - - assert default_headers["Authorization"] == "Bearer token-123" - assert default_headers["X-NMP-Principal-Id"] == "user@example.com" - - -# ============================================================================ -# AsyncModelsResource Tests -# ============================================================================ - - -# Tests for AsyncModelsResource URL builders (sync methods, no I/O) - - -def test_async_get_base_url_str_removes_trailing_slash(async_sdk): - """Test that trailing slash is removed from base URL (async resource).""" - result = async_sdk.models._get_base_url_str() - - assert result == "https://nmp.example.com" - - -def test_async_get_openai_route_base_url(async_sdk): - """Test URL generation with async resource (sync method, no I/O).""" - result = async_sdk.models.get_openai_route_base_url(workspace="default") - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - - -def test_async_get_provider_route_openai_url(async_sdk): - """Test provider URL generation with async resource (sync method, no I/O).""" - provider = MagicMock(spec=ModelProvider) - provider.workspace = "default" - provider.name = "my-provider" - provider.host_url = "https://api.example.com" - - result = async_sdk.models.get_provider_route_openai_url(provider) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1" - - -def test_async_get_model_entity_route_openai_url(async_sdk): - """Test model entity URL generation with async resource (sync method, no I/O).""" - model_entity = MagicMock(spec=ModelEntity) - model_entity.workspace = "default" - model_entity.name = "my-model" - - result = async_sdk.models.get_model_entity_route_openai_url(model_entity) - - assert result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/model/my-model/-/v1" - - -# Tests for get_async_openai_client - - -def test_get_async_openai_client_returns_async_client(async_sdk): - """Test that get_async_openai_client returns an AsyncOpenAI client.""" - with patch("openai.AsyncOpenAI") as mock_async_openai_cls: - mock_async_openai_cls.return_value = MagicMock() - - async_sdk.models.get_async_openai_client(workspace="default") - expected_headers = async_sdk.models.get_client_default_headers() - - mock_async_openai_cls.assert_called_once_with( - base_url="https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/openai/-/v1", - api_key="not-needed", - default_headers=expected_headers, - ) - - -def test_get_async_openai_client_custom_workspace(async_sdk): - """Test get_async_openai_client with custom workspace.""" - with patch("openai.AsyncOpenAI") as mock_async_openai_cls: - mock_async_openai_cls.return_value = MagicMock() - - async_sdk.models.get_async_openai_client(workspace="production") - expected_headers = async_sdk.models.get_client_default_headers() - - mock_async_openai_cls.assert_called_once_with( + r.get_openai_client(workspace="default") + headers = mock_openai_cls.call_args.kwargs["default_headers"] + assert headers["Authorization"] == "Bearer token-123" + assert headers["X-NMP-Principal-Id"] == "user@example.com" + + +def test_get_async_openai_client_returns_async_client() -> None: + r = _async_resource() + with patch("openai.AsyncOpenAI") as mock_cls: + mock_cls.return_value = MagicMock() + r.get_async_openai_client(workspace="production") + mock_cls.assert_called_once_with( base_url="https://nmp.example.com/apis/inference-gateway/v2/workspaces/production/openai/-/v1", api_key="not-needed", - default_headers=expected_headers, + default_headers=r.get_client_default_headers(), ) -def test_get_async_openai_client_includes_auth_headers(): - """Test that auth headers from the async SDK are propagated to AsyncOpenAI client.""" - sdk_with_auth = AsyncNeMoPlatform( - base_url="https://nmp.example.com/", - default_headers={"Authorization": "Bearer token-abc", "X-NMP-Principal-Id": "async-user@example.com"}, - ) +# --------------------------------------------------------------------------- +# get_provider_route_openai_url_for_deployment (drives the typed client) +# +# A real httpx client with a MockTransport is used so the client_from_platform +# bridge (which copies the SDK's transport + headers) works end to end. +# --------------------------------------------------------------------------- - with patch("openai.AsyncOpenAI") as mock_async_openai_cls: - mock_async_openai_cls.return_value = MagicMock() - sdk_with_auth.models.get_async_openai_client(workspace="default") - default_headers = mock_async_openai_cls.call_args.kwargs["default_headers"] +def _capturing_transport(response: httpx.Response) -> tuple[httpx.MockTransport, list[httpx.Request]]: + seen: list[httpx.Request] = [] - assert default_headers["Authorization"] == "Bearer token-abc" - assert default_headers["X-NMP-Principal-Id"] == "async-user@example.com" + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return response + return httpx.MockTransport(handler), seen -# Tests for async get_provider_route_openai_url_for_deployment +def _provider_payload(**extra: object) -> dict: + base = { + "id": "provider-1", + "name": "my-provider", + "workspace": "default", + "host_url": "https://api.example.com", + "created_at": "2020-01-01T00:00:00Z", + "updated_at": "2020-01-01T00:00:00Z", + } + base.update(extra) + return base -@pytest.mark.asyncio -async def test_async_get_provider_route_openai_url_for_deployment(async_sdk): - """Test async version fetches provider and generates URL correctly.""" - mock_provider = MagicMock(spec=ModelProvider) - mock_provider.workspace = "default" - mock_provider.name = "my-provider" - mock_provider.host_url = "https://api.example.com" + +def test_provider_route_for_deployment_fetches_provider() -> None: + transport, seen = _capturing_transport(httpx.Response(200, json=_provider_payload())) + r = _resource(http_client=httpx.Client(transport=transport)) deployment = MagicMock(spec=ModelDeployment) deployment.name = "my-deployment" deployment.model_provider_id = "default/my-provider" - with patch.object( - async_sdk.inference.providers, "retrieve", AsyncMock(return_value=mock_provider) - ) as mock_retrieve: - result = await async_sdk.models.get_provider_route_openai_url_for_deployment(deployment) - - mock_retrieve.assert_called_once_with("my-provider", workspace="default") - assert ( - result == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1" - ) + url = r.get_provider_route_openai_url_for_deployment(deployment) + assert url == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1" + assert str(seen[0].url).endswith("/apis/models/v2/workspaces/default/providers/my-provider") -@pytest.mark.asyncio -async def test_async_get_provider_route_openai_url_for_deployment_raises_when_no_provider_id(async_sdk): - """Test that ValueError is raised when deployment has no model_provider_id (async).""" +def test_provider_route_for_deployment_raises_when_no_provider_id() -> None: deployment = MagicMock(spec=ModelDeployment) deployment.name = "orphan-deployment" deployment.model_provider_id = None + with pytest.raises(ValueError, match="no associated model_provider_id"): + _resource().get_provider_route_openai_url_for_deployment(deployment) - with pytest.raises(ValueError) as exc_info: - await async_sdk.models.get_provider_route_openai_url_for_deployment(deployment) - - assert "orphan-deployment" in str(exc_info.value) - assert "no associated model_provider_id" in str(exc_info.value) - - -# ============================================================================ -# Workspace Resolution Tests (client-level fallback) -# ============================================================================ - - -# Tests for get_openai_route_base_url workspace resolution - - -def test_get_openai_route_base_url_uses_client_workspace(sdk_with_workspace): - """Test that client-level workspace is used when not explicitly provided.""" - result = sdk_with_workspace.models.get_openai_route_base_url() - - assert "/workspaces/client-ws/" in result - - -def test_get_openai_route_base_url_explicit_overrides_client(sdk_with_workspace): - """Test that explicit workspace overrides client-level workspace.""" - result = sdk_with_workspace.models.get_openai_route_base_url(workspace="override") - - assert "/workspaces/override/" in result - assert "/workspaces/client-ws/" not in result - - -def test_get_openai_route_base_url_raises_without_workspace(sdk): - """Test that ValueError is raised when no workspace is available.""" - with pytest.raises(ValueError, match="Missing workspace"): - sdk.models.get_openai_route_base_url() - - -# Tests for get_openai_client workspace resolution - - -def test_get_openai_client_uses_client_workspace(sdk_with_workspace): - """Test that client-level workspace is used when not explicitly provided.""" - with patch("openai.OpenAI") as mock_openai_cls: - mock_openai_cls.return_value = MagicMock() - - sdk_with_workspace.models.get_openai_client() - - call_args = mock_openai_cls.call_args - assert "/workspaces/client-ws/" in call_args.kwargs["base_url"] - - -def test_get_openai_client_explicit_overrides_client(sdk_with_workspace): - """Test that explicit workspace overrides client-level workspace.""" - with patch("openai.OpenAI") as mock_openai_cls: - mock_openai_cls.return_value = MagicMock() - - sdk_with_workspace.models.get_openai_client(workspace="override") - - call_args = mock_openai_cls.call_args - assert "/workspaces/override/" in call_args.kwargs["base_url"] +@pytest.mark.asyncio +async def test_async_provider_route_for_deployment_fetches_provider() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_provider_payload())) + r = _async_resource(http_client=httpx.AsyncClient(transport=transport)) -def test_get_openai_client_raises_without_workspace(sdk): - """Test that ValueError is raised when no workspace is available.""" - with pytest.raises(ValueError, match="Missing workspace"): - sdk.models.get_openai_client() + deployment = MagicMock(spec=ModelDeployment) + deployment.name = "my-deployment" + deployment.model_provider_id = "default/my-provider" + url = await r.get_provider_route_openai_url_for_deployment(deployment) + assert url == "https://nmp.example.com/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1" -# Tests for async get_openai_route_base_url workspace resolution +# --------------------------------------------------------------------------- +# Deployment status polling (drives the typed client) +# --------------------------------------------------------------------------- -def test_async_get_openai_route_base_url_uses_client_workspace(async_sdk_with_workspace): - """Test that client-level workspace is used when not explicitly provided.""" - result = async_sdk_with_workspace.models.get_openai_route_base_url() - assert "/workspaces/client-ws/" in result +def _deployment_payload(status: str) -> dict: + return { + "id": "dep-1", + "name": "my-deploy", + "workspace": "default", + "entity_version": 1, + "config": "cfg", + "config_version": 1, + "status": status, + "status_message": "", + "status_history": [], + "created_at": "2020-01-01T00:00:00Z", + "updated_at": "2020-01-01T00:00:00Z", + } -def test_async_get_openai_route_base_url_explicit_overrides_client(async_sdk_with_workspace): - """Test that explicit workspace overrides client-level workspace.""" - result = async_sdk_with_workspace.models.get_openai_route_base_url(workspace="override") +def test_wait_for_status_ready_checks_gateway() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_deployment_payload("READY"))) + r = _resource(workspace="default", http_client=httpx.Client(transport=transport)) - assert "/workspaces/override/" in result + with patch.object(r, "wait_for_gateway", return_value=True) as gw: + assert r.wait_for_status("my-deploy", "READY") is True + gw.assert_called_once() -def test_async_get_openai_route_base_url_raises_without_workspace(async_sdk): - """Test that ValueError is raised when no workspace is available.""" - with pytest.raises(ValueError, match="Missing workspace"): - async_sdk.models.get_openai_route_base_url() +def test_wait_for_status_deployment_not_ready_skips_gateway() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_deployment_payload("ERROR"))) + r = _resource(workspace="default", http_client=httpx.Client(transport=transport)) + with patch.object(r, "wait_for_gateway", return_value=True) as gw: + assert r.wait_for_status("my-deploy", "READY") is False + gw.assert_not_called() -# Tests for get_async_openai_client workspace resolution +# --------------------------------------------------------------------------- +# Provider status polling (drives the typed client) +# --------------------------------------------------------------------------- -def test_get_async_openai_client_uses_client_workspace(async_sdk_with_workspace): - """Test that client-level workspace is used when not explicitly provided.""" - with patch("openai.AsyncOpenAI") as mock_openai_cls: - mock_openai_cls.return_value = MagicMock() - async_sdk_with_workspace.models.get_async_openai_client() +def test_wait_for_provider_ready_checks_gateway() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_provider_payload(status="READY"))) + r = _resource(workspace="default", http_client=httpx.Client(transport=transport)) - call_args = mock_openai_cls.call_args - assert "/workspaces/client-ws/" in call_args.kwargs["base_url"] + with patch.object(r, "wait_for_gateway", return_value=True) as gw: + assert r.wait_for_provider("my-provider", "READY") is True + gw.assert_called_once() -def test_get_async_openai_client_explicit_overrides_client(async_sdk_with_workspace): - """Test that explicit workspace overrides client-level workspace.""" - with patch("openai.AsyncOpenAI") as mock_openai_cls: - mock_openai_cls.return_value = MagicMock() +def test_wait_for_provider_error_skips_gateway() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_provider_payload(status="ERROR"))) + r = _resource(workspace="default", http_client=httpx.Client(transport=transport)) - async_sdk_with_workspace.models.get_async_openai_client(workspace="override") + with patch.object(r, "wait_for_gateway", return_value=True) as gw: + assert r.wait_for_provider("my-provider", "READY") is False + gw.assert_not_called() - call_args = mock_openai_cls.call_args - assert "/workspaces/override/" in call_args.kwargs["base_url"] +@pytest.mark.asyncio +async def test_async_wait_for_provider_error_skips_gateway() -> None: + transport, _ = _capturing_transport(httpx.Response(200, json=_provider_payload(status="ERROR"))) + r = _async_resource(workspace="default", http_client=httpx.AsyncClient(transport=transport)) -def test_get_async_openai_client_raises_without_workspace(async_sdk): - """Test that ValueError is raised when no workspace is available.""" - with pytest.raises(ValueError, match="Missing workspace"): - async_sdk.models.get_async_openai_client() + with patch.object(r, "wait_for_gateway", return_value=True) as gw: + assert await r.wait_for_provider("my-provider", "READY") is False + gw.assert_not_called() 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 d67180020d..dd177bd7dc 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 @@ -61,7 +61,13 @@ def client_from_platform( _skip = {"accept", "accept-encoding", "connection", "user-agent", "host"} headers = {k: v for k, v in platform._client.headers.items() if k.lower() not in _skip} # type: ignore[union-attr] - retry = RetryPolicy(max_retries=platform.max_retries) + retry = RetryPolicy( + max_retries=platform.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 = _url_resolver_from_platform(platform) if isinstance(platform, AsyncNeMoPlatform): if not issubclass(client_cls, AsyncNemoClient): 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 feb69df85c..33bf046849 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 @@ -18,6 +18,7 @@ import asyncio import copy +import email.utils import inspect import json import os @@ -117,6 +118,23 @@ def _get_paginated_types( # --------------------------------------------------------------------------- +def _retry_after(response: httpx.Response) -> float | None: + """Parse a reasonable server-requested retry delay in seconds.""" + retry_after_ms = response.headers.get("retry-after-ms") + try: + delay = float(retry_after_ms) / 1000 + except (TypeError, ValueError): + retry_after = response.headers.get("retry-after") + try: + delay = float(retry_after) + except (TypeError, ValueError): + retry_date = email.utils.parsedate_tz(retry_after) + if retry_date is None: + return None + delay = float(email.utils.mktime_tz(retry_date) - time.time()) + return delay if 0 < delay <= 60 else None + + def _should_retry( response: httpx.Response | None, exc: httpx.TransportError | None, @@ -129,14 +147,33 @@ def _should_retry( Returns the sleep duration if a retry should happen, or ``None`` if the response should be returned / the exception re-raised. """ - is_last = attempt >= policy.max_retries - if is_last: + if attempt >= policy.max_retries: return None + + backoff = policy.backoff_base * (2**attempt) if exc is not None: - return policy.backoff_base * (2**attempt) - if response is not None and response.status_code in policy.retryable_status_codes: - return policy.backoff_base * (2**attempt) - return None + return backoff + if response is None: + return None + + if policy.respect_retry_decision_headers: + if response.status_code < 400: + return None + should_retry = response.headers.get("x-should-retry") + if should_retry == "true": + return (_retry_after(response) or backoff) if policy.respect_retry_after_headers else backoff + if should_retry == "false": + return None + + retryable_status = response.status_code in policy.retryable_status_codes + if policy.retry_all_server_errors and response.status_code >= 500: + retryable_status = True + if not retryable_status: + return None + + if policy.respect_retry_after_headers: + return _retry_after(response) or backoff + return backoff def _should_resolve_conflict(response: httpx.Response, request: PreparedRequest) -> bool: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py index 58c9184feb..71e14ef12b 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/method.py @@ -71,16 +71,52 @@ class EndpointMethod(Generic[P, SyncReturnT, AsyncReturnT]): the response type that ``send()`` returns for each endpoint marker. """ + # Copied from the endpoint in __init__ so help() and autodoc describe the + # endpoint. Declared here so the descriptor's introspection surface is part of + # its type rather than something callers have to discover at runtime. + __wrapped__: Callable[P, PreparedRequest] + __name__: str + __qualname__: str + __doc__: str | None + __module__: str + def __init__(self, endpoint_fn: Callable[P, PreparedRequest]) -> None: self._endpoint_fn = endpoint_fn + # Carry the endpoint's name, docstring, and annotations onto the descriptor + # so help() and autodoc describe the endpoint rather than the descriptor. + # Set directly rather than via functools.update_wrapper, which expects a + # callable wrapper; a descriptor is not one, and which would also copy + # __dict__ and with it the endpoint's __isabstractmethod__ marker. + # + # This does NOT make inspect.signature(SomeClient.method) work: signature() + # rejects a non-callable before it ever consults __wrapped__. Reach the + # parameter list via inspect.unwrap() at class level, or just read it off + # an instance, where __get__ hands back the bound function. + self.__wrapped__ = endpoint_fn + for attr in functools.WRAPPER_ASSIGNMENTS: + try: + setattr(self, attr, getattr(endpoint_fn, attr)) + except AttributeError: + pass + + @property + def endpoint(self) -> Callable[P, PreparedRequest]: + """The endpoint function this descriptor binds.""" + return self._endpoint_fn + @overload + def __get__(self, obj: None, objtype: type | None = None) -> EndpointMethod[P, SyncReturnT, AsyncReturnT]: ... @overload def __get__(self, obj: NemoClient, objtype: type | None = None) -> Callable[P, SyncReturnT]: ... @overload def __get__(self, obj: AsyncNemoClient, objtype: type | None = None) -> Callable[P, Awaitable[AsyncReturnT]]: ... def __get__(self, obj: NemoClient | AsyncNemoClient | None, objtype: type | None = None) -> object: - assert obj is not None + if obj is None: + # Class-level access. Anything that inspects a client class rather than + # an instance -- Mock(spec=...), inspect, help(), autodoc -- lands here, + # and the descriptor protocol says to hand back the descriptor itself. + return self if isinstance(obj, AsyncNemoClient): @functools.wraps(self._endpoint_fn) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py index 56c6959220..627eef150e 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/response.py @@ -79,6 +79,10 @@ class NemoResponse(Generic[ResponseT]): resp.http_response # full httpx.Response user = resp.data() # raises on non-2xx, otherwise returns body + + When several 2xx codes share one typed body (e.g. a delete that returns 202 + Accepted for async teardown or 204 No Content when already gone, both typed + ``None``), inspect ``resp.http_response.status_code`` to tell them apart. """ http_response: httpx.Response diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py index c53990986b..ad9667fb9d 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/types.py @@ -277,8 +277,10 @@ class RetryPolicy: Set as a client-level default via the ``retry`` constructor parameter, or override per-request via ``send()``'s ``retry`` keyword argument. - This is an operational concern — it does not belong in endpoint - signatures. + This is an operational concern and does not belong in endpoint + signatures. Response-header handling and broad server-error retries are + opt-in so adapters can reproduce another client's retry contract without + changing standalone client defaults. .. note:: @@ -300,6 +302,9 @@ class RetryPolicy: max_retries: int = 3 backoff_base: float = 0.5 retryable_status_codes: tuple[int, ...] = (502, 503, 504, 429) + retry_all_server_errors: bool = False + respect_retry_decision_headers: bool = False + respect_retry_after_headers: bool = False @dataclass(frozen=True, slots=True) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/client.py new file mode 100644 index 0000000000..914f50c43b --- /dev/null +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/client.py @@ -0,0 +1,420 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Typed HTTP clients for the Models service. + +Wraps the endpoint functions from ``models.endpoints`` as direct methods using +the ``method()`` descriptor (the Files/Secrets/Jobs pattern), and layers on the +Models-specific ergonomics that used to live on the vendored Stainless +``ModelsResource``: + +- OpenAI inference-gateway route builders (``get_openai_route_base_url`` and + friends) -- pure string builders, safe from sync or async code, and +- deployment/provider status polling (``wait_for_deployment_status`` / + ``wait_for_provider_status``) driven by the client's own ``get_deployment`` / + ``get_provider`` methods. + +The inference-gateway *readiness* probe (``wait_for_gateway``) lives one layer +up in ``packages/models`` because it targets the separate inference-gateway +service, not Models -- see that module and AIRCORE notes. + +Usage:: + + from nemo_platform_plugin.models.client import ModelsClient + from nemo_platform_plugin.models.types import CreateModelEntityRequest + + client = ModelsClient(base_url="...", workspace="default") + model = client.create_model(body=CreateModelEntityRequest(name="llama")).data() + for m in client.list_models().items(): + print(m.name) + client.wait_for_deployment_status("my-deploy", "READY") +""" + +from __future__ import annotations + +import asyncio +import time +from datetime import datetime +from typing import Protocol + +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.errors import NotFoundError +from nemo_platform_plugin.client.method import method +from nemo_platform_plugin.models import endpoints +from nemo_platform_plugin.models.types import ModelDeployment + +_INFERENCE_GATEWAY_PREFIX = "/apis/inference-gateway/v2/workspaces" + + +# The OpenAI-route builders only read a couple of attributes, so they accept any +# object exposing them -- the plugin ``ModelProvider`` / ``ModelEntity`` models, +# or the Stainless SDK equivalents that ``packages/models`` passes through. +# Structural typing keeps the plugin free of a dependency on the generated SDK +# types while still accepting them. + + +class ProviderLike(Protocol): + """An object identifying a model provider and its upstream host URL.""" + + @property + def workspace(self) -> str: ... + @property + def name(self) -> str: ... + @property + def host_url(self) -> str: ... + + +class ModelEntityLike(Protocol): + """An object identifying a model entity.""" + + @property + def workspace(self) -> str: ... + @property + def name(self) -> str: ... + + +class DeploymentLike(Protocol): + """An object identifying a deployment and its auto-created provider.""" + + @property + def name(self) -> str: ... + @property + def model_provider_id(self) -> str | None: ... + + +def _seconds_since_creation(entry_timestamp: datetime | str | None, created_at: datetime | None) -> int | None: + """Seconds from deployment creation to the entry timestamp, or None if not comparable.""" + if created_at is None or entry_timestamp is None: + return None + if isinstance(entry_timestamp, str): + try: + entry_timestamp = datetime.fromisoformat(entry_timestamp.replace("Z", "+00:00")) + except (ValueError, TypeError): + return None + if not hasattr(entry_timestamp, "timestamp") or not hasattr(created_at, "timestamp"): + return None + try: + return int(entry_timestamp.timestamp() - created_at.timestamp()) + except (TypeError, OSError): + return None + + +def _deployment_status(deployment: ModelDeployment) -> tuple[str, str]: + """Return ``(current_status, status_message)`` for a deployment. + + The API guarantees the last history entry is the current state; fall back to + the top-level fields when there is no history. + """ + history = deployment.status_history + if history: + last = history[-1] + return last.status.value, last.status_message or "" + return deployment.status.value, deployment.status_message or "" + + +def _print_new_history(deployment: ModelDeployment, last_history_len: int) -> int: + """Print any status-history entries not yet seen; return the new history length.""" + history = deployment.status_history + created_at = deployment.created_at + if len(history) > last_history_len: + for entry in history[last_history_len:]: + ts = entry.timestamp + ts_str = ts.strftime("%H:%M:%S") if hasattr(ts, "strftime") else str(ts) + secs = _seconds_since_creation(ts, created_at) + part = f" [{ts_str}] " + if secs is not None: + part += f"(+{secs}s) " + part += f"Status: {entry.status.value}" + if entry.status_message: + part += f" - {entry.status_message}" + print(part) + return len(history) + return last_history_len + + +class _ModelsMethods: + # Model entities + create_model = method(endpoints.create_model) + list_models = method(endpoints.list_models) + get_model = method(endpoints.get_model) + update_model = method(endpoints.update_model) + delete_model = method(endpoints.delete_model) + + # Nested adapters (base model in the path) + create_model_adapter = method(endpoints.create_model_adapter) + update_model_adapter = method(endpoints.update_model_adapter) + delete_model_adapter = method(endpoints.delete_model_adapter) + + # Top-level adapters + create_adapter = method(endpoints.create_adapter) + list_adapters = method(endpoints.list_adapters) + get_adapter = method(endpoints.get_adapter) + update_adapter = method(endpoints.update_adapter) + delete_adapter = method(endpoints.delete_adapter) + + # Model providers + create_provider = method(endpoints.create_provider) + list_providers = method(endpoints.list_providers) + get_provider = method(endpoints.get_provider) + upsert_provider = method(endpoints.upsert_provider) + update_provider_status = method(endpoints.update_provider_status) + delete_provider = method(endpoints.delete_provider) + + # Prompts + create_prompt = method(endpoints.create_prompt) + list_prompts = method(endpoints.list_prompts) + get_prompt = method(endpoints.get_prompt) + update_prompt = method(endpoints.update_prompt) + delete_prompt = method(endpoints.delete_prompt) + + # Model deployments + create_deployment = method(endpoints.create_deployment) + list_deployments = method(endpoints.list_deployments) + get_deployment = method(endpoints.get_deployment) + get_deployment_models = method(endpoints.get_deployment_models) + list_deployment_versions = method(endpoints.list_deployment_versions) + get_deployment_version = method(endpoints.get_deployment_version) + update_deployment = method(endpoints.update_deployment) + update_deployment_status = method(endpoints.update_deployment_status) + delete_deployment = method(endpoints.delete_deployment) + delete_deployment_version = method(endpoints.delete_deployment_version) + + # Model deployment configs + create_deployment_config = method(endpoints.create_deployment_config) + list_deployment_configs = method(endpoints.list_deployment_configs) + get_deployment_config = method(endpoints.get_deployment_config) + list_deployment_config_versions = method(endpoints.list_deployment_config_versions) + get_deployment_config_version = method(endpoints.get_deployment_config_version) + update_deployment_config = method(endpoints.update_deployment_config) + delete_deployment_config = method(endpoints.delete_deployment_config) + delete_deployment_config_version = method(endpoints.delete_deployment_config_version) + + +class _ModelsUrlMixin: + """Pure OpenAI-route URL builders. No I/O -- safe from sync or async code. + + Depends only on the client's ``base_url`` and default ``workspace`` (both + provided by :class:`BaseNemoClient`). + """ + + base_url: str + workspace: str | None + + def _resolve_workspace(self, workspace: str | None) -> str: + ws = workspace or self.workspace + if not ws: + raise ValueError("Missing workspace argument; either set a client-level workspace or pass workspace=...") + return ws + + def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: + """Base URL for the OpenAI proxy route (routes on the request body ``model`` field).""" + ws = self._resolve_workspace(workspace) + return f"{self.base_url}/{_INFERENCE_GATEWAY_PREFIX.lstrip('/')}/{ws}/openai/-/v1" + + def get_provider_route_openai_url(self, provider: ProviderLike) -> str: + """OpenAI SDK-compatible URL for a provider proxy route. + + Appends ``/v1`` unless the provider's ``host_url`` already ends in ``/v1``. + """ + route = ( + f"{self.base_url}/{_INFERENCE_GATEWAY_PREFIX.lstrip('/')}/{provider.workspace}/provider/{provider.name}/-" + ) + if not provider.host_url.rstrip("/").endswith("/v1"): + route = f"{route}/v1" + return route + + def get_model_entity_route_openai_url(self, model_entity: ModelEntityLike) -> str: + """OpenAI SDK-compatible URL for a model-entity proxy route (always ``/v1``).""" + return ( + f"{self.base_url}/{_INFERENCE_GATEWAY_PREFIX.lstrip('/')}/" + f"{model_entity.workspace}/model/{model_entity.name}/-/v1" + ) + + +class ModelsClient(_ModelsMethods, _ModelsUrlMixin, NemoClient): + """Sync client for the Models service API.""" + + def get_provider_route_openai_url_for_deployment(self, deployment: DeploymentLike) -> str: + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + if not deployment.model_provider_id: + raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") + workspace, name = deployment.model_provider_id.split("/", 1) + provider = self.get_provider(name=name, workspace=workspace).data() + return self.get_provider_route_openai_url(provider) + + def wait_for_deployment_status( + self, + deployment_name: str, + desired_status: str, + *, + workspace: str | None = None, + timeout: int = 1200, + poll_interval: float = 3.0, + ) -> bool: + """Poll a ModelDeployment until it reaches ``desired_status`` (or times out). + + For ``"DELETED"``, waits for the resource to be fully garbage collected + (404), not merely for the status to read DELETED. Returns False on + timeout or a terminal ERROR state. + """ + start = time.time() + last_status = "" + last_message = "" + last_history_len = 0 + print(f"Waiting for status: {desired_status}...\n") + + while time.time() - start < timeout: + try: + deployment = self.get_deployment(name=deployment_name, workspace=workspace).data() + except NotFoundError: + if desired_status == "DELETED": + print(f"Deployment {desired_status}!\n") + return True + print("Deployment not found\n") + return False + + current_status, status_message = _deployment_status(deployment) + last_status, last_message = current_status, status_message + last_history_len = _print_new_history(deployment, last_history_len) + + if current_status == desired_status and desired_status != "DELETED": + print(f"Deployment reached {desired_status} status!\n") + return True + if current_status == "ERROR": + print(f"Deployment entered ERROR state: {status_message}\n") + return False + time.sleep(poll_interval) + + detail = f"Last status: {last_status}" + if last_message: + detail += f" - {last_message}" + print(f"Timeout after {int(time.time() - start)}s. {detail}\n") + return False + + def wait_for_provider_status( + self, + provider_name: str, + desired_status: str = "READY", + *, + workspace: str | None = None, + timeout: int = 60, + poll_interval: float = 1.0, + ) -> bool: + """Poll a ModelProvider until it reaches ``desired_status`` (or times out).""" + start = time.time() + last_status = "" + print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") + + while time.time() - start < timeout: + try: + provider = self.get_provider(name=provider_name, workspace=workspace).data() + except NotFoundError: + print(f"\nProvider '{provider_name}' not found\n") + return False + + current_status = provider.status.value + if current_status != last_status: + elapsed = int(time.time() - start) + print(f" [{datetime.now().strftime('%H:%M:%S')}] ({elapsed}s) Status: {current_status}") + last_status = current_status + if current_status == desired_status: + return True + if current_status == "ERROR": + print(f"\nProvider entered ERROR state: {provider.status_message}\n") + return False + time.sleep(poll_interval) + + print(f"\nProvider timeout after {int(time.time() - start)}s. Last status: {last_status}\n") + return False + + +class AsyncModelsClient(_ModelsMethods, _ModelsUrlMixin, AsyncNemoClient): + """Async client for the Models service API.""" + + async def get_provider_route_openai_url_for_deployment(self, deployment: DeploymentLike) -> str: + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + if not deployment.model_provider_id: + raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") + workspace, name = deployment.model_provider_id.split("/", 1) + provider = (await self.get_provider(name=name, workspace=workspace)).data() + return self.get_provider_route_openai_url(provider) + + async def wait_for_deployment_status( + self, + deployment_name: str, + desired_status: str, + *, + workspace: str | None = None, + timeout: int = 1200, + poll_interval: float = 3.0, + ) -> bool: + """Async twin of :meth:`ModelsClient.wait_for_deployment_status`.""" + start = time.time() + last_status = "" + last_message = "" + last_history_len = 0 + print(f"Waiting for status: {desired_status}...\n") + + while time.time() - start < timeout: + try: + deployment = (await self.get_deployment(name=deployment_name, workspace=workspace)).data() + except NotFoundError: + if desired_status == "DELETED": + print(f"Deployment {desired_status}!\n") + return True + print("Deployment not found\n") + return False + + current_status, status_message = _deployment_status(deployment) + last_status, last_message = current_status, status_message + last_history_len = _print_new_history(deployment, last_history_len) + + if current_status == desired_status and desired_status != "DELETED": + print(f"Deployment reached {desired_status} status!\n") + return True + if current_status == "ERROR": + print(f"Deployment entered ERROR state: {status_message}\n") + return False + await asyncio.sleep(poll_interval) + + detail = f"Last status: {last_status}" + if last_message: + detail += f" - {last_message}" + print(f"Timeout after {int(time.time() - start)}s. {detail}\n") + return False + + async def wait_for_provider_status( + self, + provider_name: str, + desired_status: str = "READY", + *, + workspace: str | None = None, + timeout: int = 60, + poll_interval: float = 1.0, + ) -> bool: + """Async twin of :meth:`ModelsClient.wait_for_provider_status`.""" + start = time.time() + last_status = "" + print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") + + while time.time() - start < timeout: + try: + provider = (await self.get_provider(name=provider_name, workspace=workspace)).data() + except NotFoundError: + print(f"\nProvider '{provider_name}' not found\n") + return False + + current_status = provider.status.value + if current_status != last_status: + elapsed = int(time.time() - start) + print(f" [{datetime.now().strftime('%H:%M:%S')}] ({elapsed}s) Status: {current_status}") + last_status = current_status + if current_status == desired_status: + return True + if current_status == "ERROR": + print(f"\nProvider entered ERROR state: {provider.status_message}\n") + return False + await asyncio.sleep(poll_interval) + + print(f"\nProvider timeout after {int(time.time() - start)}s. Last status: {last_status}\n") + return False diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/endpoints.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/endpoints.py new file mode 100644 index 0000000000..ab1916510d --- /dev/null +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/endpoints.py @@ -0,0 +1,390 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Typed endpoint definitions for the Models service. + +These are the single source of truth for the HTTP contract. All paths include +the ``/apis/models`` gateway prefix. Although the Stainless SDK grouped some of +these under ``sdk.inference.*`` (deployments, providers, prompts), every route +is served by the Models service under ``/apis/models/v2/...``. + +The service exposes six resource groups: +- model entities (``/models``) and their nested adapters (``/models/{m}/adapters``), +- top-level adapters (``/adapters``), +- model providers (``/providers``), +- prompts (``/prompts``), +- model deployments (``/deployments``) with immutable versioning, and +- model deployment configs (``/deployment-configs``) with immutable versioning. +""" + +from __future__ import annotations + +from abc import abstractmethod +from typing import Any + +from nemo_platform_plugin.client.endpoint import delete, get, patch, post, put +from nemo_platform_plugin.client.types import Paginated, PreparedRequest +from nemo_platform_plugin.models.types import ( + Adapter, + CreateAdapterRequest, + CreateModelAdapterRequest, + CreateModelDeploymentConfigRequest, + CreateModelDeploymentRequest, + CreateModelEntityRequest, + CreateModelProviderRequest, + CreatePromptRequest, + GetModelQueryParams, + ListAdaptersQueryParams, + ListDeploymentConfigsQueryParams, + ListDeploymentsQueryParams, + ListModelsQueryParams, + ListPromptsQueryParams, + ListProvidersQueryParams, + ModelDeployment, + ModelDeploymentConfig, + ModelEntity, + ModelProvider, + Prompt, + UpdateAdapterRequest, + UpdateDeploymentStatusQueryParams, + UpdateModelDeploymentConfigRequest, + UpdateModelDeploymentRequest, + UpdateModelDeploymentStatusRequest, + UpdateModelEntityRequest, + UpdateModelProviderStatusRequest, + UpdatePromptRequest, + UpsertModelProviderRequest, +) + +_MODELS = "/apis/models/v2/workspaces/{workspace}" + + +# --------------------------------------------------------------------------- +# Model entities +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/models/{name}") +@abstractmethod +def get_model( + *, workspace: str | None = None, name: str, query_params: GetModelQueryParams | None = None +) -> ModelEntity: ... + + +@get(_MODELS + "/models") +@abstractmethod +def list_models( + *, workspace: str | None = None, query_params: ListModelsQueryParams | None = None +) -> Paginated[ModelEntity]: ... + + +def _get_model_on_conflict(body: CreateModelEntityRequest, workspace: str | None) -> PreparedRequest[ModelEntity]: + """Retrieve request replayed when ``create_model(exist_ok=True)`` 409s.""" + return get_model(name=body.name, workspace=workspace) + + +@post(_MODELS + "/models", get_on_conflict=_get_model_on_conflict) +@abstractmethod +def create_model( + *, workspace: str | None = None, body: CreateModelEntityRequest, exist_ok: bool = False +) -> ModelEntity: ... + + +@patch(_MODELS + "/models/{name}") +@abstractmethod +def update_model( + *, + workspace: str | None = None, + name: str, + body: UpdateModelEntityRequest, + query_params: GetModelQueryParams | None = None, +) -> ModelEntity: ... + + +@delete(_MODELS + "/models/{name}") +@abstractmethod +def delete_model(*, workspace: str | None = None, name: str) -> None: ... + + +# --------------------------------------------------------------------------- +# Nested adapters (base model in the path) +# --------------------------------------------------------------------------- + + +@post(_MODELS + "/models/{model_name}/adapters") +@abstractmethod +def create_model_adapter( + *, workspace: str | None = None, model_name: str, body: CreateModelAdapterRequest +) -> Adapter: ... + + +@patch(_MODELS + "/models/{model_name}/adapters/{adapter}") +@abstractmethod +def update_model_adapter( + *, workspace: str | None = None, model_name: str, adapter: str, body: UpdateAdapterRequest +) -> Adapter: ... + + +@delete(_MODELS + "/models/{model_name}/adapters/{adapter}") +@abstractmethod +def delete_model_adapter(*, workspace: str | None = None, model_name: str, adapter: str) -> None: ... + + +# --------------------------------------------------------------------------- +# Top-level adapters +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/adapters/{name}") +@abstractmethod +def get_adapter(*, workspace: str | None = None, name: str) -> Adapter: ... + + +@get(_MODELS + "/adapters") +@abstractmethod +def list_adapters( + *, workspace: str | None = None, query_params: ListAdaptersQueryParams | None = None +) -> Paginated[Adapter]: ... + + +def _get_adapter_on_conflict(body: CreateAdapterRequest, workspace: str | None) -> PreparedRequest[Adapter]: + return get_adapter(name=body.name, workspace=workspace) + + +@post(_MODELS + "/adapters", get_on_conflict=_get_adapter_on_conflict) +@abstractmethod +def create_adapter(*, workspace: str | None = None, body: CreateAdapterRequest, exist_ok: bool = False) -> Adapter: ... + + +@patch(_MODELS + "/adapters/{name}") +@abstractmethod +def update_adapter(*, workspace: str | None = None, name: str, body: UpdateAdapterRequest) -> Adapter: ... + + +@delete(_MODELS + "/adapters/{name}") +@abstractmethod +def delete_adapter(*, workspace: str | None = None, name: str) -> None: ... + + +# --------------------------------------------------------------------------- +# Model providers +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/providers/{name}") +@abstractmethod +def get_provider(*, workspace: str | None = None, name: str) -> ModelProvider: ... + + +@get(_MODELS + "/providers") +@abstractmethod +def list_providers( + *, workspace: str | None = None, query_params: ListProvidersQueryParams | None = None +) -> Paginated[ModelProvider]: ... + + +def _get_provider_on_conflict( + body: CreateModelProviderRequest, workspace: str | None +) -> PreparedRequest[ModelProvider]: + return get_provider(name=body.name, workspace=workspace) + + +@post(_MODELS + "/providers", get_on_conflict=_get_provider_on_conflict) +@abstractmethod +def create_provider( + *, workspace: str | None = None, body: CreateModelProviderRequest, exist_ok: bool = False +) -> ModelProvider: ... + + +@put(_MODELS + "/providers/{name}") +@abstractmethod +def upsert_provider(*, workspace: str | None = None, name: str, body: UpsertModelProviderRequest) -> ModelProvider: ... + + +@put(_MODELS + "/providers/{name}/status") +@abstractmethod +def update_provider_status( + *, workspace: str | None = None, name: str, body: UpdateModelProviderStatusRequest +) -> ModelProvider: ... + + +@delete(_MODELS + "/providers/{name}") +@abstractmethod +def delete_provider(*, workspace: str | None = None, name: str) -> None: ... + + +# --------------------------------------------------------------------------- +# Prompts +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/prompts/{name}") +@abstractmethod +def get_prompt(*, workspace: str | None = None, name: str) -> Prompt: ... + + +@get(_MODELS + "/prompts") +@abstractmethod +def list_prompts( + *, workspace: str | None = None, query_params: ListPromptsQueryParams | None = None +) -> Paginated[Prompt]: ... + + +def _get_prompt_on_conflict(body: CreatePromptRequest, workspace: str | None) -> PreparedRequest[Prompt]: + return get_prompt(name=body.name, workspace=workspace) + + +@post(_MODELS + "/prompts", get_on_conflict=_get_prompt_on_conflict) +@abstractmethod +def create_prompt(*, workspace: str | None = None, body: CreatePromptRequest, exist_ok: bool = False) -> Prompt: ... + + +@put(_MODELS + "/prompts/{name}") +@abstractmethod +def update_prompt(*, workspace: str | None = None, name: str, body: UpdatePromptRequest) -> Prompt: ... + + +@delete(_MODELS + "/prompts/{name}") +@abstractmethod +def delete_prompt(*, workspace: str | None = None, name: str) -> None: ... + + +# --------------------------------------------------------------------------- +# Model deployments +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/deployments/{name}") +@abstractmethod +def get_deployment(*, workspace: str | None = None, name: str) -> ModelDeployment: ... + + +@get(_MODELS + "/deployments") +@abstractmethod +def list_deployments( + *, workspace: str | None = None, query_params: ListDeploymentsQueryParams | None = None +) -> Paginated[ModelDeployment]: ... + + +@get(_MODELS + "/deployments/{name}/models") +@abstractmethod +def get_deployment_models(*, workspace: str | None = None, name: str) -> dict[str, Any]: ... + + +@get(_MODELS + "/deployments/{name}/versions") +@abstractmethod +def list_deployment_versions(*, workspace: str | None = None, name: str) -> list[ModelDeployment]: ... + + +@get(_MODELS + "/deployments/{deployment}/versions/{name}") +@abstractmethod +def get_deployment_version(*, workspace: str | None = None, deployment: str, name: str) -> ModelDeployment: ... + + +def _get_deployment_on_conflict( + body: CreateModelDeploymentRequest, workspace: str | None +) -> PreparedRequest[ModelDeployment]: + return get_deployment(name=body.name, workspace=workspace) + + +@post(_MODELS + "/deployments", get_on_conflict=_get_deployment_on_conflict) +@abstractmethod +def create_deployment( + *, workspace: str | None = None, body: CreateModelDeploymentRequest, exist_ok: bool = False +) -> ModelDeployment: ... + + +@post(_MODELS + "/deployments/{name}") +@abstractmethod +def update_deployment( + *, workspace: str | None = None, name: str, body: UpdateModelDeploymentRequest +) -> ModelDeployment: ... + + +@post(_MODELS + "/deployments/{name}/status") +@abstractmethod +def update_deployment_status( + *, + workspace: str | None = None, + name: str, + body: UpdateModelDeploymentStatusRequest, + query_params: UpdateDeploymentStatusQueryParams | None = None, +) -> ModelDeployment: ... + + +@delete(_MODELS + "/deployments/{name}") +@abstractmethod +def delete_deployment(*, workspace: str | None = None, name: str) -> None: + """Delete a deployment. + + Returns 202 Accepted when teardown is asynchronous (the deployment enters + DELETING while infrastructure is torn down) or 204 No Content when the + delete is synchronous (already hard-deleted). Both are success and the typed + body is ``None``; read ``resp.http_response.status_code`` to distinguish them. + """ + + +@delete(_MODELS + "/deployments/{deployment}/versions/{name}") +@abstractmethod +def delete_deployment_version(*, workspace: str | None = None, deployment: str, name: str) -> None: + """Delete a single deployment version (202 async / 204 synchronous; body ``None``). + + See :func:`delete_deployment` for the 202-vs-204 distinction. + """ + + +# --------------------------------------------------------------------------- +# Model deployment configs +# --------------------------------------------------------------------------- + + +@get(_MODELS + "/deployment-configs/{name}") +@abstractmethod +def get_deployment_config(*, workspace: str | None = None, name: str) -> ModelDeploymentConfig: ... + + +@get(_MODELS + "/deployment-configs") +@abstractmethod +def list_deployment_configs( + *, workspace: str | None = None, query_params: ListDeploymentConfigsQueryParams | None = None +) -> Paginated[ModelDeploymentConfig]: ... + + +@get(_MODELS + "/deployment-configs/{name}/versions") +@abstractmethod +def list_deployment_config_versions(*, workspace: str | None = None, name: str) -> list[ModelDeploymentConfig]: ... + + +@get(_MODELS + "/deployment-configs/{config}/versions/{name}") +@abstractmethod +def get_deployment_config_version(*, workspace: str | None = None, config: str, name: str) -> ModelDeploymentConfig: ... + + +def _get_deployment_config_on_conflict( + body: CreateModelDeploymentConfigRequest, workspace: str | None +) -> PreparedRequest[ModelDeploymentConfig]: + return get_deployment_config(name=body.name, workspace=workspace) + + +@post(_MODELS + "/deployment-configs", get_on_conflict=_get_deployment_config_on_conflict) +@abstractmethod +def create_deployment_config( + *, workspace: str | None = None, body: CreateModelDeploymentConfigRequest, exist_ok: bool = False +) -> ModelDeploymentConfig: ... + + +@post(_MODELS + "/deployment-configs/{name}") +@abstractmethod +def update_deployment_config( + *, workspace: str | None = None, name: str, body: UpdateModelDeploymentConfigRequest +) -> ModelDeploymentConfig: ... + + +@delete(_MODELS + "/deployment-configs/{name}") +@abstractmethod +def delete_deployment_config(*, workspace: str | None = None, name: str) -> None: ... + + +@delete(_MODELS + "/deployment-configs/{config}/versions/{name}") +@abstractmethod +def delete_deployment_config_version(*, workspace: str | None = None, config: str, name: str) -> None: ... diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/types.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/types.py new file mode 100644 index 0000000000..0ce274132d --- /dev/null +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/models/types.py @@ -0,0 +1,1466 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Typed request/response models for the Models service client. + +These models mirror the HTTP contract for model entities, adapters, model +providers, prompts, model deployments, and model deployment configs. The +Models service remains the authoritative wire-schema owner; this plugin module +keeps client DTOs independent of the Stainless-generated SDK. + +The plugin package must stay free of an ``nmp_common`` dependency because that +would create a reverse service dependency. Server-only pieces are handled per +the data-vs-behavior split documented in ``client/MIGRATION.md``: + +- **Constants** (name regex, max lengths) that live in + ``nmp.common.entities.constants`` are inlined here with a comment pointing at + the origin (matching the ``secrets.types`` boundary). +- **``AuthContext``** is a pure-data mirror of ``nmp.common.auth.AuthContext`` + with the same wire shape and no ``from_principal``/``to_principal`` behavior. +- **``InferenceParams``** is a faithful replica of + ``nmp.common.inference.InferenceParams`` (pure pydantic, no server imports). +- **``BackendFormat``** is reused from ``nemo_platform_plugin.inference_middleware``. +- The Jinja2 ``auth_header_format`` validator needs ``jinja2`` (not a plugin + dependency) so it stays server-side; the plugin field is a plain string. +- Entity-store ``Filter`` subclasses are genuinely server-only and are not + mirrored here; the client passes filters as a ``filter`` query-param string. +""" + +from __future__ import annotations + +import re +from datetime import datetime +from enum import Enum, StrEnum +from typing import Any, NotRequired, Self, TypedDict + +from nemo_platform_plugin.inference_middleware import BackendFormat +from pydantic import AnyUrl, BaseModel, ConfigDict, Field, field_validator, model_validator + +# --------------------------------------------------------------------------- +# Inlined constants +# +# Mirror ``nmp.common.entities.constants`` and ``nmp.core.models.constants``. +# Inlined rather than imported so this package stays free of an ``nmp_common`` +# (or models-service) dependency -- see the module docstring. +# --------------------------------------------------------------------------- + +_NAME_REGEX = r"^[\w\-.]+$" # constants.REGEX_WORD_CHARACTER_DOT_DASH +_NAME_SLASH_REGEX = r"^[\w\-./]+$" # constants.REGEX_WORD_CHARACTER_DOT_DASH_SLASH +_NAME_DESC = "Allowed characters: letters (a-z, A-Z), digits (0-9), underscores, hyphens, and dots." +_MAX_LEN_255 = 255 # constants.MAX_LENGTH_255 + +# nmp.core.models.constants -- adapter/model reference rules. +_MODEL_REF_NAME_SEGMENT = r"[a-z](?!.*--)[a-z0-9\-@.+_]{1,62}(? bool: + """True if *value* matches :data:`MODEL_REF_PATTERN` (entity NAME rules per segment).""" + return _MODEL_REF_RE.fullmatch(value) is not None + + +# --------------------------------------------------------------------------- +# Auth context (data-only mirror of nmp.common.auth.AuthContext) +# --------------------------------------------------------------------------- + + +class AuthContext(BaseModel): + """Auth context captured at resource creation for delegated access. + + This is the wire/data shape. The server's ``nmp.common.auth.AuthContext`` + adds ``from_principal`` / ``to_principal`` behavior on top of the same + fields, and the server response models re-type this field to that class. + """ + + principal_id: str = Field(..., description="The principal's unique identifier") + principal_email: str | None = Field(default=None, description="The principal's email address") + principal_groups: list[str] = Field(default_factory=list, description="Groups the principal belongs to") + principal_on_behalf_of: str | None = Field( + default=None, description="If acting on behalf of another principal, their principal ID" + ) + principal_on_behalf_of_groups: list[str] | None = Field( + default=None, description="Groups the on-behalf-of principal belongs to" + ) + principal_on_behalf_of_email: str | None = Field( + default=None, description="The on-behalf-of principal's email address" + ) + + +# --------------------------------------------------------------------------- +# Inference parameters (replica of nmp.common.inference.InferenceParams) +# --------------------------------------------------------------------------- + + +class InferenceParams(BaseModel): + """Parameters for model inference. + + Extra fields can be supplied for additional options applied to the inference + request directly. Fields not supported by the model may cause inference + errors during evaluation. + """ + + model_config = ConfigDict(extra="allow") + + model: str | None = Field(default=None, description="Model identifier") + temperature: float | None = Field( + default=None, + ge=0, + le=2, + description="Float value between 0 and 1. temp of 0 indicates greedy decoding, " + "where the token with highest prob is chosen. Temperature can't be set to 0.0 currently", + ) + max_tokens: int | None = Field(default=None, ge=1, description="Max tokens to generate") + max_completion_tokens: int | None = Field(default=None, ge=1, description="Max tokens to generate") + top_p: float | None = Field( + default=None, + ge=0, + le=1, + description="Float value between 0 and 1; limits to the top tokens within a certain " + "probability. top_p=0 means the model will only consider the single most likely " + "token for the next prediction", + ) + stop: list[str] | None = Field(default=None) + + @model_validator(mode="after") + def check_max_tokens(self) -> Self: + if self.max_tokens and self.max_completion_tokens: + raise ValueError( + "max_tokens and max_completion_tokens cannot both be configured. " + "Choose the appropriate tokens parameter for the model." + ) + return self + + +# --------------------------------------------------------------------------- +# Value types +# --------------------------------------------------------------------------- + + +class ModelPrecision(str, Enum): + """Type of model precision.""" + + INT8 = "int8" + BF16 = "bf16" + FP16 = "fp16" + FP32 = "fp32" + FP8_MIXED = "fp8-mixed" + BF16_MIXED = "bf16-mixed" + + +class FinetuningType(str, Enum): + """Finetuning types.""" + + LORA_MERGED = "lora_merged" + ALL_WEIGHTS = "all_weights" + + LAST_LAYER = "last_layer" + TOP_LAYERS = "top_layers" + GRADUAL_UNFREEZING = "gradual_unfreezing" + BIAS_ONLY = "bias_only" # BitFit + ATTENTION_ONLY = "attention_only" + + LORA = "lora" + QLORA = "qlora" + ADALORA = "adalora" + DORA = "dora" + LORA_PLUS = "lora_plus" + + PROMPT_TUNING = "prompt_tuning" + PREFIX_TUNING = "prefix_tuning" + P_TUNING = "p_tuning" + P_TUNING_V2 = "p_tuning_v2" + SOFT_PROMPT = "soft_prompt" + + PPO = "ppo" + DPO = "dpo" + CDPO = "cdpo" + IPO = "ipo" + ORPO = "orpo" + KTO = "kto" + RRHF = "rrhf" + GRPO = "grpo" + + +class MoEConfig(BaseModel): + """Mixture of Experts configuration.""" + + num_experts: int = Field(description="Total number of routed experts (sharded by EP)") + num_experts_per_tok: int = Field(description="Number of experts activated per token (top-k routing)") + num_expert_layers: int = Field(description="Number of layers with MoE") + expert_ffn_size: int | None = Field(default=None, description="FFN size for experts (if different from main FFN)") + num_shared_experts: int = Field(default=0, description="Number of shared experts (replicated, not sharded by EP)") + + +class MambaConfig(BaseModel): + """Mamba/State Space Model configuration.""" + + is_hybrid: bool = Field(description="Whether model is Mamba-Transformer hybrid") + num_mamba_layers: int = Field(description="Number of Mamba/SSM layers") + num_attention_layers: int = Field(default=0, description="Number of attention layers (for hybrids)") + num_mlp_layers: int = Field( + default=0, description="Number of standalone MLP layers (for interleaved architectures)" + ) + state_size: int = Field(default=16, description="SSM state expansion factor (d_state)") + conv_kernel: int = Field(default=4, description="Convolution kernel size for Mamba (d_conv)") + + +class SlidingWindowConfig(BaseModel): + """Sliding window attention configuration.""" + + window_size: int = Field(description="Sliding window size (attends to last N tokens)") + + +class ToolCallConfig(BaseModel): + """Configuration for tool calling support in NIM deployments.""" + + tool_call_parser: str | None = Field( + default=None, + description="Name of the tool call parser to use (e.g., 'openai', 'hermes', 'pythonic', 'llama3_json', 'mistral').", + max_length=_MAX_LEN_255, + ) + tool_call_plugin: str | None = Field( + default=None, + description="Reference to a fileset containing the custom tool call plugin Python file. " + "Expected format: '{workspace}/{fileset_name}'. The fileset is mounted separately from " + "the model checkpoint at deployment time.", + max_length=_MAX_LEN_255, + ) + auto_tool_choice: bool | None = Field( + default=None, + description="Whether to enable automatic tool choice. When enabled, the model can decide to call tools " + "without explicit user instruction.", + ) + + +class LinearLayerSpec(BaseModel): + """Specification for a single linear layer in the model.""" + + name: str = Field(description="Module name (e.g., 'model.layers.0.self_attn.q_proj')") + in_features: int = Field(description="Input feature dimension") + out_features: int = Field(description="Output feature dimension") + + +class ModelSpec(BaseModel): + """Detailed specification for a model.""" + + context_size: int | None = Field(None, description="Context window size") + num_virtual_tokens: int | None = Field(None, description="Number of virtual tokens for prompt tuning") + is_chat: bool | None = Field(None, description="Whether this is a chat model") + is_embedding_model: bool = Field(False, description="Whether this is an embedding model") + + # Basic model information + checkpoint_model_name: str = Field(description="Checkpoint Model identifier or model path") + family: str = Field(description="Model architecture family (e.g., 'llama', 'mixtral', 'gpt2')") + + # Architecture dimensions + num_layers: int = Field(description="Number of transformer layers") + hidden_size: int = Field(description="Hidden dimension size") + num_attention_heads: int = Field(description="Number of attention heads") + num_kv_heads: int = Field(description="Number of key-value heads (for GQA/MQA)") + ffn_hidden_size: int = Field(description="FFN intermediate size") + vocab_size: int = Field(description="Vocabulary size") + + # Model properties + tied_embeddings: bool = Field(description="Whether embeddings are tied") + gated_mlp: bool = Field(description="Whether MLP uses gated activation") + base_num_parameters: int = Field(description="Total model parameters") + precision: str = Field(description="Model precision (e.g., 'float16', 'bfloat16', 'float32', 'int8', 'int4')") + + # Optional configurations + moe_config: MoEConfig | None = Field(default=None, description="MoE configuration if applicable") + mamba_config: MambaConfig | None = Field(default=None, description="Mamba/SSM configuration if applicable") + sliding_window_config: SlidingWindowConfig | None = Field( + default=None, description="Sliding window attention config if applicable" + ) + + # LoRA-specific metadata (pre-computed to avoid model instantiation) + linear_layers: list[LinearLayerSpec] | None = Field( + default=None, + description="List of all linear/Conv1D layers with their dimensions. " + "Used for LoRA parameter estimation without requiring model instantiation. " + "Each entry contains the module name, in_features, and out_features.", + ) + + # Deployment configuration + chat_template: str | None = Field( + default=None, + description="Jinja2 chat template string for the model. Used by NIM to format chat completions. " + "If not set, the model's built-in tokenizer template is used.", + ) + tool_call_config: ToolCallConfig | None = Field( + default=None, + description="Tool calling configuration for NIM deployments. Controls how the model handles " + "function/tool calling in chat completions.", + ) + + # GPU requirements (auto-calculated) + minimum_gpus_all_weights: int | None = Field( + default=None, + description="Minimum GPUs required for full fine-tuning using default configurations.", + ) + minimum_gpus_lora: int | None = Field( + default=None, + description="Minimum GPUs required for LoRA fine-tuning using default configurations.", + ) + + def model_precision(self) -> ModelPrecision: + """Convert the precision string to a :class:`ModelPrecision` enum.""" + precision_map = { + "bf16-mixed": ModelPrecision.BF16_MIXED, + "bf16": ModelPrecision.BF16, + "bfloat16": ModelPrecision.BF16, + "float16": ModelPrecision.FP16, + "float32": ModelPrecision.FP32, + "fp16": ModelPrecision.FP16, + "fp32": ModelPrecision.FP32, + "fp8-mixed": ModelPrecision.FP8_MIXED, + "int4": ModelPrecision.INT8, # Map int4 to int8 as int4 is not in the enum + "int8": ModelPrecision.INT8, + } + if self.precision in precision_map: + return precision_map[self.precision] + return ModelPrecision.BF16 + + +class Lora(BaseModel): + alpha: int | None = Field(None, description="Alpha scaling used for this adapter") + rank: int = Field(..., description="LoRA Rank") + + +class APIEndpointData(BaseModel): + """Data about an inference endpoint.""" + + url: AnyUrl | None = Field(None, description="Endpoint URL") + model_id: str | None = Field(None, description="Model identifier at the endpoint") + api_key: str | None = Field(None, description="API key for authentication") + format: str | None = Field(None, description="API format (e.g., openai, nvidia)") + + +class PromptData(BaseModel): + """Configuration for prompt engineering.""" + + system_prompt: str | None = Field(None, description="System prompt template") + icl_few_shot_examples: str | None = Field(None, description="In-context learning examples") + inference_params: InferenceParams | None = Field( + default=None, description="Inference parameters that should be overridden." + ) + system_prompt_template: str | None = Field( + default=None, + title="System Prompt Template", + description="The template which will be used to compile the final prompt used for prompting the LLM. Currently supports only {{icl_few_shot_examples}}", + ) + + +# --------------------------------------------------------------------------- +# Base model +# --------------------------------------------------------------------------- + + +class ModelEntityBaseModel(BaseModel): + """Base model for all Models service domain objects.""" + + id: str = Field(..., description="Autogenerated id") + name: str = Field( + description=f"Name of the entity. Name/workspace combo must be unique across all entities. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["llama-3.1-8b", "my-custom-model"], + ) + workspace: str = Field( + description=f"The workspace of the entity. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this entity.", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + created_at: datetime = Field(..., description="The timestamp of model entity creation") + updated_at: datetime = Field(..., description="The timestamp of the last model entity update") + + +# --------------------------------------------------------------------------- +# ModelProvider +# --------------------------------------------------------------------------- + +_AUTH_HEADER_FORMAT_DESCRIPTION = ( + "Jinja2 template string controlling how the API key secret is sent to the upstream. " + "Must contain exactly one variable named `auth_secret`, which is substituted with the " + "resolved secret value at request time. " + "Example: `'X-Api-Key: {{ auth_secret }}'`. " + "If not set, defaults to `'Authorization: Bearer {{ auth_secret }}'`." +) + + +class ModelProviderStatus(str, Enum): + """Status enum for ModelProvider objects.""" + + UNKNOWN = "UNKNOWN" + CREATED = "CREATED" + PENDING = "PENDING" + READY = "READY" + ERROR = "ERROR" + DELETING = "DELETING" + DELETED = "DELETED" + LOST = "LOST" + + +class ServedModelMapping(BaseModel): + """Mapping between a Model Entity and how it's served by this provider.""" + + model_entity_id: str = Field( + description="Model Entity identifier as workspace/name (e.g., 'my-ws/my-model')", + max_length=_MAX_LEN_255, + ) + served_model_name: str = Field( + description="The actual model name to send to the backend endpoint in the 'model' field", + max_length=_MAX_LEN_255, + ) + + +class ModelProvider(ModelEntityBaseModel): + """A reachable network endpoint that provides inference for one or more Model Entities. + + The unique identifier for a ModelProvider is the combination of workspace/name. + """ + + id: str = Field(default="", description="Unique identifier for the model provider") + description: str | None = Field( + default=None, + description="Optional description of the model provider", + max_length=1000, + ) + host_url: str = Field( + description="The network endpoint URL for the model provider", + max_length=2048, + ) + api_key_secret_name: str | None = Field( + default=None, + description="Reference to the API key stored in Secrets service", + max_length=_MAX_LEN_255, + ) + served_models: list[ServedModelMapping] | None = Field( + default_factory=list, + description="List of models served by this provider with routing information for IGW", + ) + enabled_models: list[str] | None = Field( + default=None, + description="Optional list of specific models to enable from this provider. If not set, all discovered models are enabled.", + ) + status: ModelProviderStatus = Field( + default=ModelProviderStatus.UNKNOWN, + description="Current status of the model provider, populated by models service", + ) + status_message: str = Field( + default="", + description="Detailed status message, populated by models service", + max_length=1000, + ) + default_extra_body: dict[str, Any] | None = Field( + default=None, + description="Default body parameters for inference requests. Can be overridden by user requests.", + ) + default_extra_headers: dict[str, str] | None = Field( + default=None, + description="Default headers for inference requests. Can be overridden by user requests.", + ) + required_extra_body: dict[str, Any] | None = Field( + default=None, + description="Required body parameters for inference requests. Cannot be overridden by user requests.", + ) + required_extra_headers: dict[str, str] | None = Field( + default=None, + description="Required headers for inference requests. Cannot be overridden by user requests.", + ) + model_deployment_id: str | None = Field( + default=None, + description="Optional reference to the ModelDeployment ID if this provider was auto-created for a deployment", + max_length=_MAX_LEN_255, + ) + auth_context: AuthContext | None = Field(default=None, description="Auth context captured at provider creation.") + auth_header_format: str | None = Field( + default=None, + description=_AUTH_HEADER_FORMAT_DESCRIPTION, + max_length=1024, + ) + + +class ModelProviderSort(StrEnum): + """Sort fields for ModelProvider queries.""" + + NAME_ASC = "name" + NAME_DESC = "-name" + CREATED_AT_ASC = "created_at" + CREATED_AT_DESC = "-created_at" + UPDATED_AT_ASC = "updated_at" + UPDATED_AT_DESC = "-updated_at" + STATUS_ASC = "status" + STATUS_DESC = "-status" + + +class CreateModelProviderRequest(BaseModel): + """Request model for creating a ModelProvider.""" + + name: str = Field( + description=f"Name of the model provider. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["my-nim-provider", "openai-endpoint"], + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this model provider", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field( + default=None, + description="Optional description of the model provider", + max_length=1000, + ) + host_url: str = Field( + description="The network endpoint URL for the model provider", + max_length=2048, + ) + api_key_secret_name: str | None = Field( + default=None, + description="Reference to an API key secret stored in the Secrets service. " + "Create the secret first via secrets API, then pass the secret name here.", + max_length=_MAX_LEN_255, + ) + enabled_models: list[str] | None = Field( + default=None, description="Optional list of specific models to enable from this provider" + ) + default_extra_body: dict[str, Any] | None = Field( + default=None, + description="Default body parameters for inference requests. Can be overridden by user requests.", + ) + default_extra_headers: dict[str, str] | None = Field( + default=None, + description="Default headers for inference requests. Can be overridden by user requests.", + ) + required_extra_body: dict[str, Any] | None = Field( + default=None, + description="Required body parameters for inference requests. Cannot be overridden by user requests.", + ) + required_extra_headers: dict[str, str] | None = Field( + default=None, + description="Required headers for inference requests. Cannot be overridden by user requests.", + ) + model_deployment_id: str | None = Field( + default=None, + description="Optional reference to the ModelDeployment ID if this provider is being auto-created for a deployment", + max_length=_MAX_LEN_255, + ) + status: ModelProviderStatus | None = Field(default=None, description="Status of the model provider") + status_message: str | None = Field( + default=None, + description="Status message", + max_length=1000, + ) + auth_header_format: str | None = Field( + default=None, + description=_AUTH_HEADER_FORMAT_DESCRIPTION, + max_length=1024, + ) + + +class UpsertModelProviderRequest(BaseModel): + """Request model for upserting a ModelProvider (PUT). + + All fields must be provided - partial updates are not supported for security reasons. + Use PUT /status endpoint to update status-related fields only. + """ + + project: str | None = Field( + default=None, + description="The URN of the project associated with this model provider", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field( + default=None, + description="Optional description of the model provider", + max_length=1000, + ) + host_url: str = Field( + description="The network endpoint URL for the model provider", + max_length=2048, + ) + api_key_secret_name: str | None = Field( + default=None, + description="Reference to an API key secret stored in the Secrets service. " + "Create the secret first via secrets API, then pass the secret name here.", + max_length=_MAX_LEN_255, + ) + enabled_models: list[str] | None = Field( + default=None, description="Optional list of specific models to enable from this provider" + ) + default_extra_body: dict[str, Any] | None = Field( + default=None, + description="Default body parameters for inference requests. Can be overridden by user requests.", + ) + default_extra_headers: dict[str, str] | None = Field( + default=None, + description="Default headers for inference requests. Can be overridden by user requests.", + ) + required_extra_body: dict[str, Any] | None = Field( + default=None, + description="Required body parameters for inference requests. Cannot be overridden by user requests.", + ) + required_extra_headers: dict[str, str] | None = Field( + default=None, + description="Required headers for inference requests. Cannot be overridden by user requests.", + ) + model_deployment_id: str | None = Field( + default=None, + description="Optional reference to the ModelDeployment ID if this provider is associated with a deployment", + max_length=_MAX_LEN_255, + ) + status: ModelProviderStatus | None = Field(default=None, description="Status of the model provider") + status_message: str | None = Field( + default=None, + description="Status message", + max_length=1000, + ) + auth_header_format: str | None = Field( + default=None, + description=_AUTH_HEADER_FORMAT_DESCRIPTION, + max_length=1024, + ) + + +class UpdateModelProviderStatusRequest(BaseModel): + """Request model for updating ModelProvider status and autodiscovery fields.""" + + model_deployment_id: str | None = Field( + default=None, + description="Reference to the ModelDeployment ID if this provider is associated with a deployment", + max_length=_MAX_LEN_255, + ) + served_models: list[ServedModelMapping] | None = Field( + default=None, description="List of models served by this provider with routing information for IGW" + ) + status: ModelProviderStatus | None = Field(default=None, description="Status of the model provider") + status_message: str | None = Field( + default=None, + description="Status message. If status is provided without status_message, defaults to empty string.", + max_length=1000, + ) + + +# --------------------------------------------------------------------------- +# Prompt +# --------------------------------------------------------------------------- + + +class PromptMessageRole(StrEnum): + """Role of a message author in a chat prompt.""" + + SYSTEM = "system" + DEVELOPER = "developer" + USER = "user" + ASSISTANT = "assistant" + + +class PromptMessage(BaseModel): + """A single templated message in a chat prompt.""" + + role: PromptMessageRole = Field(description="The role of the message author.") + content: str = Field(description="Templated message content. May contain template variables.") + + +class FunctionDefinition(BaseModel): + """An OpenAI-compatible function definition for tool calling.""" + + name: str = Field( + description="The name of the function to be called.", + max_length=_MAX_LEN_255, + ) + description: str | None = Field( + default=None, + description="A description of what the function does, used by the model to decide when and how to call it.", + ) + parameters: dict[str, Any] | None = Field( + default=None, + description="The parameters the function accepts, described as a JSON Schema object.", + ) + strict: bool | None = Field( + default=None, + description="Whether to enforce strict schema adherence when generating the function call.", + ) + + +class ChatCompletionTool(BaseModel): + """An OpenAI-compatible tool definition (currently always a function tool).""" + + type: str = Field(description="The type of the tool. Currently only 'function' is supported.") + function: FunctionDefinition = Field(description="The function definition for this tool.") + + @field_validator("type") + @classmethod + def _validate_type(cls, v: str) -> str: + if v != "function": + raise ValueError("Only 'function' tools are supported") + return v + + +class Prompt(ModelEntityBaseModel): + """A reusable, stored chat prompt. The unique identifier is workspace/name.""" + + id: str = Field(default="", description="Unique identifier for the prompt.") + description: str | None = Field( + default=None, + description="Optional description of the prompt.", + max_length=1000, + ) + messages: list[PromptMessage] = Field( + default_factory=list, + description="Ordered list of chat messages that make up the prompt.", + ) + input_variables: list[str] = Field( + default_factory=list, + description="Names of the Jinja2 template variables the prompt expects.", + ) + tools: list[ChatCompletionTool] | None = Field( + default=None, + description="Optional OpenAI-compatible tool definitions to send with the prompt.", + ) + tool_choice: str | dict[str, Any] | None = Field( + default=None, + description="Controls which (if any) tool is called: 'none', 'auto', 'required', or a named-tool object.", + ) + response_format: dict[str, Any] | None = Field( + default=None, + description="Optional OpenAI-compatible response_format, e.g. a json_schema structured-output spec.", + ) + inference_params: InferenceParams | None = Field( + default=None, + description="Optional default model and sampling parameters (temperature, top_p, max_tokens, ...).", + ) + tags: list[str] = Field( + default_factory=list, + description="Optional free-form tags for organizing prompts.", + ) + + +class PromptSort(StrEnum): + """Sort fields for Prompt queries.""" + + NAME_ASC = "name" + NAME_DESC = "-name" + CREATED_AT_ASC = "created_at" + CREATED_AT_DESC = "-created_at" + UPDATED_AT_ASC = "updated_at" + UPDATED_AT_DESC = "-updated_at" + + +class CreatePromptRequest(BaseModel): + """Request model for creating a Prompt.""" + + name: str = Field( + description=f"Name of the prompt. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["support-bot-system", "summarizer"], + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this prompt.", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field(default=None, max_length=1000) + messages: list[PromptMessage] = Field(default_factory=list) + input_variables: list[str] = Field(default_factory=list) + tools: list[ChatCompletionTool] | None = Field(default=None) + tool_choice: str | dict[str, Any] | None = Field(default=None) + response_format: dict[str, Any] | None = Field(default=None) + inference_params: InferenceParams | None = Field(default=None) + tags: list[str] | None = Field(default=None) + + +class UpdatePromptRequest(BaseModel): + """Request model for replacing a Prompt's mutable fields (full update). + + The prompt name and workspace come from the URL path and cannot be changed. + """ + + project: str | None = Field( + default=None, + description="The URN of the project associated with this prompt.", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field(default=None, max_length=1000) + messages: list[PromptMessage] = Field(default_factory=list) + input_variables: list[str] = Field(default_factory=list) + tools: list[ChatCompletionTool] | None = Field(default=None) + tool_choice: str | dict[str, Any] | None = Field(default=None) + response_format: dict[str, Any] | None = Field(default=None) + inference_params: InferenceParams | None = Field(default=None) + tags: list[str] | None = Field(default=None) + + +# --------------------------------------------------------------------------- +# Model entity + Adapter +# --------------------------------------------------------------------------- + + +class Adapter(BaseModel): + name: str = Field( + ..., + description=f"Name of the adapter. Name must be unique in the workspace for all Adapters and match the following regex: {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["lora-adapter-v1", "my-finetune"], + ) + workspace: str = Field( + ..., + description=f"Workspace of the adapter. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + ) + description: str | None = Field( + default=None, + description="Optional description of the adapter", + max_length=1000, + ) + fileset: str = Field( + ..., + description="Fileset where the adapter files are stored expected format {workspace}/{fileset_name}", + ) + finetuning_type: FinetuningType = Field(..., description="Type of finetuning (LORA, P_TUNING, etc.)") + enabled: bool = Field( + default=True, + description="Whether to make this adapter available for inference post training", + ) + lora_config: Lora | None = Field(None, description="Lora configuration specifics") + model: str | None = Field( + default=None, + description=f"Parent model entity reference. {MODEL_REF_PATTERN_DESCRIPTION}", + max_length=MODEL_REF_MAX_LEN, + ) + created_at: datetime = Field(default_factory=datetime.now) + updated_at: datetime = Field(default_factory=datetime.now) + + @field_validator("model") + @classmethod + def validate_model(cls, v: str | None) -> str | None: + if v is not None and not is_valid_model_ref(v): + raise ValueError(MODEL_REF_PATTERN_DESCRIPTION) + return v + + +class ModelEntity(ModelEntityBaseModel): + """A versioned model registered within the platform.""" + + project: str | None = Field( + default=None, + description="The URN of the project associated with this model entity.", + max_length=_MAX_LEN_255, + ) + description: str | None = Field( + default=None, + description="Optional description of the model.", + max_length=1000, + ) + spec: ModelSpec | None = Field(default=None, description="Detailed specification for the model") + finetuning_type: FinetuningType | None = Field(None, description="Set for full weight finetuned models") + fileset: str | None = Field( + default=None, + description="A set of checkpoint files, configs, and other auxiliary info associated with this model - expected format {workspace}/{fileset_name}", + ) + trust_remote_code: bool = Field( + default=False, + description="Whether to trust remote code to load this model checkpoint.", + ) + base_model: str | None = Field( + default=None, description="Link to another model which is used as a base for the current model" + ) + api_endpoint: APIEndpointData | None = Field( + default=None, description="Data about the inference endpoint for this model" + ) + backend_format: BackendFormat | None = Field( + default=None, + description=( + "Inference API wire format expected by the backend. If unset, inference routing treats the model as " + "OPENAI_CHAT." + ), + json_schema_extra={"nullable": True}, + ) + adapters: list[Adapter] | None = Field( + default=None, + description="Adapters that have been created against this model", + ) + prompt: PromptData | None = Field(default=None, description="Configuration for prompt engineering") + custom_fields: dict[str, Any] = Field(default_factory=dict, description="Custom fields for additional metadata") + ownership: dict[str, Any] | None = Field(default=None, description="Ownership information for the model") + model_providers: list[str] = Field( + default_factory=list, + description="List of ModelProvider workspace/name resource names that provide inference for this Model Entity", + ) + + +class CreateModelEntityRequest(BaseModel): + """Request model for creating a Model Entity.""" + + name: str = Field( + description=f"Name of the model entity. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["llama-3.1-8b", "my-custom-model"], + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this model entity", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field( + default=None, + description="Optional description of the model", + max_length=1000, + ) + spec: ModelSpec | None = Field( + default=None, + description="Detailed specification for the model - Automatically generated by the platform at creation when fileset provided.", + ) + finetuning_type: FinetuningType | None = Field(None, description="Set for full weight finetuned models") + fileset: str | None = Field( + default=None, + description="A set of checkpoint files, configs, and other auxiliary info associated with this model - expected format {workspace}/{fileset_name}", + ) + base_model: str | None = Field( + default=None, description="Link to another model which is used as a base for the current model" + ) + api_endpoint: APIEndpointData | None = Field( + default=None, description="Data about the inference endpoint for this model" + ) + backend_format: BackendFormat | None = Field( + default=None, + description=( + "Inference API wire format expected by the backend. If unset, inference routing treats the model as " + "OPENAI_CHAT." + ), + json_schema_extra={"nullable": True}, + ) + prompt: PromptData | None = Field(default=None, description="Configuration for prompt engineering") + custom_fields: dict[str, Any] | None = Field(default=None, description="Custom fields for additional metadata") + ownership: dict[str, Any] | None = Field(default=None, description="Ownership information for the model") + model_providers: list[str] | None = Field( + default_factory=list, + description="List of ModelProvider workspace/name resource names that provide inference for this Model Entity", + ) + trust_remote_code: bool = Field( + default=False, + description="Whether to trust remote code for the checkpoint.", + ) + + +class CreateModelAdapterRequest(BaseModel): + """Request body for nested Adapter creation. The base model comes from the URL path, not the body.""" + + name: str = Field( + ..., + description=f"Name of the adapter. Name must be unique in the workspace. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["lora-adapter-v1", "my-finetune"], + ) + description: str | None = Field( + default=None, + description="Optional description of the adapter", + max_length=1000, + ) + fileset: str = Field( + ..., + description="Location where adapter files are stored - expected format {workspace}/{fileset_name}", + ) + finetuning_type: FinetuningType = Field(..., description="Type of finetuning (LORA, P_TUNING, etc.)") + enabled: bool = Field( + default=True, + description="Whether to make this adapter available for inference post training", + ) + lora_config: Lora | None = Field(None, description="Lora configuration specifics") + + +class CreateAdapterRequest(CreateModelAdapterRequest): + """Request body for Adapter creation.""" + + model: str = Field( + ..., + max_length=MODEL_REF_MAX_LEN, + description=( + f"Base model entity. Use `{{workspace}}/{{model_name}}` to reference a model in any workspace, " + f"or a single `{{model_name}}` resolved in the path workspace. {MODEL_REF_PATTERN_DESCRIPTION}" + ), + examples=["llama-3-8b-instruct", "shared-tenant/base-llm"], + ) + + @field_validator("model") + @classmethod + def validate_model(cls, v: str | None) -> str | None: + if v is not None and not is_valid_model_ref(v): + raise ValueError(MODEL_REF_PATTERN_DESCRIPTION) + return v + + +class UpdateModelEntityRequest(BaseModel): + """Request model for updating Model Entity metadata.""" + + description: str | None = Field( + default=None, + description="Optional description of the model", + max_length=1000, + ) + spec: ModelSpec | None = Field(default=None, description="Detailed specification for the model") + fileset: str | None = Field( + default=None, + description="A set of checkpoint files, configs, and other auxiliary info associated with this model - expected format {workspace}/{fileset_name}", + ) + finetuning_type: FinetuningType | None = Field(None, description="Set for full weight finetuned models") + base_model: str | None = Field( + default=None, description="Link to another model which is used as a base for the current model" + ) + api_endpoint: APIEndpointData | None = Field( + default=None, description="Data about the inference endpoint for this model" + ) + backend_format: BackendFormat | None = Field( + default=None, + description=( + "Inference API wire format expected by the backend. If unset, inference routing treats the model as " + "OPENAI_CHAT." + ), + json_schema_extra={"nullable": True}, + ) + prompt: PromptData | None = Field(default=None, description="Configuration for prompt engineering") + custom_fields: dict[str, Any] | None = Field(default=None, description="Custom fields for additional metadata") + ownership: dict[str, Any] | None = Field(default=None, description="Ownership information for the model") + model_providers: list[str] | None = Field( + default=None, + description="List of ModelProvider workspace/name resource names that provide inference for this Model Entity", + ) + trust_remote_code: bool | None = Field( + default=None, + description="Whether to trust remote code for the checkpoint.", + ) + + +class UpdateAdapterRequest(BaseModel): + """Request model for updating Adapter Sub Entity metadata.""" + + description: str | None = Field( + default=None, + description="Optional description of the adapter", + max_length=1000, + ) + enabled: bool | None = Field( + default=None, + description="Whether to make this adapter available for inference post training", + ) + fileset: str | None = Field( + default=None, + description="Updated fileset for the adapter", + ) + + +class ModelEntitySortField(StrEnum): + """Sort fields for Model Entity queries.""" + + NAME_ASC = "name" + NAME_DESC = "-name" + CREATED_AT_ASC = "created_at" + CREATED_AT_DESC = "-created_at" + UPDATED_AT_ASC = "updated_at" + UPDATED_AT_DESC = "-updated_at" + + +# --------------------------------------------------------------------------- +# ModelDeploymentConfig + ModelDeployment +# --------------------------------------------------------------------------- + + +class ModelType(str, Enum): + """Model type enum for NIM deployments.""" + + LLM = "llm" + EMBED = "embed" + OTHER = "other" + + +class K8sNIMOperatorConfig(BaseModel): + """Kubernetes configuration for NIM deployment via k8s-nim-operator.""" + + resources: dict[str, Any] | None = Field( + default=None, + description="Kubernetes resource requirements including requests and limits. " + "Example: {'requests': {'cpu': '2', 'memory': '8Gi'}, 'limits': {'memory': '16Gi'}}", + ) + tolerations: list[dict[str, Any]] | None = Field( + default=None, + description="Kubernetes tolerations for pod scheduling. " + "Example: [{'key': 'nvidia.com/gpu', 'operator': 'Exists', 'effect': 'NoSchedule'}]", + ) + node_selector: dict[str, str] | None = Field( + default=None, + description="Kubernetes node selector for pod placement. " + "Example: {'node-type': 'gpu-node', 'zone': 'us-west1-a'}", + ) + startup_probe_grace_seconds: int | None = Field( + default=None, + description="Grace period in seconds for NIM startup. " + "Determines how long Kubernetes will wait for the NIM to become ready before restarting it. " + "Example: 600 (10 minutes). " + "Must be a positive integer.", + gt=0, + ) + + +class Engine(str, Enum): + """Inference engine selecting the compiler path for a deployment.""" + + NIM = "nim" + VLLM = "vllm" + GENERIC = "generic" + + +class ModelDeploymentConfigModelSpec(BaseModel): + """What model to serve and how -- independent of the executor it runs on.""" + + model_type: ModelType | None = Field(default=None, description="Type of model being deployed") + model_namespace: str | None = Field( + default=None, + description="Model repository namespace - organization/user namespace as it exists in repo_id.", + max_length=_MAX_LEN_255, + ) + model_name: str | None = Field( + default=None, + description="Model name - model repository name for model weights.", + max_length=_MAX_LEN_255, + ) + model_revision: str | None = Field( + default=None, + description="Model revision (branch, tag, or commit). If not specified, parsed from model_name @revision suffix or defaults to 'main'", + max_length=_MAX_LEN_255, + ) + chat_template: str | None = Field( + default=None, + description="Jinja2 chat template string for the model. Overrides the chat_template from ModelEntity.spec " + "if both are set. Used by the engine to format chat completions.", + ) + tool_call_config: ToolCallConfig | None = Field( + default=None, + description="Tool calling configuration for the deployment. Overrides tool_call_config from " + "ModelEntity.spec if both are set. Controls how the model handles function/tool calling.", + ) + lora_enabled: bool = Field(default=False, description="Whether to enable LoRA support") + + +class ContainerExecutorConfig(BaseModel): + """Compute + container settings shared by the docker and k8s executors.""" + + gpu: int = Field(description="Number of GPUs required for the deployment. 0 = CPU-only.", ge=0) + disk_size: str = Field(default="50Gi", description="Disk size for the deployment") + image_name: str | None = Field( + default=None, + description="Container image name. If not specified, defaults to the engine's configured image " + "(e.g. default_vllm_image / default_nimservice_image). Required for engine='generic'.", + max_length=_MAX_LEN_255, + ) + image_tag: str | None = Field( + default=None, + description="Container image tag. If not specified, defaults to the engine's configured image tag.", + max_length=_MAX_LEN_255, + ) + health_check_path: str | None = Field( + default=None, + description="HTTP path used for the container readiness probe. If not specified, defaults to the " + "engine's standard health endpoint (e.g. '/v1/health/ready' for NIM, '/health' for vLLM). " + "Set this for engine='generic' containers that expose a non-standard health endpoint.", + max_length=_MAX_LEN_255, + ) + run_as_user: int | None = Field( + default=None, + ge=0, + description="Pod securityContext runAsUser (uid) for the serving container (k8s backend only). " + "If unset, the engine default applies (vLLM pins its image's user; generic uses the image's " + "own user). Ignored by the docker backend.", + ) + run_as_group: int | None = Field( + default=None, + ge=0, + description="Pod securityContext runAsGroup (gid) for the serving container (k8s backend only). " + "If unset, the engine default applies. Ignored by the docker backend.", + ) + additional_envs: dict[str, str] | None = Field( + default=None, description="Additional environment variables for the deployment" + ) + additional_args: list[str] = Field( + default_factory=list, + description="Raw container/`serve` args appended verbatim to the container's arg vector.", + ) + k8s_nim_operator_config: K8sNIMOperatorConfig | None = Field( + default=None, + description="Typed Kubernetes configuration for common NIMService Spec fields (NIM engine on k8s). " + "Applied after defaults but before override_config. Ignored by non-NIM engines.", + ) + override_config: dict[str, Any] | None = Field( + default=None, + description="Raw NIMService spec configuration that takes precedence over generated config (NIM engine " + "on k8s). Allows advanced configuration options directly. Ignored by non-NIM engines.", + ) + + +class ModelDeploymentStatus(str, Enum): + """Status enum for ModelDeployment objects.""" + + UNKNOWN = "UNKNOWN" # Terminal + CREATED = "CREATED" + PENDING = "PENDING" + READY = "READY" + ERROR = "ERROR" # Terminal + DELETING = "DELETING" + DELETED = "DELETED" # Terminal + LOST = "LOST" # Terminal + + +class ModelDeploymentStatusHistoryItem(BaseModel): + """Record of a status change in ModelDeployment history.""" + + timestamp: datetime = Field(description="When this status was recorded") + status: ModelDeploymentStatus = Field(description="The status at this point in time") + status_message: str = Field(default="", description="Status message", max_length=1000) + + +class ModelDeploymentConfig(ModelEntityBaseModel): + """Immutable, automatically-versioned deployment config. + + The unique identifier is the combination of workspace/name/entity_version. + """ + + id: str = Field(default="", description="Unique identifier for the deployment config") + entity_version: int = Field(description="Version of this deployment config. Automatically managed.") + description: str | None = Field( + default=None, + description="Optional description of the deployment configuration", + max_length=1000, + ) + engine: Engine = Field(description="Inference engine selecting the compiler path (nim/vllm/generic)") + model_spec: ModelDeploymentConfigModelSpec = Field( + description="What model to serve and how -- independent of the executor it runs on" + ) + executor_config: ContainerExecutorConfig = Field( + description="Compute + container settings for the executor the deployment runs on" + ) + model_entity_id: str | None = Field( + default=None, + description="Optional reference to the base model entity ID for this deployment", + max_length=_MAX_LEN_255, + ) + + +class ModelDeployment(ModelEntityBaseModel): + """A deployed instance of a model with a specific configuration. + + The unique identifier is the combination of workspace/name/entity_version. + """ + + id: str = Field(default="", description="Unique identifier for the deployment") + entity_version: int = Field(description="Version of this deployment. Automatically managed.") + config: str = Field( + description="Reference to the ModelDeploymentConfig name", + max_length=_MAX_LEN_255, + ) + config_version: int = Field(description="Reference to the specific ModelDeploymentConfig version") + status: ModelDeploymentStatus = Field( + default=ModelDeploymentStatus.UNKNOWN, + description="Current status of the deployment, populated by models controller", + ) + status_message: str = Field( + default="", + description="Detailed status message, populated by models controller", + max_length=1000, + ) + status_history: list[ModelDeploymentStatusHistoryItem] = Field( + default_factory=list, + description="History of status changes, ordered chronologically (oldest first)", + ) + model_provider_id: str | None = Field( + default=None, + description="Optional reference to the auto-created ModelProvider workspace/name (format: workspace/name)", + max_length=_MAX_LEN_255, + ) + auth_context: AuthContext | None = Field(default=None, description="Auth context captured at deployment creation. ") + + +class CreateModelDeploymentConfigRequest(BaseModel): + """Request model for creating a ModelDeploymentConfig.""" + + name: str = Field( + description=f"Name of the deployment configuration. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["nim-config-v1", "production-config"], + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this deployment configuration", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + description: str | None = Field( + default=None, + description="Optional description of the deployment configuration", + max_length=1000, + ) + engine: Engine = Field(description="Inference engine selecting the compiler path (nim/vllm/generic)") + model_spec: ModelDeploymentConfigModelSpec = Field( + description="What model to serve and how -- independent of the executor it runs on" + ) + executor_config: ContainerExecutorConfig = Field( + description="Compute + container settings for the executor the deployment runs on" + ) + model_entity_id: str | None = Field( + default=None, + description="Optional reference to the base model entity ID for this deployment", + max_length=_MAX_LEN_255, + ) + + +class UpdateModelDeploymentConfigRequest(BaseModel): + """Request model for updating a ModelDeploymentConfig (creates new version).""" + + description: str | None = Field( + default=None, + description="Optional description of the deployment configuration", + max_length=1000, + ) + engine: Engine = Field(description="Inference engine selecting the compiler path (nim/vllm/generic)") + model_spec: ModelDeploymentConfigModelSpec = Field( + description="What model to serve and how -- independent of the executor it runs on" + ) + executor_config: ContainerExecutorConfig = Field( + description="Compute + container settings for the executor the deployment runs on" + ) + model_entity_id: str | None = Field( + default=None, + description="Optional reference to the base model entity ID for this deployment", + max_length=_MAX_LEN_255, + ) + + +class CreateModelDeploymentRequest(BaseModel): + """Request model for creating a ModelDeployment.""" + + name: str = Field( + description=f"Name of the deployment. {_NAME_DESC}", + max_length=_MAX_LEN_255, + pattern=_NAME_REGEX, + examples=["llama-deploy-v1", "production-nim"], + ) + project: str | None = Field( + default=None, + description="The URN of the project associated with this deployment", + max_length=_MAX_LEN_255, + pattern=_NAME_SLASH_REGEX, + ) + config: str = Field( + description="Reference to the ModelDeploymentConfig name", + max_length=_MAX_LEN_255, + ) + config_version: int | None = Field( + default=None, + description="Reference to a specific ModelDeploymentConfig version. If not specified, uses latest.", + ) + + +class UpdateModelDeploymentRequest(BaseModel): + """Request model for updating a ModelDeployment (creates new version).""" + + config: str = Field( + description="Reference to the ModelDeploymentConfig name", + max_length=_MAX_LEN_255, + ) + config_version: int | None = Field( + default=None, + description="Reference to a specific ModelDeploymentConfig version. If not specified, uses latest.", + ) + + +class UpdateModelDeploymentStatusRequest(BaseModel): + """Request model for updating ModelDeployment status.""" + + status: ModelDeploymentStatus = Field(description="New status for the deployment") + status_message: str = Field(default="", description="Detailed status message", max_length=1000) + model_provider_id: str | None = Field( + default=None, + description="Optional reference to the auto-created ModelProvider workspace/name (format: workspace/name)", + max_length=_MAX_LEN_255, + ) + + +# --------------------------------------------------------------------------- +# Query parameter types +# --------------------------------------------------------------------------- + + +class ListModelsQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + filter: NotRequired[str] + verbose: NotRequired[bool] + + +class GetModelQueryParams(TypedDict, total=False): + verbose: NotRequired[bool] + + +class ListAdaptersQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + filter: NotRequired[str] + + +class ListProvidersQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + filter: NotRequired[str] + + +class ListPromptsQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + filter: NotRequired[str] + + +class ListDeploymentsQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + all_versions: NotRequired[bool] + filter: NotRequired[str] + + +class ListDeploymentConfigsQueryParams(TypedDict, total=False): + page: NotRequired[int] + page_size: NotRequired[int] + sort: NotRequired[str] + filter: NotRequired[str] + + +class UpdateDeploymentStatusQueryParams(TypedDict, total=False): + version: NotRequired[str] diff --git a/packages/nemo_platform_plugin/tests/client/test_adapter.py b/packages/nemo_platform_plugin/tests/client/test_adapter.py index e51c4fd19d..7d3f3d9e6f 100644 --- a/packages/nemo_platform_plugin/tests/client/test_adapter.py +++ b/packages/nemo_platform_plugin/tests/client/test_adapter.py @@ -6,11 +6,12 @@ import httpx from nemo_platform import NeMoPlatform from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.types import RetryPolicy from nemo_platform_plugin.jobs import endpoints from nemo_platform_plugin.jobs.client import JobsClient -def test_client_from_platform_preserves_retry_count_with_nemoclient_defaults() -> None: +def test_client_from_platform_preserves_stainless_retry_policy() -> None: http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) platform = NeMoPlatform( base_url="http://test", @@ -22,8 +23,13 @@ def test_client_from_platform_preserves_retry_count_with_nemoclient_defaults() - client = client_from_platform(platform, JobsClient) assert client.retry is not None - assert client.retry.max_retries == 4 - assert client.retry.retryable_status_codes == (502, 503, 504, 429) + assert client.retry == RetryPolicy( + max_retries=4, + retryable_status_codes=(408, 409, 429), + retry_all_server_errors=True, + respect_retry_decision_headers=True, + respect_retry_after_headers=True, + ) def test_client_from_platform_prefers_platform_request_router() -> None: 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 b0dee0fdf9..1dc4b83e94 100644 --- a/packages/nemo_platform_plugin/tests/client/test_client_options.py +++ b/packages/nemo_platform_plugin/tests/client/test_client_options.py @@ -5,7 +5,7 @@ from __future__ import annotations -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -23,6 +23,14 @@ from pydantic import BaseModel BASE = "http://test:8000" +STAINLESS_RETRY = RetryPolicy( + max_retries=1, + backoff_base=0.25, + retryable_status_codes=(408, 409, 429), + retry_all_server_errors=True, + respect_retry_decision_headers=True, + respect_retry_after_headers=True, +) class ItemRequest(BaseModel): @@ -459,6 +467,180 @@ def test_no_retry_without_policy(self) -> None: assert exc_info.value.status_code == 503 assert mock_http.request.call_count == 1 + def test_standalone_retry_policy_defaults_are_unchanged(self) -> None: + policy = RetryPolicy() + + assert policy.retryable_status_codes == (502, 503, 504, 429) + assert policy.retry_all_server_errors is False + assert policy.respect_retry_decision_headers is False + assert policy.respect_retry_after_headers is False + + @pytest.mark.parametrize("status_code", [408, 409, 500]) + def test_standalone_policy_does_not_add_stainless_statuses(self, status_code: int) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + status_code, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "error"}, + ) + client = NemoClient( + base_url=BASE, + http_client=mock_http, + retry=RetryPolicy(max_retries=1, backoff_base=0.0), + ) + + with pytest.raises(NemoHTTPError): + client.send(GET_ITEM(name="alice")) + + assert mock_http.request.call_count == 1 + + @pytest.mark.parametrize("status_code", [408, 409, 500]) + def test_stainless_policy_retries_all_expected_statuses(self, status_code: int) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + status_code, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "error"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with patch("nemo_platform_plugin.client.client.time.sleep"): + response = client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 2 + + def test_stainless_true_header_forces_retry(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + 400, + headers={"x-should-retry": "true"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "retry"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with patch("nemo_platform_plugin.client.client.time.sleep"): + response = client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 2 + + def test_stainless_true_header_does_not_retry_success(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 200, + headers={"x-should-retry": "true"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ) + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + response = client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 1 + + def test_stainless_false_header_suppresses_retry(self) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.return_value = httpx.Response( + 500, + headers={"x-should-retry": "false"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "do not retry"}, + ) + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with pytest.raises(NemoHTTPError): + client.send(GET_ITEM(name="alice")) + + assert mock_http.request.call_count == 1 + + @pytest.mark.parametrize( + ("headers", "expected_delay"), + [ + ({"retry-after-ms": "1250"}, 1.25), + ({"retry-after-ms": "invalid", "retry-after": "3"}, 3.0), + ({"retry-after": "2.5"}, 2.5), + ({"retry-after": "Mon, 12 Jan 1970 13:47:10 GMT"}, 30.0), + ({"retry-after": "60"}, 60.0), + ], + ) + def test_stainless_policy_honors_retry_after(self, headers: dict[str, str], expected_delay: float) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + 500, + headers=headers, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "retry"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with ( + patch("nemo_platform_plugin.client.client.time.time", return_value=1_000_000), + patch("nemo_platform_plugin.client.client.time.sleep") as sleep, + ): + client.send(GET_ITEM(name="alice")) + + sleep.assert_called_once_with(expected_delay) + + @pytest.mark.parametrize( + "headers", + [ + {"retry-after-ms": "0"}, + {"retry-after-ms": "60001"}, + {"retry-after": "-1"}, + {"retry-after": "60.1"}, + {"retry-after": "not-a-delay"}, + {"retry-after": "Mon, 12 Jan 1970 13:46:39 GMT"}, + ], + ) + def test_stainless_policy_falls_back_for_unreasonable_retry_after(self, headers: dict[str, str]) -> None: + mock_http = MagicMock(spec=httpx.Client) + mock_http.request.side_effect = [ + httpx.Response( + 500, + headers=headers, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "retry"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = NemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with ( + patch("nemo_platform_plugin.client.client.time.time", return_value=1_000_000), + patch("nemo_platform_plugin.client.client.time.sleep") as sleep, + ): + client.send(GET_ITEM(name="alice")) + + sleep.assert_called_once_with(STAINLESS_RETRY.backoff_base) + def test_binary_stream_retries_before_returning_content(self) -> None: attempts = 0 @@ -540,6 +722,86 @@ async def test_retry_on_503_async(self) -> None: assert resp.body.name == "alice" assert mock_http.request.call_count == 2 + @pytest.mark.asyncio + @pytest.mark.parametrize("status_code", [408, 409, 500]) + async def test_stainless_policy_retries_expected_statuses_async(self, status_code: int) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.side_effect = [ + httpx.Response( + status_code, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "retry"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = AsyncNemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with patch("nemo_platform_plugin.client.client.asyncio.sleep", new_callable=AsyncMock): + response = await client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 2 + + @pytest.mark.asyncio + async def test_stainless_true_header_forces_retry_async(self) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.side_effect = [ + httpx.Response( + 400, + headers={"x-should-retry": "true"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "retry"}, + ), + httpx.Response( + 200, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ), + ] + client = AsyncNemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with patch("nemo_platform_plugin.client.client.asyncio.sleep", new_callable=AsyncMock): + response = await client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 2 + + @pytest.mark.asyncio + async def test_stainless_true_header_does_not_retry_success_async(self) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.return_value = httpx.Response( + 200, + headers={"x-should-retry": "true"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"id": 1, "name": "alice"}, + ) + client = AsyncNemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + response = await client.send(GET_ITEM(name="alice")) + + assert response.body.name == "alice" + assert mock_http.request.call_count == 1 + + @pytest.mark.asyncio + async def test_stainless_false_header_suppresses_retry_async(self) -> None: + mock_http = AsyncMock(spec=httpx.AsyncClient) + mock_http.request.return_value = httpx.Response( + 500, + headers={"x-should-retry": "false"}, + request=httpx.Request("GET", f"{BASE}/apis/test/v2/items/alice"), + json={"detail": "do not retry"}, + ) + client = AsyncNemoClient(base_url=BASE, http_client=mock_http, retry=STAINLESS_RETRY) + + with pytest.raises(NemoHTTPError): + await client.send(GET_ITEM(name="alice")) + + assert mock_http.request.call_count == 1 + @pytest.mark.asyncio async def test_exhausted_transport_error_is_wrapped_async(self) -> None: mock_http = AsyncMock(spec=httpx.AsyncClient) diff --git a/packages/nemo_platform_plugin/tests/client/test_method.py b/packages/nemo_platform_plugin/tests/client/test_method.py new file mode 100644 index 0000000000..b54ed711d6 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_method.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for the ``method()`` descriptor that binds endpoints onto client classes. + +Instance-level dispatch is exercised throughout the client and Models suites. What +is pinned here is *class*-level access, which nothing else touches and which every +introspection tool performs: ``Mock(spec=SomeClient)``, ``help()``, ``pydoc``, +autodoc, and anything walking ``dir()``. +""" + +from __future__ import annotations + +import inspect +import pydoc +from unittest.mock import MagicMock, create_autospec + +import pytest +from nemo_platform_plugin.client.method import EndpointMethod +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient + +BASE = "http://test:8000" + + +def _descriptor(name: str) -> EndpointMethod: + return inspect.getattr_static(AsyncModelsClient, name) + + +# --------------------------------------------------------------------------- +# Class-level access +# --------------------------------------------------------------------------- + + +def test_class_level_access_returns_the_descriptor() -> None: + """``__get__`` with no instance hands back the descriptor, per the protocol. + + It used to assert ``obj is not None``, so every one of these raised. + """ + assert isinstance(ModelsClient.create_model, EndpointMethod) + assert isinstance(AsyncModelsClient.create_model, EndpointMethod) + assert ModelsClient.create_model is inspect.getattr_static(ModelsClient, "create_model") + + +def test_mock_spec_against_a_client_class_builds() -> None: + """``Mock(spec=...)`` reads every attribute off the class to classify it. + + This is the failure that surfaced the bug: the models controller test suite + could not spec a mock against its own client. + """ + mock = MagicMock(spec=AsyncModelsClient) + + assert callable(mock.create_model) + with pytest.raises(AttributeError): + mock.create_modle # noqa: B018 a typo must not be silently mockable + + +def test_pydoc_lists_endpoints_with_their_docstrings() -> None: + """Endpoint docs reach help() through the copied ``__doc__``. + + pydoc swallows per-member errors, so this does not discriminate the + class-access fix; what it pins is the attribute copying. + """ + # plain() strips pydoc's backspace-overstrike bolding. + rendered = pydoc.plain(pydoc.render_doc(ModelsClient)) + + assert "delete_deployment" in rendered + assert "Delete a deployment" in rendered + + +def test_create_autospec_does_not_raise() -> None: + """autospec also walks the class. + + The resulting endpoint stubs are *not* callable, because the descriptor is not. + That is a real limitation of this design and is pinned here so it is a + deliberate trade-off rather than a surprise: use ``spec=`` for client mocks. + """ + auto = create_autospec(AsyncModelsClient) + + assert not callable(auto.create_model) + + +# --------------------------------------------------------------------------- +# Identity carried from the endpoint +# --------------------------------------------------------------------------- + + +def test_descriptor_carries_endpoint_identity() -> None: + descriptor = _descriptor("delete_deployment") + + assert descriptor.__name__ == "delete_deployment" + assert descriptor.__doc__ == descriptor.endpoint.__doc__ + assert descriptor.__doc__ # the endpoint really does carry one + assert descriptor.__wrapped__ is descriptor.endpoint + + +def test_signature_is_reachable_by_unwrapping_at_class_level() -> None: + """``inspect.signature`` rejects the descriptor; ``unwrap`` gets past it. + + Pinned because the obvious reading of ``__wrapped__`` is that ``signature()`` + follows it. It does not: it refuses a non-callable before ever looking. + """ + # Both calls are typed as taking a callable; passing the descriptor is the + # behaviour under test, hence the suppressions rather than a cast. + with pytest.raises(TypeError): + inspect.signature(AsyncModelsClient.create_model) # ty: ignore[invalid-argument-type] + + signature = inspect.signature(inspect.unwrap(AsyncModelsClient.create_model)) # ty: ignore[invalid-argument-type] + + assert set(signature.parameters) == {"workspace", "body", "exist_ok"} + assert all(p.kind is inspect.Parameter.KEYWORD_ONLY for p in signature.parameters.values()) + + +def test_descriptor_does_not_leak_the_endpoint_abstractmethod_marker() -> None: + """Endpoints are ``@abstractmethod`` stubs; that marker must not ride along. + + ``functools.update_wrapper`` would copy ``__dict__`` and with it + ``__isabstractmethod__``, which would make any ABCMeta-based client class + uninstantiable. The attributes are copied one by one to avoid exactly that. + """ + descriptor = _descriptor("create_model") + + assert getattr(descriptor.endpoint, "__isabstractmethod__", False) is True + assert not getattr(descriptor, "__isabstractmethod__", False) + assert not getattr(AsyncModelsClient, "__abstractmethods__", frozenset()) + ModelsClient(base_url=BASE, workspace="default") # constructs + + +# --------------------------------------------------------------------------- +# Instance-level dispatch (unchanged, but nothing states it outright) +# --------------------------------------------------------------------------- + + +def test_sync_client_binds_a_plain_callable() -> None: + bound = ModelsClient(base_url=BASE, workspace="default").create_model + + assert callable(bound) + assert not inspect.iscoroutinefunction(bound) + + +def test_async_client_binds_a_coroutine_function() -> None: + bound = AsyncModelsClient(base_url=BASE, workspace="default").create_model + + assert inspect.iscoroutinefunction(bound) + + +def test_bound_method_exposes_the_real_signature() -> None: + """Instance access hands back the wrapped function, so signature() works there.""" + signature = inspect.signature(ModelsClient(base_url=BASE, workspace="default").get_model) + + assert set(signature.parameters) == {"workspace", "name", "query_params"} diff --git a/packages/nemo_platform_plugin/tests/models/test_client.py b/packages/nemo_platform_plugin/tests/models/test_client.py new file mode 100644 index 0000000000..fc78574e27 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/models/test_client.py @@ -0,0 +1,405 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for ModelsClient / AsyncModelsClient via mocked httpx transport.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from nemo_platform_plugin.client.errors import ConflictError, NotFoundError +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient +from nemo_platform_plugin.models.types import ( + CreateModelDeploymentRequest, + CreateModelEntityRequest, + CreateModelProviderRequest, + ModelDeployment, + ModelEntity, + ModelProvider, +) + +BASE = "http://test:8000" + + +def _model_json(name: str = "llama", workspace: str = "default", **extra: object) -> dict: + base = { + "id": f"model-{name}", + "name": name, + "workspace": workspace, + "created_at": "2020-01-01T00:00:00Z", + "updated_at": "2020-01-01T00:00:00Z", + } + base.update(extra) + return base + + +def _provider_json(status: str = "PENDING", name: str = "my-provider", **extra: object) -> dict: + base = { + "id": f"provider-{name}", + "name": name, + "workspace": "default", + "host_url": "https://api.example.com", + "status": status, + "status_message": "", + "created_at": "2020-01-01T00:00:00Z", + "updated_at": "2020-01-01T00:00:00Z", + } + base.update(extra) + return base + + +def _deployment_json(status: str = "PENDING", history: list | None = None, **extra: object) -> dict: + base = { + "id": "dep-1", + "name": "my-deploy", + "workspace": "default", + "entity_version": 1, + "config": "cfg", + "config_version": 1, + "status": status, + "status_message": "", + "status_history": history if history is not None else [], + "created_at": "2020-01-01T00:00:00Z", + "updated_at": "2020-01-01T00:00:00Z", + } + base.update(extra) + return base + + +# --------------------------------------------------------------------------- +# Round-trips +# --------------------------------------------------------------------------- + + +def test_create_model_round_trip() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(201, request=httpx.Request("POST", BASE), json=_model_json()) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + out = client.create_model(body=CreateModelEntityRequest(name="llama")).data() + + assert isinstance(out, ModelEntity) + assert out.name == "llama" + args, kwargs = http.request.call_args + assert args == ("POST", f"{BASE}/apis/models/v2/workspaces/default/models") + assert kwargs["content"] == b'{"name":"llama"}' + + +def test_list_models_paginates_across_pages() -> None: + http = MagicMock(spec=httpx.Client) + page1 = { + "data": [_model_json("a"), _model_json("b")], + "pagination": { + "page": 1, + "page_size": 2, + "current_page_size": 2, + "total_pages": 2, + "total_results": 3, + }, + } + page2 = { + "data": [_model_json("c")], + "pagination": { + "page": 2, + "page_size": 2, + "current_page_size": 1, + "total_pages": 2, + "total_results": 3, + }, + } + http.request.side_effect = [ + httpx.Response(200, request=httpx.Request("GET", BASE), json=page1), + httpx.Response(200, request=httpx.Request("GET", BASE), json=page2), + ] + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + names = [m.name for m in client.list_models().items()] + assert names == ["a", "b", "c"] + + +def test_get_model_not_found_raises() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(404, request=httpx.Request("GET", BASE), json={"detail": "not found"}) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + with pytest.raises(NotFoundError) as exc: + client.get_model(name="missing") + assert exc.value.status_code == 404 + + +def test_delete_deployment_returns_none_on_202() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(202, request=httpx.Request("DELETE", BASE)) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + resp = client.delete_deployment(name="d") + assert resp.data() is None + assert resp.http_response.status_code == 202 + + +def test_delete_deployment_returns_none_on_204() -> None: + """Synchronous hard-delete: 204 No Content is success with a ``None`` body.""" + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(204, request=httpx.Request("DELETE", BASE)) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + resp = client.delete_deployment(name="d") + assert resp.data() is None + # The 202/204 distinction is only observable via the raw status code. + assert resp.http_response.status_code == 204 + + +def test_delete_deployment_version_accepts_202_and_204() -> None: + http = MagicMock(spec=httpx.Client) + http.request.side_effect = [ + httpx.Response(202, request=httpx.Request("DELETE", BASE)), + httpx.Response(204, request=httpx.Request("DELETE", BASE)), + ] + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.delete_deployment_version(deployment="d", name="1").data() is None + assert client.delete_deployment_version(deployment="d", name="2").data() is None + + +def test_create_deployment_conflict_without_exist_ok_raises() -> None: + """Default exist_ok=False: a 409 surfaces as ConflictError (no GET replay).""" + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(409, request=httpx.Request("POST", BASE), json={"detail": "exists"}) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + with pytest.raises(ConflictError) as exc: + client.create_deployment(body=CreateModelDeploymentRequest(name="d", config="cfg")) + assert exc.value.status_code == 409 + # A single POST was made; no conflict-resolving GET replay happened. + assert http.request.call_count == 1 + + +def test_create_provider_exist_ok_resolves_conflict() -> None: + """exist_ok=True: a 409 replays the linked GET and returns the existing entity.""" + http = MagicMock(spec=httpx.Client) + conflict = httpx.Response(409, request=httpx.Request("POST", BASE), json={"detail": "exists"}) + existing = httpx.Response( + 200, + request=httpx.Request("GET", BASE), + json=_model_json("p", host_url="http://x") | {"host_url": "http://x"}, + ) + http.request.side_effect = [conflict, existing] + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + out = client.create_provider(body=CreateModelProviderRequest(name="p", host_url="http://x"), exist_ok=True).data() + + assert isinstance(out, ModelProvider) + assert out.name == "p" + # Second call is the GET replay for the existing provider. + assert http.request.call_args_list[1].args[0] == "GET" + assert http.request.call_args_list[1].args[1].endswith("/providers/p") + + +# --------------------------------------------------------------------------- +# URL builders +# --------------------------------------------------------------------------- + + +def test_openai_route_base_url() -> None: + client = ModelsClient(base_url=BASE + "/", workspace="default") + assert client.get_openai_route_base_url() == f"{BASE}/apis/inference-gateway/v2/workspaces/default/openai/-/v1" + assert ( + client.get_openai_route_base_url(workspace="other") + == f"{BASE}/apis/inference-gateway/v2/workspaces/other/openai/-/v1" + ) + + +def test_openai_route_base_url_missing_workspace_raises() -> None: + client = ModelsClient(base_url=BASE) + with pytest.raises(ValueError, match="Missing workspace"): + client.get_openai_route_base_url() + + +def test_provider_route_appends_v1_conditionally() -> None: + client = ModelsClient(base_url=BASE, workspace="default") + p_openai = ModelProvider.model_validate(_model_json("p", host_url="https://api.openai.com")) + p_nim = ModelProvider.model_validate(_model_json("p", host_url="https://nim.example.com/v1")) + + assert client.get_provider_route_openai_url(p_openai).endswith("/provider/p/-/v1") + assert client.get_provider_route_openai_url(p_nim).endswith("/provider/p/-") + + +def test_model_entity_route_always_v1() -> None: + client = ModelsClient(base_url=BASE, workspace="default") + me = ModelEntity.model_validate(_model_json("m")) + assert client.get_model_entity_route_openai_url(me).endswith("/model/m/-/v1") + + +def test_provider_route_for_deployment_fetches_provider() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response( + 200, + request=httpx.Request("GET", BASE), + json=_model_json("my-provider", host_url="https://api.example.com"), + ) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + deployment = ModelDeployment.model_validate(_deployment_json(model_provider_id="default/my-provider")) + + url = client.get_provider_route_openai_url_for_deployment(deployment) + assert url.endswith("/workspaces/default/provider/my-provider/-/v1") + assert http.request.call_args.args[1].endswith("/providers/my-provider") + + +def test_provider_route_for_deployment_without_provider_id_raises() -> None: + client = ModelsClient(base_url=BASE, workspace="default") + deployment = ModelDeployment.model_validate(_deployment_json(model_provider_id=None)) + with pytest.raises(ValueError, match="no associated model_provider_id"): + client.get_provider_route_openai_url_for_deployment(deployment) + + +# --------------------------------------------------------------------------- +# Deployment polling +# --------------------------------------------------------------------------- + + +def test_wait_for_deployment_status_reaches_ready() -> None: + http = MagicMock(spec=httpx.Client) + http.request.side_effect = [ + httpx.Response(200, request=httpx.Request("GET", BASE), json=_deployment_json("PENDING")), + httpx.Response(200, request=httpx.Request("GET", BASE), json=_deployment_json("READY")), + ] + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_deployment_status("my-deploy", "READY", poll_interval=0.0) is True + + +def test_wait_for_deployment_status_deleted_on_404() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(404, request=httpx.Request("GET", BASE), json={"detail": "x"}) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_deployment_status("my-deploy", "DELETED", poll_interval=0.0) is True + + +def test_wait_for_deployment_status_error_returns_false() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response( + 200, request=httpx.Request("GET", BASE), json=_deployment_json("ERROR", status_message="boom") + ) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_deployment_status("my-deploy", "READY", poll_interval=0.0) is False + + +def test_wait_for_deployment_status_uses_history_tail() -> None: + history = [ + {"timestamp": "2020-01-01T00:00:01Z", "status": "PENDING", "status_message": ""}, + {"timestamp": "2020-01-01T00:00:05Z", "status": "READY", "status_message": "up"}, + ] + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response( + 200, request=httpx.Request("GET", BASE), json=_deployment_json("PENDING", history=history) + ) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + # Top-level status is PENDING but history tail is READY -> reached. + assert client.wait_for_deployment_status("my-deploy", "READY", poll_interval=0.0) is True + + +def test_wait_for_deployment_status_timeout() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response( + 200, request=httpx.Request("GET", BASE), json=_deployment_json("PENDING") + ) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_deployment_status("my-deploy", "READY", timeout=0, poll_interval=0.0) is False + + +# --------------------------------------------------------------------------- +# Provider polling +# --------------------------------------------------------------------------- + + +def test_wait_for_provider_status_reaches_ready() -> None: + http = MagicMock(spec=httpx.Client) + http.request.side_effect = [ + httpx.Response(200, request=httpx.Request("GET", BASE), json=_provider_json("PENDING")), + httpx.Response(200, request=httpx.Request("GET", BASE), json=_provider_json("READY")), + ] + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_provider_status("my-provider", "READY", poll_interval=0.0) is True + + +def test_wait_for_provider_status_error_returns_false() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response( + 200, request=httpx.Request("GET", BASE), json=_provider_json("ERROR", status_message="boom") + ) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_provider_status("my-provider", "READY", poll_interval=0.0) is False + + +def test_wait_for_provider_status_not_found_returns_false() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(404, request=httpx.Request("GET", BASE), json={"detail": "x"}) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_provider_status("my-provider", "READY", poll_interval=0.0) is False + + +def test_wait_for_provider_status_timeout() -> None: + http = MagicMock(spec=httpx.Client) + http.request.return_value = httpx.Response(200, request=httpx.Request("GET", BASE), json=_provider_json("PENDING")) + client = ModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert client.wait_for_provider_status("my-provider", "READY", timeout=0, poll_interval=0.0) is False + + +# --------------------------------------------------------------------------- +# Async +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_async_create_model() -> None: + http = AsyncMock(spec=httpx.AsyncClient) + http.request.return_value = httpx.Response(201, request=httpx.Request("POST", BASE), json=_model_json()) + client = AsyncModelsClient(base_url=BASE, workspace="default", http_client=http) + + out = (await client.create_model(body=CreateModelEntityRequest(name="llama"))).data() + assert out.name == "llama" + + +@pytest.mark.asyncio +async def test_async_wait_for_deployment_status_ready() -> None: + http = AsyncMock(spec=httpx.AsyncClient) + http.request.side_effect = [ + httpx.Response(200, request=httpx.Request("GET", BASE), json=_deployment_json("PENDING")), + httpx.Response(200, request=httpx.Request("GET", BASE), json=_deployment_json("READY")), + ] + client = AsyncModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert await client.wait_for_deployment_status("my-deploy", "READY", poll_interval=0.0) is True + + +@pytest.mark.asyncio +async def test_async_wait_for_provider_status_ready() -> None: + http = AsyncMock(spec=httpx.AsyncClient) + http.request.side_effect = [ + httpx.Response(200, request=httpx.Request("GET", BASE), json=_provider_json("PENDING")), + httpx.Response(200, request=httpx.Request("GET", BASE), json=_provider_json("READY")), + ] + client = AsyncModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert await client.wait_for_provider_status("my-provider", "READY", poll_interval=0.0) is True + + +@pytest.mark.asyncio +async def test_async_wait_for_provider_status_error_returns_false() -> None: + http = AsyncMock(spec=httpx.AsyncClient) + http.request.return_value = httpx.Response( + 200, request=httpx.Request("GET", BASE), json=_provider_json("ERROR", status_message="boom") + ) + client = AsyncModelsClient(base_url=BASE, workspace="default", http_client=http) + + assert await client.wait_for_provider_status("my-provider", "READY", poll_interval=0.0) is False diff --git a/packages/nemo_platform_plugin/tests/models/test_endpoints.py b/packages/nemo_platform_plugin/tests/models/test_endpoints.py new file mode 100644 index 0000000000..b45eb4b6af --- /dev/null +++ b/packages/nemo_platform_plugin/tests/models/test_endpoints.py @@ -0,0 +1,312 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for Models service endpoint definitions (PreparedRequest shape).""" + +from __future__ import annotations + +import json +from typing import get_origin + +from nemo_platform_plugin.client.types import Paginated, PreparedRequest +from nemo_platform_plugin.models import endpoints +from nemo_platform_plugin.models.types import ( + Adapter, + ContainerExecutorConfig, + CreateAdapterRequest, + CreateModelDeploymentConfigRequest, + CreateModelDeploymentRequest, + CreateModelEntityRequest, + CreateModelProviderRequest, + CreatePromptRequest, + Engine, + FinetuningType, + ModelDeployment, + ModelDeploymentConfig, + ModelDeploymentConfigModelSpec, + ModelDeploymentStatus, + ModelEntity, + ModelProvider, + ModelProviderStatus, + Prompt, + UpdateAdapterRequest, + UpdateModelDeploymentRequest, + UpdateModelDeploymentStatusRequest, + UpdateModelEntityRequest, + UpdateModelProviderStatusRequest, + UpdatePromptRequest, + UpsertModelProviderRequest, +) + +_PREFIX = "/apis/models/v2/workspaces/{workspace}" + + +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) + + +# --------------------------------------------------------------------------- +# Model entities +# --------------------------------------------------------------------------- + + +def test_create_model() -> None: + prepared = endpoints.create_model(workspace="default", body=CreateModelEntityRequest(name="llama")) + assert isinstance(prepared, PreparedRequest) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/models" + assert prepared.path_params == {"workspace": "default"} + assert prepared.content_type == "application/json" + assert prepared.response_type is ModelEntity + # exist_ok wiring: conflict resolver prebuilt, but not requested by default. + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.method == "GET" + assert prepared.on_conflict_get.path_params == {"workspace": "default", "name": "llama"} + + +def test_create_model_excludes_unset() -> None: + prepared = endpoints.create_model(workspace="w", body=CreateModelEntityRequest(name="m")) + assert _json_body(prepared) == {"name": "m"} + + +def test_create_model_workspace_optional() -> None: + prepared = endpoints.create_model(body=CreateModelEntityRequest(name="m")) + assert prepared.path_params == {} + + +def test_list_models_paginated_with_query() -> None: + prepared = endpoints.list_models( + workspace="default", query_params={"page": 2, "sort": "-created_at", "verbose": True} + ) + assert prepared.method == "GET" + assert prepared.path_template == _PREFIX + "/models" + assert get_origin(prepared.response_type) is Paginated + assert prepared.query_params == {"page": 2, "sort": "-created_at", "verbose": True} + + +def test_get_model_with_verbose() -> None: + prepared = endpoints.get_model(workspace="default", name="m", query_params={"verbose": True}) + assert prepared.method == "GET" + assert prepared.path_params == {"workspace": "default", "name": "m"} + assert prepared.query_params == {"verbose": True} + assert prepared.response_type is ModelEntity + + +def test_update_model_patch_with_verbose() -> None: + prepared = endpoints.update_model( + workspace="default", name="m", body=UpdateModelEntityRequest(description="d"), query_params={"verbose": False} + ) + assert prepared.method == "PATCH" + assert prepared.path_params == {"workspace": "default", "name": "m"} + assert _json_body(prepared) == {"description": "d"} + assert prepared.query_params == {"verbose": False} + + +def test_delete_model_returns_none() -> None: + prepared = endpoints.delete_model(workspace="default", name="m") + assert prepared.method == "DELETE" + assert prepared.content is None + assert prepared.response_type is None + + +# --------------------------------------------------------------------------- +# Adapters (nested + top-level) +# --------------------------------------------------------------------------- + + +def test_create_model_adapter_nested_path() -> None: + prepared = endpoints.create_model_adapter( + workspace="w", + model_name="base", + body=__import__( + "nemo_platform_plugin.models.types", fromlist=["CreateModelAdapterRequest"] + ).CreateModelAdapterRequest(name="a", fileset="w/fs", finetuning_type=FinetuningType.LORA), + ) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/models/{model_name}/adapters" + assert prepared.path_params == {"workspace": "w", "model_name": "base"} + assert prepared.response_type is Adapter + + +def test_update_model_adapter_path() -> None: + prepared = endpoints.update_model_adapter( + workspace="w", model_name="base", adapter="a", body=UpdateAdapterRequest(enabled=False) + ) + assert prepared.method == "PATCH" + assert prepared.path_params == {"workspace": "w", "model_name": "base", "adapter": "a"} + assert _json_body(prepared) == {"enabled": False} + + +def test_delete_model_adapter_path() -> None: + prepared = endpoints.delete_model_adapter(workspace="w", model_name="base", adapter="a") + assert prepared.method == "DELETE" + assert prepared.path_template == _PREFIX + "/models/{model_name}/adapters/{adapter}" + assert prepared.response_type is None + + +def test_create_adapter_top_level_conflict_resolver() -> None: + body = CreateAdapterRequest(name="a", fileset="w/fs", finetuning_type=FinetuningType.LORA, model="ws/base") + prepared = endpoints.create_adapter(workspace="w", body=body) + assert prepared.path_template == _PREFIX + "/adapters" + assert prepared.response_type is Adapter + assert prepared.on_conflict_get is not None + assert prepared.on_conflict_get.path_params == {"workspace": "w", "name": "a"} + + +def test_list_adapters_paginated() -> None: + prepared = endpoints.list_adapters(workspace="w", query_params={"filter": "name:a"}) + assert get_origin(prepared.response_type) is Paginated + assert prepared.query_params == {"filter": "name:a"} + + +def test_get_and_delete_adapter() -> None: + assert endpoints.get_adapter(workspace="w", name="a").response_type is Adapter + assert endpoints.delete_adapter(workspace="w", name="a").method == "DELETE" + + +# --------------------------------------------------------------------------- +# Model providers +# --------------------------------------------------------------------------- + + +def test_create_provider() -> None: + prepared = endpoints.create_provider(workspace="w", body=CreateModelProviderRequest(name="p", host_url="http://x")) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/providers" + assert prepared.response_type is ModelProvider + assert prepared.on_conflict_get is not None + + +def test_upsert_provider_is_put() -> None: + prepared = endpoints.upsert_provider(workspace="w", name="p", body=UpsertModelProviderRequest(host_url="http://x")) + assert prepared.method == "PUT" + assert prepared.path_template == _PREFIX + "/providers/{name}" + assert prepared.response_type is ModelProvider + + +def test_update_provider_status_is_put_status_path() -> None: + prepared = endpoints.update_provider_status( + workspace="w", name="p", body=UpdateModelProviderStatusRequest(status=ModelProviderStatus.READY) + ) + assert prepared.method == "PUT" + assert prepared.path_template == _PREFIX + "/providers/{name}/status" + + +def test_list_get_delete_provider() -> None: + assert get_origin(endpoints.list_providers(workspace="w").response_type) is Paginated + assert endpoints.get_provider(workspace="w", name="p").response_type is ModelProvider + assert endpoints.delete_provider(workspace="w", name="p").response_type is None + + +# --------------------------------------------------------------------------- +# Prompts +# --------------------------------------------------------------------------- + + +def test_prompt_crud_paths() -> None: + assert endpoints.create_prompt(workspace="w", body=CreatePromptRequest(name="p")).method == "POST" + assert endpoints.update_prompt(workspace="w", name="p", body=UpdatePromptRequest()).method == "PUT" + assert endpoints.get_prompt(workspace="w", name="p").response_type is Prompt + assert get_origin(endpoints.list_prompts(workspace="w").response_type) is Paginated + assert endpoints.delete_prompt(workspace="w", name="p").response_type is None + + +# --------------------------------------------------------------------------- +# Deployments +# --------------------------------------------------------------------------- + + +def _create_deployment_body() -> CreateModelDeploymentRequest: + return CreateModelDeploymentRequest(name="d", config="cfg") + + +def test_create_deployment() -> None: + prepared = endpoints.create_deployment(workspace="w", body=_create_deployment_body()) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/deployments" + assert prepared.response_type is ModelDeployment + + +def test_update_deployment_is_post_name_path() -> None: + prepared = endpoints.update_deployment(workspace="w", name="d", body=UpdateModelDeploymentRequest(config="cfg")) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/deployments/{name}" + + +def test_update_deployment_status_with_version_query() -> None: + prepared = endpoints.update_deployment_status( + workspace="w", + name="d", + body=UpdateModelDeploymentStatusRequest(status=ModelDeploymentStatus.READY), + query_params={"version": "2"}, + ) + assert prepared.method == "POST" + assert prepared.path_template == _PREFIX + "/deployments/{name}/status" + assert prepared.query_params == {"version": "2"} + + +def test_deployment_versions_and_models() -> None: + assert endpoints.list_deployment_versions(workspace="w", name="d").response_type == list[ModelDeployment] + assert endpoints.get_deployment_version(workspace="w", deployment="d", name="2").response_type is ModelDeployment + models_ep = endpoints.get_deployment_models(workspace="w", name="d") + assert models_ep.path_template == _PREFIX + "/deployments/{name}/models" + + +def test_delete_deployment_and_version_return_none() -> None: + assert endpoints.delete_deployment(workspace="w", name="d").response_type is None + assert ( + endpoints.delete_deployment_version(workspace="w", deployment="d", name="2").path_template + == _PREFIX + "/deployments/{deployment}/versions/{name}" + ) + + +# --------------------------------------------------------------------------- +# Deployment configs +# --------------------------------------------------------------------------- + + +def _create_config_body() -> CreateModelDeploymentConfigRequest: + return CreateModelDeploymentConfigRequest( + name="cfg", + engine=Engine.VLLM, + model_spec=ModelDeploymentConfigModelSpec(model_name="llama"), + executor_config=ContainerExecutorConfig(gpu=1), + ) + + +def test_deployment_config_crud_paths() -> None: + create = endpoints.create_deployment_config(workspace="w", body=_create_config_body()) + assert create.method == "POST" + assert create.path_template == _PREFIX + "/deployment-configs" + assert create.response_type is ModelDeploymentConfig + assert create.on_conflict_get is not None + + update = endpoints.update_deployment_config( + workspace="w", + name="cfg", + body=__import__( + "nemo_platform_plugin.models.types", fromlist=["UpdateModelDeploymentConfigRequest"] + ).UpdateModelDeploymentConfigRequest( + engine=Engine.VLLM, + model_spec=ModelDeploymentConfigModelSpec(model_name="llama"), + executor_config=ContainerExecutorConfig(gpu=1), + ), + ) + assert update.method == "POST" + assert update.path_template == _PREFIX + "/deployment-configs/{name}" + + assert ( + endpoints.list_deployment_config_versions(workspace="w", name="cfg").response_type + == list[ModelDeploymentConfig] + ) + assert ( + endpoints.get_deployment_config_version(workspace="w", config="cfg", name="1").response_type + is ModelDeploymentConfig + ) + assert endpoints.delete_deployment_config(workspace="w", name="cfg").response_type is None + assert ( + endpoints.delete_deployment_config_version(workspace="w", config="cfg", name="1").path_template + == _PREFIX + "/deployment-configs/{config}/versions/{name}" + ) diff --git a/packages/nmp_customization_common/src/nmp/customization_common/service/platform_client.py b/packages/nmp_customization_common/src/nmp/customization_common/service/platform_client.py index 1398dbd2e9..ad65242be7 100644 --- a/packages/nmp_customization_common/src/nmp/customization_common/service/platform_client.py +++ b/packages/nmp_customization_common/src/nmp/customization_common/service/platform_client.py @@ -9,12 +9,11 @@ """ from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import NotFoundError, PermissionDeniedError -from nemo_platform.types.models import ModelEntity from nemo_platform_plugin.client.adapter import client_from_platform -from nemo_platform_plugin.client.errors import NotFoundError as ClientNotFoundError -from nemo_platform_plugin.client.errors import PermissionDeniedError as ClientPermissionDeniedError +from nemo_platform_plugin.client.errors import NotFoundError, PermissionDeniedError from nemo_platform_plugin.files.client import AsyncFilesClient +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ModelEntity from nmp.common.entities.utils import parse_entity_ref from nmp.customization_common.schemas.file_io import FileSetRef @@ -31,9 +30,9 @@ async def check_dataset_access(sdk: AsyncNeMoPlatform, dataset_uri: str, default files = client_from_platform(sdk, AsyncFilesClient) try: await files.get_fileset(workspace=workspace, name=ref.name) - except ClientPermissionDeniedError: + except PermissionDeniedError: raise PermissionError(f"Access denied to dataset fileset '{workspace}/{ref.name}'") from None - except ClientNotFoundError: + except NotFoundError: raise ValueError( f"Dataset fileset '{ref.name}' not found in workspace '{workspace}'. Verify the dataset exists." ) from None @@ -46,8 +45,14 @@ async def fetch_model_entity( ) -> ModelEntity: """Retrieve a model entity by reference string.""" resolved_ref = parse_entity_ref(model_ref, default_workspace) + models = client_from_platform(sdk, AsyncModelsClient) try: - return await sdk.models.retrieve(name=resolved_ref.name, workspace=resolved_ref.workspace, verbose=True) + response = await models.get_model( + name=resolved_ref.name, + workspace=resolved_ref.workspace, + query_params={"verbose": True}, + ) + return response.data() except PermissionDeniedError: raise PermissionError(f"Access denied to model '{resolved_ref.workspace}/{resolved_ref.name}'") from None except NotFoundError: diff --git a/packages/nmp_customization_common/src/nmp/customization_common/tasks/model_entity/run.py b/packages/nmp_customization_common/src/nmp/customization_common/tasks/model_entity/run.py index 1da52d3c1a..727d1582f5 100644 --- a/packages/nmp_customization_common/src/nmp/customization_common/tasks/model_entity/run.py +++ b/packages/nmp_customization_common/src/nmp/customization_common/tasks/model_entity/run.py @@ -14,27 +14,38 @@ import time from pathlib import Path -import httpx -from nemo_platform import ( - APIConnectionError, - APITimeoutError, +from nemo_platform import NeMoPlatform +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import ( ConflictError, InternalServerError, - NeMoPlatform, + NemoTransportError, NotFoundError, ) -from nemo_platform.types.inference import ( - ContainerExecutorConfigParam, +from nemo_platform_plugin.files.client import FilesClient +from nemo_platform_plugin.models.client import ModelsClient +from nemo_platform_plugin.models.types import ( + ContainerExecutorConfig, + CreateModelAdapterRequest, + CreateModelDeploymentConfigRequest, + CreateModelDeploymentRequest, + CreateModelEntityRequest, + Engine, + ListDeploymentConfigsQueryParams, + ListDeploymentsQueryParams, + Lora, ModelDeploymentConfig, - ModelDeploymentConfigFilterParam, - ModelDeploymentConfigModelSpecParam, - ModelDeploymentFilterParam, + ModelDeploymentConfigModelSpec, + ModelDeploymentStatus, + ModelEntity, + ToolCallConfig, + UpdateAdapterRequest, + UpdateModelDeploymentConfigRequest, + UpdateModelEntityRequest, +) +from nemo_platform_plugin.models.types import ( + FinetuningType as ModelsFinetuningType, ) -from nemo_platform.types.models import LoraParam, ModelEntity -from nemo_platform.types.shared_params.tool_call_config import ToolCallConfig as ToolCallConfigParam -from nemo_platform_plugin.client.adapter import client_from_platform -from nemo_platform_plugin.client.errors import InternalServerError as ClientInternalServerError -from nemo_platform_plugin.files.client import FilesClient from nmp.common.sdk_factory import get_task_sdk from nmp.customization_common.schemas.model_entity import ( DeploymentParameters, @@ -51,19 +62,14 @@ INITIAL_BACKOFF_SECONDS = 1.0 MAX_BACKOFF_SECONDS = 30.0 -ACTIVE_DEPLOYMENT_STATUSES = frozenset({"CREATED", "PENDING", "READY"}) +ACTIVE_DEPLOYMENT_STATUSES = frozenset( + {ModelDeploymentStatus.CREATED, ModelDeploymentStatus.PENDING, ModelDeploymentStatus.READY} +) SPEC_POLL_INTERVAL_SECONDS = 10 SPEC_POLL_TIMEOUT_SECONDS = 600 -TRANSIENT_RETRYABLE_EXCEPTIONS = ( - InternalServerError, - APITimeoutError, - APIConnectionError, - ClientInternalServerError, - httpx.TimeoutException, - httpx.ConnectError, -) +TRANSIENT_RETRYABLE_EXCEPTIONS = (InternalServerError, NemoTransportError) def get_config(config_path: Path) -> ModelEntityTaskConfig: @@ -84,6 +90,7 @@ class ModelEntityRunner: def __init__(self, sdk: NeMoPlatform, job_ctx: NMPJobContext): self.sdk = sdk + self.models = client_from_platform(sdk, ModelsClient) self.job_ctx = job_ctx def _wait_for_spec(self, workspace: str, name: str) -> ModelEntity: @@ -93,7 +100,7 @@ def _wait_for_spec(self, workspace: str, name: str) -> ModelEntity: while time.monotonic() - start < SPEC_POLL_TIMEOUT_SECONDS: try: - target = self.sdk.models.retrieve(name=name, workspace=workspace) + target = self.models.get_model(name=name, workspace=workspace).data() spec = target.spec if spec is not None: family = getattr(spec, "family", None) @@ -132,7 +139,7 @@ def get_model_entity(self, model_entity: str, fileset_workspace: str) -> ModelEn ) try: - me: ModelEntity = self.sdk.models.retrieve(name=me_name, workspace=me_workspace) + me = self.models.get_model(name=me_name, workspace=me_workspace).data() except NotFoundError as e: raise ModelEntityCreationError(f"Model entity {me_workspace}/{me_name} not found") from e @@ -182,19 +189,21 @@ def _create_or_update_adapter( """Create or update a LoRA adapter on ``base_me``. Returns (result, base_me).""" assert config.peft is not None try: - output_me = self.sdk.models.adapters.create( + output_me = self.models.create_model_adapter( model_name=base_me.name, workspace=base_me.workspace, - name=config.name, - description=config.description, - fileset=fileset_ref, - finetuning_type=config.peft.type.value, - lora_config=LoraParam( - alpha=config.peft.alpha, - rank=config.peft.rank, + body=CreateModelAdapterRequest( + name=config.name, + description=config.description, + fileset=fileset_ref, + finetuning_type=ModelsFinetuningType(config.peft.type.value), + lora_config=Lora( + alpha=config.peft.alpha, + rank=config.peft.rank, + ), + enabled=True, ), - enabled=True, - ) + ).data() return output_me.model_dump(), base_me except ConflictError: logger.warning( @@ -202,14 +211,16 @@ def _create_or_update_adapter( f"{base_me.workspace}/{base_me.name}, updating with new fileset" ) try: - output_me = self.sdk.models.adapters.update( + output_me = self.models.update_model_adapter( adapter=config.name, model_name=base_me.name, workspace=base_me.workspace, - fileset=fileset_ref, - description=config.description, - enabled=True, - ) + body=UpdateAdapterRequest( + fileset=fileset_ref, + description=config.description, + enabled=True, + ), + ).data() logger.info( f"Successfully updated adapter: {base_me.workspace}/{config.name} " f"for base model {base_me.workspace}/{base_me.name}" @@ -237,29 +248,34 @@ def _create_or_update_full_entity( """Create or update a full / merged model entity. Returns (result, output_me).""" ft_type = config.peft.type.value if config.peft else FinetuningType.ALL_WEIGHTS.value - request_body: dict = { - "name": config.name, - "description": config.description, - "fileset": fileset_ref, - "finetuning_type": ft_type, - "trust_remote_code": config.trust_remote_code, - } - if config.base_model: - request_body["base_model"] = config.base_model + create_request = CreateModelEntityRequest( + name=config.name, + description=config.description, + fileset=fileset_ref, + finetuning_type=ModelsFinetuningType(ft_type), + trust_remote_code=config.trust_remote_code, + base_model=config.base_model, + ) try: - output_me = self.sdk.models.create(workspace=workspace, **request_body) + output_me = self.models.create_model(workspace=workspace, body=create_request).data() logger.info(f"Successfully created model entity: {output_me.workspace}/{output_me.name}") return output_me.model_dump(), output_me except ConflictError: logger.warning(f"Model entity already exists: {workspace}/{config.name}, updating existing model") try: - update_body = {k: v for k, v in request_body.items() if k != "name"} - output_me = self.sdk.models.update( + update_request = UpdateModelEntityRequest( + description=config.description, + fileset=fileset_ref, + finetuning_type=ModelsFinetuningType(ft_type), + trust_remote_code=config.trust_remote_code, + base_model=config.base_model, + ) + output_me = self.models.update_model( name=config.name, workspace=workspace, - **update_body, - ) + body=update_request, + ).data() logger.info(f"Successfully updated model entity: {output_me.workspace}/{output_me.name}") return output_me.model_dump(), output_me except TRANSIENT_RETRYABLE_EXCEPTIONS: @@ -298,15 +314,22 @@ def launch_model(self, config: ModelEntityTaskConfig, me: ModelEntity) -> None: def _has_active_deployment(self, me: ModelEntity) -> bool: """Check if the model entity already has an active deployment.""" - deployment_configs = self.sdk.inference.deployment_configs.list( + config_query = ListDeploymentConfigsQueryParams( + filter=json.dumps({"model_entity_id": f"{me.workspace}/{me.name}"}) + ) + deployment_configs = self.models.list_deployment_configs( workspace=me.workspace, - filter=ModelDeploymentConfigFilterParam(model_entity_id=f"{me.workspace}/{me.name}"), - ).data + query_params=config_query, + ).items() for c in deployment_configs: - deployments = self.sdk.inference.deployments.list( - filter=ModelDeploymentFilterParam(config=c.name, workspace=me.workspace) - ).data + deployment_query = ListDeploymentsQueryParams( + filter=json.dumps({"config": c.name, "workspace": me.workspace}) + ) + deployments = self.models.list_deployments( + workspace=me.workspace, + query_params=deployment_query, + ).items() for d in deployments: if d.status in ACTIVE_DEPLOYMENT_STATUSES: logger.info(f"Active deployment (status={d.status}) exists for config {c.name}, skipping") @@ -327,7 +350,7 @@ def _resolve_config_ref(self, config_ref: str, me_workspace: str) -> ModelDeploy ) try: - return self.sdk.inference.deployment_configs.retrieve(workspace=workspace, name=name) + return self.models.get_deployment_config(workspace=workspace, name=name).data() except Exception as e: raise ModelEntityCreationError( f"Failed to resolve deployment config '{config_ref}' in workspace '{workspace}': {e}" @@ -335,12 +358,12 @@ def _resolve_config_ref(self, config_ref: str, me_workspace: str) -> ModelDeploy def _create_deployment_config(self, deploy_params: DeploymentParameters, me: ModelEntity) -> ModelDeploymentConfig: """Create (or update) a ``ModelDeploymentConfig`` from inline parameters.""" - model_spec = ModelDeploymentConfigModelSpecParam( + model_spec = ModelDeploymentConfigModelSpec( model_name=me.name, model_namespace=me.workspace, lora_enabled=deploy_params.lora_enabled, ) - executor_config = ContainerExecutorConfigParam( + executor_config = ContainerExecutorConfig( image_name=deploy_params.image_name, image_tag=deploy_params.image_tag, gpu=deploy_params.gpu, @@ -348,28 +371,32 @@ def _create_deployment_config(self, deploy_params: DeploymentParameters, me: Mod ) if deploy_params.tool_call_config: - model_spec["tool_call_config"] = ToolCallConfigParam( - **deploy_params.tool_call_config.model_dump(exclude_none=True) + model_spec.tool_call_config = ToolCallConfig.model_validate( + deploy_params.tool_call_config.model_dump(exclude_none=True) ) deployment_cfg_name = sanitize_name("sft-cfg", me.name) try: - return self.sdk.inference.deployment_configs.create( + return self.models.create_deployment_config( workspace=me.workspace, - name=deployment_cfg_name, - engine="nim", - model_spec=model_spec, - executor_config=executor_config, - ) + body=CreateModelDeploymentConfigRequest( + name=deployment_cfg_name, + engine=Engine.NIM, + model_spec=model_spec, + executor_config=executor_config, + ), + ).data() except ConflictError: logger.info(f"Deployment config {me.workspace}/{deployment_cfg_name} already exists, updating") - return self.sdk.inference.deployment_configs.update( + return self.models.update_deployment_config( workspace=me.workspace, name=deployment_cfg_name, - engine="nim", - model_spec=model_spec, - executor_config=executor_config, - ) + body=UpdateModelDeploymentConfigRequest( + engine=Engine.NIM, + model_spec=model_spec, + executor_config=executor_config, + ), + ).data() def _create_deployment(self, deployment_config: ModelDeploymentConfig, me: ModelEntity) -> None: """Create a deployment from the given ``ModelDeploymentConfig``.""" @@ -380,23 +407,25 @@ def _create_deployment(self, deployment_config: ModelDeploymentConfig, me: Model deployment_name = sanitize_name("sft-deploy", me.name) try: - deployment = self.sdk.inference.deployments.create( + deployment = self.models.create_deployment( workspace=deployment_config.workspace, - name=deployment_name, - config=deployment_config.name, - ) + body=CreateModelDeploymentRequest( + name=deployment_name, + config=deployment_config.name, + ), + ).data() logger.info(f"Deployment created: {deployment.workspace}/{deployment.name}") except ConflictError: logger.info(f"Deployment {deployment_config.workspace}/{deployment_name} already exists") - deployment = self.sdk.inference.deployments.retrieve( + deployment = self.models.get_deployment( workspace=deployment_config.workspace, name=deployment_name, - ) + ).data() - deployment_status = self.sdk.inference.deployments.retrieve( + deployment_status = self.models.get_deployment( workspace=deployment.workspace, name=deployment.name, - ) + ).data() logger.info( f"Deployment {deployment_status.workspace}/{deployment_status.name} status: {deployment_status.status}" ) diff --git a/packages/nmp_customization_common/tests/test_models_client_migration.py b/packages/nmp_customization_common/tests/test_models_client_migration.py new file mode 100644 index 0000000000..c38ddccb22 --- /dev/null +++ b/packages/nmp_customization_common/tests/test_models_client_migration.py @@ -0,0 +1,61 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from nemo_platform_plugin.client.errors import NemoTransportError +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient +from nmp.customization_common.service.platform_client import fetch_model_entity +from nmp.customization_common.tasks.model_entity.run import ( + TRANSIENT_RETRYABLE_EXCEPTIONS, + ModelEntityRunner, +) + + +@pytest.mark.asyncio +async def test_fetch_model_entity_uses_async_models_client() -> None: + sdk = MagicMock() + entity = MagicMock() + response = MagicMock() + response.data.return_value = entity + models = MagicMock() + models.get_model = AsyncMock(return_value=response) + + with patch( + "nmp.customization_common.service.platform_client.client_from_platform", + return_value=models, + ) as make_client: + result = await fetch_model_entity("other/model", "default", sdk) + + make_client.assert_called_once_with(sdk, AsyncModelsClient) + models.get_model.assert_awaited_once_with( + name="model", + workspace="other", + query_params={"verbose": True}, + ) + assert result is entity + + +def test_model_entity_runner_uses_models_client() -> None: + sdk = MagicMock() + entity = MagicMock() + response = MagicMock() + response.data.return_value = entity + models = MagicMock() + models.get_model.return_value = response + + with patch( + "nmp.customization_common.tasks.model_entity.run.client_from_platform", + return_value=models, + ) as make_client: + runner = ModelEntityRunner(sdk, MagicMock()) + result = runner.get_model_entity("other/model", "default") + + make_client.assert_called_once_with(sdk, ModelsClient) + models.get_model.assert_called_once_with(name="model", workspace="other") + assert result is entity + + +def test_models_transport_errors_remain_retryable() -> None: + assert NemoTransportError in TRANSIENT_RETRYABLE_EXCEPTIONS diff --git a/plugins/nemo-unsloth/tests/test_jobs.py b/plugins/nemo-unsloth/tests/test_jobs.py index 121963df0f..b251232481 100644 --- a/plugins/nemo-unsloth/tests/test_jobs.py +++ b/plugins/nemo-unsloth/tests/test_jobs.py @@ -19,11 +19,13 @@ import asyncio from types import SimpleNamespace -from typing import Any +from typing import Any, Callable from unittest.mock import AsyncMock, MagicMock, patch import pytest +from nemo_platform_plugin.files.client import AsyncFilesClient from nemo_platform_plugin.jobs.exceptions import PlatformJobCompilationError +from nemo_platform_plugin.models.client import AsyncModelsClient from nemo_unsloth_plugin.jobs.jobs import UnslothJob from nemo_unsloth_plugin.schema import UnslothJobInput from nmp.unsloth.schemas import UnslothJobOutput @@ -40,34 +42,42 @@ def _input_dict(**overrides: Any) -> dict[str, Any]: def _stub_async_sdk() -> SimpleNamespace: - """Async SDK used by ``to_spec`` (validates refs).""" - me = SimpleNamespace( + """Async SDK transport placeholder passed through the contributor API.""" + return SimpleNamespace() + + +def _mock_client_factory() -> Callable[[object, type[object]], MagicMock]: + """Build typed Models and Files clients for platform reference validation.""" + model_entity = SimpleNamespace( name="base", workspace="default", spec=None, fileset="base-fs", trust_remote_code=False, ) - return SimpleNamespace( - models=SimpleNamespace(retrieve=AsyncMock(return_value=me)), - files=SimpleNamespace( - filesets=SimpleNamespace(retrieve=AsyncMock(return_value=SimpleNamespace())), - ), - ) + model_response = MagicMock() + model_response.data.return_value = model_entity + models = MagicMock() + models.get_model = AsyncMock(return_value=model_response) + + files = MagicMock() + files.get_fileset = AsyncMock(return_value=MagicMock()) + def factory(_sdk: object, client_type: type[object]) -> MagicMock: + if client_type is AsyncModelsClient: + return models + if client_type is AsyncFilesClient: + return files + raise AssertionError(f"Unexpected client type: {client_type}") -def _mock_files_client() -> AsyncMock: - """Build a mock AsyncFilesClient for check_dataset_access.""" - mock = AsyncMock() - mock.get_fileset.return_value = MagicMock() - return mock + return factory def _make_canonical(workspace: str = "default", **overrides: Any) -> UnslothJobOutput: spec = UnslothJobInput.model_validate(_input_dict(**overrides)) with patch( "nmp.customization_common.service.platform_client.client_from_platform", - return_value=_mock_files_client(), + side_effect=_mock_client_factory(), ): return asyncio.run( UnslothJob.to_spec( diff --git a/plugins/nemo-unsloth/tests/test_schema.py b/plugins/nemo-unsloth/tests/test_schema.py index bbc4e662db..67b1362b6d 100644 --- a/plugins/nemo-unsloth/tests/test_schema.py +++ b/plugins/nemo-unsloth/tests/test_schema.py @@ -9,9 +9,12 @@ import json from pathlib import Path from types import SimpleNamespace +from typing import Callable from unittest.mock import AsyncMock, MagicMock, patch import pytest +from nemo_platform_plugin.files.client import AsyncFilesClient +from nemo_platform_plugin.models.client import AsyncModelsClient from nemo_unsloth_plugin.schema import ( DatasetSpec, LoRAParams, @@ -26,9 +29,14 @@ from pydantic import ValidationError -def _stub_sdk(*, is_embedding: bool = False) -> SimpleNamespace: - """Build a minimal async SDK that resolves model + dataset refs.""" - spec = SimpleNamespace(is_embedding_model=is_embedding) if is_embedding else None +def _stub_sdk() -> SimpleNamespace: + """Build the async SDK transport placeholder accepted by the contributor API.""" + return SimpleNamespace() + + +def _mock_client_factory(*, is_embedding: bool = False) -> Callable[[object, type[object]], MagicMock]: + """Build typed Models and Files clients for platform reference validation.""" + spec = SimpleNamespace(is_embedding_model=True) if is_embedding else None model_entity = SimpleNamespace( name="m", workspace="default", @@ -36,25 +44,28 @@ def _stub_sdk(*, is_embedding: bool = False) -> SimpleNamespace: fileset="m", trust_remote_code=False, ) - return SimpleNamespace( - models=SimpleNamespace(retrieve=AsyncMock(return_value=model_entity)), - files=SimpleNamespace( - filesets=SimpleNamespace(retrieve=AsyncMock(return_value=SimpleNamespace())), - ), - ) + model_response = MagicMock() + model_response.data.return_value = model_entity + models = MagicMock() + models.get_model = AsyncMock(return_value=model_response) + + files = MagicMock() + files.get_fileset = AsyncMock(return_value=MagicMock()) + def factory(_sdk: object, client_type: type[object]) -> MagicMock: + if client_type is AsyncModelsClient: + return models + if client_type is AsyncFilesClient: + return files + raise AssertionError(f"Unexpected client type: {client_type}") -def _mock_files_client() -> AsyncMock: - """Build a mock AsyncFilesClient for check_dataset_access.""" - mock = AsyncMock() - mock.get_fileset.return_value = MagicMock() - return mock + return factory def _run_transform(spec: UnslothJobInput) -> UnslothJobOutput: with patch( "nmp.customization_common.service.platform_client.client_from_platform", - return_value=_mock_files_client(), + side_effect=_mock_client_factory(), ): return asyncio.run(transform_input_to_output(spec, "default", _stub_sdk())) @@ -234,12 +245,12 @@ def test_merged_inferred_as_model_type(self) -> None: assert out.output.save_method == "merged_4bit" def test_embedding_model_rejected(self) -> None: - sdk = _stub_sdk(is_embedding=True) + sdk = _stub_sdk() spec = UnslothJobInput.model_validate(_minimal_payload()) with ( patch( "nmp.customization_common.service.platform_client.client_from_platform", - return_value=_mock_files_client(), + side_effect=_mock_client_factory(is_embedding=True), ), pytest.raises(ValueError, match="Embedding-model SFT"), ): diff --git a/sdk/python/nemo-platform/src/nemo_platform/models/resources.py b/sdk/python/nemo-platform/src/nemo_platform/models/resources.py index 773857a961..57c9ae33ee 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/models/resources.py +++ b/sdk/python/nemo-platform/src/nemo_platform/models/resources.py @@ -1,41 +1,63 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Extended ModelsResource with high-level helper methods.""" +"""High-level Models helpers built on the typed :class:`ModelsClient`. + +Historically ``ModelsResource`` / ``AsyncModelsResource`` extended the +Stainless-generated ``ModelsResource`` base to inherit its CRUD surface and +layered convenience helpers on top. As part of AIRCORE-876 the Stainless- +resource *inheritance* is removed: these classes now hold a ``NeMoPlatform`` +SDK and drive the typed ``nemo_platform_plugin.nemo_platform.models.client.ModelsClient`` +(built from that SDK via ``client_from_platform``) for their own genuine public +surface -- the inference-gateway route builders, OpenAI client factories, and +deployment/provider status polling. + +CRUD (``retrieve`` / ``create`` / ``list`` / adapter sub-resource) is +intentionally *not* re-implemented here: reproducing the Stainless resource +method/param shapes would be a compatibility proxy. Callers that need CRUD +should use the typed client directly, e.g.:: + + from nemo_platform_plugin.client.adapter import client_from_platform + from nemo_platform_plugin.nemo_platform.models.client import ModelsClient + + models = client_from_platform(sdk, ModelsClient) + entity = nemo_platform.models.get_model(name="llama", workspace="default").data() + +.. warning:: + **Vendoring gate.** This package is vendored into the SDK as + ``nemo_platform.models`` (``sdk.models`` resolves to this ``ModelsResource`` + via ``packages/models`` ``[tool.vendor-package]``). Because the Stainless + CRUD inheritance is dropped above, running ``make vendor`` will remove + ``sdk.nemo_platform.models.retrieve`` / ``create`` / ``list`` / ``adapters.*`` from the + vendored SDK and break every consumer still on that surface (automodel + compiler, provider/deployment reconcilers, models_controller, adapter + sidecar, model_spec task, ``nmp_customization_common``, evaluator resolver, + inference-gateway model cache, generated CLI). Those call sites must first be + migrated to the typed client and repointed from ``nemo_platform`` exceptions + to ``nemo_platform_plugin.client.errors`` (``NotFoundError`` / ``ConflictError`` + are distinct classes). Do not run ``make vendor`` for this package until that + consumer migration lands -- it is a separate follow-up under the AIRCORE-827 + migration umbrella. + +The inference-gateway *readiness* probe (:meth:`wait_for_gateway`) still calls +through the ``NeMoPlatform`` SDK because it targets the separate +inference-gateway service, which has not yet been migrated to a typed client. +""" + +from __future__ import annotations -import asyncio import time from datetime import datetime -from nemo_platform import NotFoundError -from nemo_platform.resources.models import AsyncModelsResource as BaseAsyncModelsResource -from nemo_platform.resources.models import ModelsResource as BaseModelsResource +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform, NotFoundError from nemo_platform.types.inference import ModelDeployment, ModelProvider from nemo_platform.types.models import ModelEntity +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.models.client import AsyncModelsClient, ModelsClient -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: - return None - if isinstance(entry_timestamp, str): - try: - entry_timestamp = datetime.fromisoformat(entry_timestamp.replace("Z", "+00:00")) - except (ValueError, TypeError): - return None - if not hasattr(entry_timestamp, "timestamp") or not hasattr(created_at, "timestamp"): - return None - try: - return int(entry_timestamp.timestamp() - created_at.timestamp()) - except (TypeError, OSError): - return None - - -class ModelsResource(BaseModelsResource): - """Extended ModelsResource with high-level helper methods. - - All existing methods (create, retrieve, list, etc.) work unchanged. - Adds convenience methods for OpenAI integration and deployment management. +class ModelsResource: + """Sync Models helpers backed by a typed :class:`ModelsClient`. Example: >>> sdk = NeMoPlatform(base_url="http://nmp-host", workspace="default") @@ -43,152 +65,56 @@ class ModelsResource(BaseModelsResource): >>> sdk.nemo_platform.models.wait_for_status("my-deployment", "READY") """ + def __init__(self, client: NeMoPlatform) -> None: + self._client = client + self._typed: ModelsClient | None = None + + @property + def models(self) -> ModelsClient: + """The typed Models client sharing this SDK's transport (built lazily).""" + if self._typed is None: + self._typed = client_from_platform(self._client, ModelsClient) + return self._typed + def _get_base_url_str(self) -> str: """Get the base URL as a string with trailing slash removed.""" return str(self._client.base_url).rstrip("/") - def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: - """ - Generate the base URL for the OpenAI proxy route. - - This route uses the `model` field in the request body for routing, - formatted as `workspace/model_entity_name`. + # -- OpenAI inference-gateway route builders (delegate to the typed client) -- - Args: - workspace: The workspace identifier - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> base_url = sdk.nemo_platform.models.get_openai_route_base_url() - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/openai/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - base_url = self._get_base_url_str() - return f"{base_url}/apis/inference-gateway/v2/workspaces/{workspace}/openai/-/v1" + def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: + """Base URL for the OpenAI proxy route (routes on the request body ``model``).""" + return self.models.get_openai_route_base_url(workspace=workspace) def get_client_default_headers(self) -> dict[str, str]: - """Get string-only default headers for third-party client libraries. + """String-only default headers for third-party client libraries (OpenAI SDK, LiteLLM). - Use this helper when constructing external clients (for example OpenAI - SDK or LiteLLM) so auth and identity headers from the SDK are forwarded. - This is required for successful inference requests when platform auth/ - authorization is enabled. + Forwards the SDK's auth/identity headers, required for inference when + platform authorization is enabled. """ return {key: value for key, value in self._client.default_headers.items() if isinstance(value, str)} def get_openai_client(self, *, workspace: str | None = None): - """ - Get a sync OpenAI client configured for NeMo Platform's inference gateway. - - This method returns an OpenAI client with the base_url set to the - OpenAI proxy route for the specified workspace. The client can be - used directly with the standard OpenAI SDK interface. - - Args: - workspace: The workspace identifier - - Returns: - An OpenAI client configured for the inference gateway - - Example: - >>> client = sdk.nemo_platform.models.get_openai_client() - >>> response = client.chat.completions.create( - ... model="default/meta_llama-3.2-1b-instruct", - ... messages=[{"role": "user", "content": "Hello!"}] - ... ) - """ + """A sync OpenAI client configured for NeMo Platform's inference gateway.""" import openai base_url = self.get_openai_route_base_url(workspace=workspace) - # Preserve auth and identity headers from the parent SDK client. default_headers = self.get_client_default_headers() return openai.OpenAI(base_url=base_url, api_key="not-needed", default_headers=default_headers) def get_provider_route_openai_url(self, provider: ModelProvider) -> str: - """ - Generate an OpenAI SDK-compatible URL for the provider proxy route. - - Handles the conditional /v1 suffix based on the provider's host_url: - - If host_url ends with /v1, no suffix is added - - Otherwise, /v1 is appended - - Args: - provider: The ModelProvider object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> provider = sdk.inference.providers.retrieve("my-provider", workspace="default") - >>> base_url = sdk.nemo_platform.models.get_provider_route_openai_url(provider) - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/provider/my-provider/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - base_url = self._get_base_url_str() - route = f"{base_url}/apis/inference-gateway/v2/workspaces/{provider.workspace}/provider/{provider.name}/-" - - host_url_normalized = provider.host_url.rstrip("/") - if not host_url_normalized.endswith("/v1"): - route = f"{route}/v1" - - return route + """OpenAI SDK-compatible URL for a provider proxy route (conditional ``/v1``).""" + return self.models.get_provider_route_openai_url(provider) def get_provider_route_openai_url_for_deployment(self, deployment: ModelDeployment) -> str: - """ - Generate an OpenAI SDK-compatible URL for a deployment's model provider. - - This is a convenience method that fetches the ModelProvider associated - with the deployment and returns the provider route URL. - - Args: - deployment: The ModelDeployment object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Raises: - ValueError: If the deployment has no associated model_provider_id - - Example: - >>> deployment = sdk.inference.deployments.retrieve("my-deployment", workspace="default") - >>> base_url = sdk.nemo_platform.models.get_provider_route_openai_url_for_deployment(deployment) - >>> openai_client = OpenAI(base_url=base_url) - """ - if not deployment.model_provider_id: - raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") - - # model_provider_id is in "workspace/name" format - workspace, name = deployment.model_provider_id.split("/", 1) - provider = self._client.inference.providers.retrieve(name, workspace=workspace) - return self.get_provider_route_openai_url(provider) + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + return self.models.get_provider_route_openai_url_for_deployment(deployment) def get_model_entity_route_openai_url(self, model_entity: ModelEntity) -> str: - """ - Generate an OpenAI SDK-compatible URL for the model entity proxy route. - - Always appends /v1 suffix since the client doesn't interact directly - with the provider's host_url. + """OpenAI SDK-compatible URL for a model-entity proxy route (always ``/v1``).""" + return self.models.get_model_entity_route_openai_url(model_entity) - Args: - model_entity: The ModelEntity object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Example: - >>> entity = sdk.nemo_platform.models.retrieve("my-model", workspace="default") - >>> base_url = sdk.nemo_platform.models.get_model_entity_route_openai_url(entity) - >>> # Returns: {base_url}/apis/inference-gateway/v2/workspaces/default/model/my-model/-/v1 - >>> openai_client = OpenAI(base_url=base_url) - """ - base_url = self._get_base_url_str() - return ( - f"{base_url}/apis/inference-gateway/v2/workspaces/{model_entity.workspace}/model/{model_entity.name}/-/v1" - ) + # -- Deployment / provider status polling -- def wait_for_status( self, @@ -199,85 +125,31 @@ def wait_for_status( timeout: int = 1200, check_gateway: bool = True, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. - - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", verify the - gateway can route to the provider before returning (default: True). + """Wait for a ModelDeployment to reach ``desired_status``. - Returns: - True if desired status reached, False if timeout + When ``desired_status`` is ``"READY"`` and ``check_gateway`` is set, also + waits for the inference gateway to be able to route to the provider. """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - deployment_status = self.wait_for_deployment_status( + if not self.models.wait_for_deployment_status( deployment_name, desired_status, workspace=workspace, timeout=timeout - ) - if not deployment_status: + ): return False - - # Verify gateway can route to the provider if desired_status == "READY" and check_gateway: - gateway_ready = self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) - if not gateway_ready: - return False - + return self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) return True - def wait_for_gateway( + def wait_for_deployment_status( self, - provider_name: str, + deployment_name: str, + desired_status: str, *, workspace: str | None = None, - timeout: int = 60, + timeout: int = 1200, ) -> bool: - """ - Wait for the inference gateway to be able to route to a provider. - - Polls the gateway's /ready endpoint until it returns success, indicating - the gateway has refreshed its cache and is aware of the provider. - - Args: - provider_name: Name of the model provider - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - - Returns: - True if gateway is ready, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - - print("Waiting for gateway to be ready...") - - while time.time() - start_time < timeout: - try: - self._client.inference.gateway.provider.ready( - provider_name, - workspace=workspace, - ) - timestamp = datetime.now().strftime("%H:%M:%S") - print(f" [{timestamp}] Gateway is ready!\n") - return True - except NotFoundError: - # Gateway doesn't know about the provider yet, keep waiting - time.sleep(1) - except Exception: - # Connection error or other issue, keep waiting - time.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"Gateway timeout after {elapsed}s\n") - return False + """Wait for a ModelDeployment to reach ``desired_status`` (or 404 for ``DELETED``).""" + return self.models.wait_for_deployment_status( + deployment_name, desired_status, workspace=workspace, timeout=timeout + ) def wait_for_provider( self, @@ -288,300 +160,93 @@ def wait_for_provider( timeout: int = 60, check_gateway: bool = True, ) -> bool: - """ - Wait for a provider to reach the desired status. - - This is useful for external providers (like NVIDIA Build or OpenAI) where - you need to wait for the provider to be ready before making inference calls. - - Args: - provider_name: Name of the provider - desired_status: Target status (default: "READY") - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", also verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - last_status = "" - - print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") - - while time.time() - start_time < timeout: - try: - provider = self._client.inference.providers.retrieve( - provider_name, - workspace=workspace, - ) - current_status = provider.status - - if current_status != last_status: - timestamp = datetime.now().strftime("%H:%M:%S") - elapsed = int(time.time() - start_time) - print(f" [{timestamp}] ({elapsed}s) Status: {current_status}") - last_status = current_status - - if current_status == desired_status: - if desired_status == "READY" and check_gateway: - return self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) - print() - return True - - if current_status == "ERROR": - print(f"\nProvider entered ERROR state: {provider.status_message}\n") - return False - - except NotFoundError: - print(f"\nProvider '{provider_name}' not found\n") - return False - - time.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"\nProvider timeout after {elapsed}s. Last status: {last_status}\n") - return False + """Wait for a ModelProvider to reach ``desired_status`` (optionally gateway-ready).""" + if not self.models.wait_for_provider_status( + provider_name, desired_status, workspace=workspace, timeout=timeout + ): + return False + if desired_status == "READY" and check_gateway: + return self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) + return True - def wait_for_deployment_status( + def wait_for_gateway( self, - deployment_name: str, - desired_status: str, + provider_name: str, *, workspace: str | None = None, - timeout: int = 1200, + timeout: int = 60, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. + """Wait for the inference gateway to be able to route to a provider. - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - - Returns: - True if desired status reached, False if timeout + Targets the separate inference-gateway service, which is not yet migrated + to a typed client, so this still calls through the ``NeMoPlatform`` SDK. """ if workspace is None: workspace = self._client._get_workspace_path_param() start_time = time.time() - last_status = "" - last_message = "" - last_history_len = 0 - - print(f"Waiting for status: {desired_status}...\n") - + print("Waiting for gateway to be ready...") while time.time() - start_time < timeout: try: - deployment = self._client.inference.deployments.retrieve(deployment_name, workspace=workspace) - history = getattr(deployment, "status_history", None) - created_at = getattr(deployment, "created_at", None) - # API guarantees last history entry is current state; fall back to top-level fields if no history - if history and len(history) > 0: - last_entry = history[-1] - current_status = getattr(last_entry, "status", deployment.status) - status_message = getattr(last_entry, "status_message", "") or "" - else: - current_status = deployment.status - status_message = deployment.status_message or "" - last_status = current_status - last_message = status_message - - # Only print status from history; elapsed shown is seconds since deployment creation - if history and len(history) > last_history_len: - for i in range(last_history_len, len(history)): - entry = history[i] - ts = getattr(entry, "timestamp", None) - ts_str = ts.strftime("%H:%M:%S") if hasattr(ts, "strftime") else str(ts) if ts else "" - st = getattr(entry, "status", "") - msg = getattr(entry, "status_message", "") or "" - secs = _seconds_since_creation(ts, created_at) - part = f" [{ts_str}] " - if secs is not None: - part += f"(+{secs}s) " - part += f"Status: {st}" - if msg: - part += f" - {msg}" - print(part) - last_history_len = len(history) - - # Check if we've reached the desired status - # For DELETED status, we need to wait for the actual 404 (garbage collection) - if current_status == desired_status and desired_status != "DELETED": - print(f"Deployment reached {desired_status} status!\n") - return True - - # Handle error states - if current_status == "ERROR": - print(f"Deployment entered ERROR state: {status_message}\n") - return False - + self._client.inference.gateway.provider.ready(provider_name, workspace=workspace) + print(f" [{datetime.now().strftime('%H:%M:%S')}] Gateway is ready!\n") + return True except NotFoundError: - # For DELETED status, not found means success - if desired_status == "DELETED": - print(f"Deployment {desired_status}!\n") - return True - # For other statuses, not found is an error - print("Deployment not found\n") - return False - - time.sleep(3) - - # Timeout reached (wait_elapsed is time since we started polling) - wait_elapsed = int(time.time() - start_time) - detail = f"Last status: {last_status}" - if last_message: - detail += f" - {last_message}" - print(f"Timeout after {wait_elapsed}s. {detail}\n") + time.sleep(1) + except Exception: + time.sleep(1) + print(f"Gateway timeout after {int(time.time() - start_time)}s\n") return False -class AsyncModelsResource(BaseAsyncModelsResource): - """Extended AsyncModelsResource with high-level helper methods. +class AsyncModelsResource: + """Async twin of :class:`ModelsResource`. - All existing async methods (create, retrieve, list, etc.) work unchanged. - Adds convenience methods for OpenAI integration and deployment management. + Route builders are synchronous (no I/O) and safe to call from async code; + methods that perform I/O are async. + """ - URL builder methods are synchronous (no I/O) and safe to call from async code. - Methods that perform I/O are properly async. + def __init__(self, client: AsyncNeMoPlatform) -> None: + self._client = client + self._typed: AsyncModelsClient | None = None - Example: - >>> sdk = AsyncNeMoPlatform(base_url="http://nmp-host", workspace="default") - >>> sdk.nemo_platform.models.get_openai_route_base_url() - >>> await sdk.nemo_platform.models.wait_for_status("my-deployment", "READY") - """ + @property + def models(self) -> AsyncModelsClient: + """The typed async Models client sharing this SDK's transport (built lazily).""" + if self._typed is None: + self._typed = client_from_platform(self._client, AsyncModelsClient) + return self._typed def _get_base_url_str(self) -> str: """Get the base URL as a string with trailing slash removed.""" return str(self._client.base_url).rstrip("/") def get_openai_route_base_url(self, *, workspace: str | None = None) -> str: - """ - Generate the base URL for the OpenAI proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - workspace: The workspace identifier - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - base_url = self._get_base_url_str() - return f"{base_url}/apis/inference-gateway/v2/workspaces/{workspace}/openai/-/v1" + """Base URL for the OpenAI proxy route. Synchronous (no I/O).""" + return self.models.get_openai_route_base_url(workspace=workspace) def get_client_default_headers(self) -> dict[str, str]: - """Get string-only default headers for third-party client libraries. - - Use this helper when constructing external clients (for example OpenAI - SDK or LiteLLM) so auth and identity headers from the SDK are forwarded. - This is required for successful inference requests when platform auth/ - authorization is enabled. - """ + """String-only default headers for third-party client libraries.""" return {key: value for key, value in self._client.default_headers.items() if isinstance(value, str)} def get_async_openai_client(self, *, workspace: str | None = None): - """ - Get an async OpenAI client configured for NeMo Platform's inference gateway. - - This method returns an AsyncOpenAI client with the base_url set to the - OpenAI proxy route for the specified workspace. - - Args: - workspace: The workspace identifier - - Returns: - An AsyncOpenAI client configured for the inference gateway - - Example: - >>> client = sdk.nemo_platform.models.get_async_openai_client() - >>> response = await client.chat.completions.create( - ... model="default/meta_llama-3.2-1b-instruct", - ... messages=[{"role": "user", "content": "Hello!"}] - ... ) - """ + """An async OpenAI client configured for NeMo Platform's inference gateway.""" import openai base_url = self.get_openai_route_base_url(workspace=workspace) - # Preserve auth and identity headers from the parent SDK client. default_headers = self.get_client_default_headers() return openai.AsyncOpenAI(base_url=base_url, api_key="not-needed", default_headers=default_headers) def get_provider_route_openai_url(self, provider: ModelProvider) -> str: - """ - Generate an OpenAI SDK-compatible URL for the provider proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - provider: The ModelProvider object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - base_url = self._get_base_url_str() - route = f"{base_url}/apis/inference-gateway/v2/workspaces/{provider.workspace}/provider/{provider.name}/-" - - host_url_normalized = provider.host_url.rstrip("/") - if not host_url_normalized.endswith("/v1"): - route = f"{route}/v1" - - return route + """OpenAI SDK-compatible URL for a provider proxy route. Synchronous (no I/O).""" + return self.models.get_provider_route_openai_url(provider) async def get_provider_route_openai_url_for_deployment(self, deployment: ModelDeployment) -> str: - """ - Generate an OpenAI SDK-compatible URL for a deployment's model provider. - - This is an async method that fetches the ModelProvider associated - with the deployment and returns the provider route URL. - - Args: - deployment: The ModelDeployment object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - - Raises: - ValueError: If the deployment has no associated model_provider_id - - Example: - >>> deployment = await sdk.inference.deployments.retrieve("my-deployment", workspace="default") - >>> base_url = await sdk.nemo_platform.models.get_provider_route_openai_url_for_deployment(deployment) - >>> openai_client = AsyncOpenAI(base_url=base_url) - """ - if not deployment.model_provider_id: - raise ValueError(f"Deployment '{deployment.name}' has no associated model_provider_id") - - # model_provider_id is in "workspace/name" format - workspace, name = deployment.model_provider_id.split("/", 1) - provider = await self._client.inference.providers.retrieve(name, workspace=workspace) - return self.get_provider_route_openai_url(provider) + """Fetch a deployment's ModelProvider and return its OpenAI route URL.""" + return await self.models.get_provider_route_openai_url_for_deployment(deployment) def get_model_entity_route_openai_url(self, model_entity: ModelEntity) -> str: - """ - Generate an OpenAI SDK-compatible URL for the model entity proxy route. - - This is a synchronous method (no I/O) and safe to call from async code. - - Args: - model_entity: The ModelEntity object from the SDK - - Returns: - A URL string suitable for use as OpenAI client's base_url - """ - base_url = self._get_base_url_str() - return ( - f"{base_url}/apis/inference-gateway/v2/workspaces/{model_entity.workspace}/model/{model_entity.name}/-/v1" - ) + """OpenAI SDK-compatible URL for a model-entity proxy route. Synchronous (no I/O).""" + return self.models.get_model_entity_route_openai_url(model_entity) async def wait_for_status( self, @@ -592,85 +257,27 @@ async def wait_for_status( timeout: int = 1200, check_gateway: bool = True, ) -> bool: - """ - Wait for a ModelDeployment and ModelProvider to reach the desired status. - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. - - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - deployment_status = await self.wait_for_deployment_status( + """Wait for a ModelDeployment to reach ``desired_status`` (optionally gateway-ready).""" + if not await self.models.wait_for_deployment_status( deployment_name, desired_status, workspace=workspace, timeout=timeout - ) - if not deployment_status: + ): return False - - # Verify gateway can route to the provider if desired_status == "READY" and check_gateway: - gateway_ready = await self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) - if not gateway_ready: - return False - + return await self.wait_for_gateway(deployment_name, workspace=workspace, timeout=timeout) return True - async def wait_for_gateway( + async def wait_for_deployment_status( self, - provider_name: str, + deployment_name: str, + desired_status: str, *, workspace: str | None = None, - timeout: int = 60, + timeout: int = 1200, ) -> bool: - """ - Wait for the inference gateway to be able to route to a provider. - - Polls the gateway's /ready endpoint until it returns success, indicating - the gateway has refreshed its cache and is aware of the provider. - - Args: - provider_name: Name of the model provider - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - - Returns: - True if gateway is ready, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - - print("Waiting for gateway to be ready...") - - while time.time() - start_time < timeout: - try: - await self._client.inference.gateway.provider.ready( - provider_name, - workspace=workspace, - ) - timestamp = datetime.now().strftime("%H:%M:%S") - print(f" [{timestamp}] Gateway is ready!\n") - return True - except NotFoundError: - # Gateway doesn't know about the provider yet, keep waiting - await asyncio.sleep(1) - except Exception: - # Connection error or other issue, keep waiting - await asyncio.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"Gateway timeout after {elapsed}s\n") - return False + """Wait for a ModelDeployment to reach ``desired_status`` (or 404 for ``DELETED``).""" + return await self.models.wait_for_deployment_status( + deployment_name, desired_status, workspace=workspace, timeout=timeout + ) async def wait_for_provider( self, @@ -681,156 +288,41 @@ async def wait_for_provider( timeout: int = 60, check_gateway: bool = True, ) -> bool: - """ - Wait for a provider to reach the desired status (async version). - - This is useful for external providers (like NVIDIA Build or OpenAI) where - you need to wait for the provider to be ready before making inference calls. - - Args: - provider_name: Name of the provider - desired_status: Target status (default: "READY") - workspace: Workspace of the provider - timeout: Maximum time to wait in seconds - check_gateway: When True and desired_status is "READY", also verify the - gateway can route to the provider before returning (default: True). - - Returns: - True if desired status reached, False if timeout - """ - if workspace is None: - workspace = self._client._get_workspace_path_param() - start_time = time.time() - last_status = "" - - print(f"Waiting for provider '{provider_name}' to reach status: {desired_status}...") - - while time.time() - start_time < timeout: - try: - provider = await self._client.inference.providers.retrieve( - provider_name, - workspace=workspace, - ) - current_status = provider.status - - if current_status != last_status: - timestamp = datetime.now().strftime("%H:%M:%S") - elapsed = int(time.time() - start_time) - print(f" [{timestamp}] ({elapsed}s) Status: {current_status}") - last_status = current_status - - if current_status == desired_status: - if desired_status == "READY" and check_gateway: - return await self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) - print() - return True - - if current_status == "ERROR": - print(f"\nProvider entered ERROR state: {provider.status_message}\n") - return False - - except NotFoundError: - print(f"\nProvider '{provider_name}' not found\n") - return False - - await asyncio.sleep(1) - - elapsed = int(time.time() - start_time) - print(f"\nProvider timeout after {elapsed}s. Last status: {last_status}\n") - return False + """Wait for a ModelProvider to reach ``desired_status`` (optionally gateway-ready).""" + if not await self.models.wait_for_provider_status( + provider_name, desired_status, workspace=workspace, timeout=timeout + ): + return False + if desired_status == "READY" and check_gateway: + return await self.wait_for_gateway(provider_name, workspace=workspace, timeout=timeout) + return True - async def wait_for_deployment_status( + async def wait_for_gateway( self, - deployment_name: str, - desired_status: str, + provider_name: str, *, workspace: str | None = None, - timeout: int = 1200, + timeout: int = 60, ) -> bool: - """ - Wait for a ModelDeployment to reach the desired status (async version). - - For "DELETED" status, this function waits for the resource to be fully garbage - collected (404 NotFoundError), not just for the status to show as DELETED. + """Wait for the inference gateway to be able to route to a provider. - Args: - deployment_name: Name of the deployment - desired_status: Target status ("READY", "DELETED", etc.) - workspace: Workspace of the deployment - timeout: Maximum time to wait in seconds - - Returns: - True if desired status reached, False if timeout + Targets the separate inference-gateway service (not yet migrated to a + typed client), so this still calls through the ``AsyncNeMoPlatform`` SDK. """ + import asyncio + if workspace is None: workspace = self._client._get_workspace_path_param() start_time = time.time() - last_status = "" - last_message = "" - last_history_len = 0 - - print(f"Waiting for status: {desired_status}...\n") - + print("Waiting for gateway to be ready...") while time.time() - start_time < timeout: try: - deployment = await self._client.inference.deployments.retrieve(deployment_name, workspace=workspace) - history = getattr(deployment, "status_history", None) - created_at = getattr(deployment, "created_at", None) - # API guarantees last history entry is current state; fall back to top-level fields if no history - if history and len(history) > 0: - last_entry = history[-1] - current_status = getattr(last_entry, "status", deployment.status) - status_message = getattr(last_entry, "status_message", "") or "" - else: - current_status = deployment.status - status_message = deployment.status_message or "" - last_status = current_status - last_message = status_message - - # Only print status from history; elapsed shown is seconds since deployment creation - if history and len(history) > last_history_len: - for i in range(last_history_len, len(history)): - entry = history[i] - ts = getattr(entry, "timestamp", None) - ts_str = ts.strftime("%H:%M:%S") if hasattr(ts, "strftime") else str(ts) if ts else "" - st = getattr(entry, "status", "") - msg = getattr(entry, "status_message", "") or "" - secs = _seconds_since_creation(ts, created_at) - part = f" [{ts_str}] " - if secs is not None: - part += f"(+{secs}s) " - part += f"Status: {st}" - if msg: - part += f" - {msg}" - print(part) - last_history_len = len(history) - - # Check if we've reached the desired status - # For DELETED status, we need to wait for the actual 404 (garbage collection) - if current_status == desired_status and desired_status != "DELETED": - print(f"Deployment reached {desired_status} status!\n") - return True - - # Handle error states - if current_status == "ERROR": - print(f"Deployment entered ERROR state: {status_message}\n") - return False - + await self._client.inference.gateway.provider.ready(provider_name, workspace=workspace) + print(f" [{datetime.now().strftime('%H:%M:%S')}] Gateway is ready!\n") + return True except NotFoundError: - # For DELETED status, not found means success - if desired_status == "DELETED": - print(f"Deployment {desired_status}!\n") - return True - # For other statuses, not found is an error - print("Deployment not found\n") - return False - - await asyncio.sleep(3) - - # Timeout reached (wait_elapsed is time since we started polling) - wait_elapsed = int(time.time() - start_time) - detail = f"Last status: {last_status}" - if last_message: - detail += f" - {last_message}" - print(f"Timeout after {wait_elapsed}s. {detail}\n") + await asyncio.sleep(1) + except Exception: + await asyncio.sleep(1) + print(f"Gateway timeout after {int(time.time() - start_time)}s\n") return False diff --git a/services/automodel/src/nmp/automodel/app/jobs/compiler.py b/services/automodel/src/nmp/automodel/app/jobs/compiler.py index 78a99751f5..10bd79a10d 100644 --- a/services/automodel/src/nmp/automodel/app/jobs/compiler.py +++ b/services/automodel/src/nmp/automodel/app/jobs/compiler.py @@ -5,8 +5,9 @@ import logging -from nemo_platform import AsyncNeMoPlatform, NotFoundError -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import NotFoundError from nemo_platform_plugin.jobs.api_factory import ( ContainerSpec, CPUExecutionProviderSpec, @@ -17,6 +18,8 @@ ResourcesRequestsSpec, ResourcesSpec, ) +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ModelDeploymentConfig, ModelEntity from nmp.automodel.api.v2.jobs.schemas import ( CustomizationJobOutput, DeploymentParams, @@ -266,11 +269,13 @@ async def _resolve_deployment_config_ref( config_ref: str, workspace: str, sdk: AsyncNeMoPlatform, -): +) -> ModelDeploymentConfig: """Resolve a ``name`` or ``workspace/name`` string to a ModelDeploymentConfig.""" ref = parse_entity_ref(config_ref, default_workspace=workspace) + models = client_from_platform(sdk, AsyncModelsClient) try: - return await sdk.inference.deployment_configs.retrieve(name=ref.name, workspace=ref.workspace) + response = await models.get_deployment_config(name=ref.name, workspace=ref.workspace) + return response.data() except NotFoundError as e: raise PlatformJobCompilationError( f"deployment_config references '{config_ref}' which does not exist in workspace '{ref.workspace}'." @@ -315,7 +320,7 @@ async def _validate_deployment_config( resolved_config = await _resolve_deployment_config_ref(dc, workspace, sdk) # LoRA job referencing a config that has lora_enabled=False - if is_lora and resolved_config.nim_deployment and resolved_config.nim_deployment.lora_enabled is False: + if is_lora and resolved_config.model_spec.lora_enabled is False: raise PlatformJobCompilationError( f"deployment_config references '{dc}' which has lora_enabled=false, " "but this is a LoRA training job. The deployment would not load LoRA adapters. " @@ -326,7 +331,9 @@ async def _validate_deployment_config( if produces_new_model: output_name = transformed_spec.output.name try: - existing_me = await sdk.models.retrieve(name=output_name, workspace=workspace) + models = client_from_platform(sdk, AsyncModelsClient) + response = await models.get_model(name=output_name, workspace=workspace) + existing_me = response.data() except NotFoundError: # Output model entity doesn't exist yet, so a string # ref is inherently invalid -- it was created for a different model. @@ -338,9 +345,9 @@ async def _validate_deployment_config( # Output model entity already exists (retraining to create a new FileSet). # Verify the config actually targets this model entity. - nim = resolved_config.nim_deployment + model_spec = resolved_config.model_spec config_targets_model = (resolved_config.model_entity_id == f"{existing_me.workspace}/{existing_me.name}") or ( - nim and nim.model_name == existing_me.name and nim.model_namespace == existing_me.workspace + model_spec.model_name == existing_me.name and model_spec.model_namespace == existing_me.workspace ) if not config_targets_model: raise PlatformJobCompilationError( diff --git a/services/automodel/src/nmp/automodel/app/jobs/training/compiler.py b/services/automodel/src/nmp/automodel/app/jobs/training/compiler.py index 011c92219e..b7142e4df0 100644 --- a/services/automodel/src/nmp/automodel/app/jobs/training/compiler.py +++ b/services/automodel/src/nmp/automodel/app/jobs/training/compiler.py @@ -5,7 +5,6 @@ import logging -from nemo_platform.types.models.model_entity import ModelEntity from nemo_platform_plugin.jobs.api_factory import ( ContainerSpec, DistributedGPUExecutionProviderSpec, @@ -15,6 +14,7 @@ ResourcesSpec, StepLifecycle, ) +from nemo_platform_plugin.models.types import ModelEntity from nmp.automodel.api.v2.jobs.schemas import ( AnyTraining, CustomizationJobOutput, diff --git a/services/automodel/tests/test_compiler.py b/services/automodel/tests/test_compiler.py index 60464c8841..0d148fa799 100644 --- a/services/automodel/tests/test_compiler.py +++ b/services/automodel/tests/test_compiler.py @@ -10,7 +10,7 @@ import pytest from nemo_platform import AsyncNeMoPlatform -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.models.types import ModelEntity from nmp.automodel.adapter import automodel_spec_to_compiler_output from nmp.automodel.api.v2.jobs.schemas import CustomizationJobOutput, LoRAParams, OutputResponse, SFTTraining from nmp.automodel.app.jobs.compiler import _build_file_download_config @@ -39,13 +39,7 @@ def _make_mock_model_entity( @pytest.fixture def mock_sdk(): - sdk = Mock(spec=AsyncNeMoPlatform) - sdk.models = Mock() - sdk.models.retrieve = AsyncMock( - side_effect=lambda name, workspace, verbose=True: _make_mock_model_entity(workspace=workspace, name=name), - ) - sdk.files = Mock() - return sdk + return Mock(spec=AsyncNeMoPlatform) def _make_job_output() -> CustomizationJobOutput: diff --git a/services/automodel/tests/test_integrations_compiler.py b/services/automodel/tests/test_integrations_compiler.py index 8982a53416..1c93844375 100644 --- a/services/automodel/tests/test_integrations_compiler.py +++ b/services/automodel/tests/test_integrations_compiler.py @@ -4,8 +4,8 @@ from datetime import datetime import pytest -from nemo_platform.types.models.model_entity import ModelEntity from nemo_platform_plugin.integrations import IntegrationsSpec +from nemo_platform_plugin.models.types import ModelEntity from nmp.automodel.api.v2.jobs.schemas import ( CustomizationJobOutput, LoRAParams, diff --git a/services/automodel/tests/test_models_client_migration.py b/services/automodel/tests/test_models_client_migration.py new file mode 100644 index 0000000000..2d1db93956 --- /dev/null +++ b/services/automodel/tests/test_models_client_migration.py @@ -0,0 +1,25 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from nemo_platform_plugin.models.client import AsyncModelsClient +from nmp.automodel.app.jobs.compiler import _resolve_deployment_config_ref + + +@pytest.mark.asyncio +async def test_resolve_deployment_config_uses_models_client() -> None: + sdk = MagicMock() + deployment_config = MagicMock() + response = MagicMock() + response.data.return_value = deployment_config + models = MagicMock() + models.get_deployment_config = AsyncMock(return_value=response) + + with patch("nmp.automodel.app.jobs.compiler.client_from_platform", return_value=models) as make_client: + result = await _resolve_deployment_config_ref("other/config", "default", sdk) + + make_client.assert_called_once_with(sdk, AsyncModelsClient) + models.get_deployment_config.assert_awaited_once_with(name="config", workspace="other") + assert result is deployment_config diff --git a/services/core/entities/src/nmp/core/entities/controllers/workspace_cleanup.py b/services/core/entities/src/nmp/core/entities/controllers/workspace_cleanup.py index 538d419ecc..fb7198ecdc 100644 --- a/services/core/entities/src/nmp/core/entities/controllers/workspace_cleanup.py +++ b/services/core/entities/src/nmp/core/entities/controllers/workspace_cleanup.py @@ -10,6 +10,7 @@ from nemo_platform_plugin.files.client import AsyncFilesClient from nemo_platform_plugin.jobs.client import AsyncJobsClient from nemo_platform_plugin.jobs.schemas import PlatformJobStatus +from nemo_platform_plugin.models.client import AsyncModelsClient from nmp.common.api.filter import ComparisonOperation, FilterOperator from nmp.common.controller.controller import Controller, HeartbeatMixin from nmp.common.observability import start_span_with_ctx @@ -161,13 +162,14 @@ async def _cleanup_jobs(self, workspace: Workspace) -> None: async def _cleanup_deployments(self, workspace: Workspace) -> None: logger.info(f"Cleaning up deployments for workspace: {workspace.name}") try: - deployments_response = await self._nmp_sdk.inference.deployments.list(workspace=workspace.name) - deployments = [deployment async for deployment in deployments_response] + models_client = client_from_platform(self._nmp_sdk, AsyncModelsClient) + deployments_response = await models_client.list_deployments(workspace=workspace.name) + deployments = [deployment async for deployment in deployments_response.items()] for deployment in deployments: try: logger.info(f"Deleting deployment: {deployment.name}") - await self._nmp_sdk.inference.deployments.delete( + await models_client.delete_deployment( name=deployment.name, workspace=workspace.name, ) diff --git a/services/core/entities/tests/controllers/test_workspace_cleanup.py b/services/core/entities/tests/controllers/test_workspace_cleanup.py index 2e0f294a39..7944ccc223 100644 --- a/services/core/entities/tests/controllers/test_workspace_cleanup.py +++ b/services/core/entities/tests/controllers/test_workspace_cleanup.py @@ -21,22 +21,6 @@ def _make_workspace(name: str = "test-workspace") -> Workspace: ) -class _AsyncIterator: - """Helper to mock async iterators returned by the NeMo Platform SDK.""" - - def __init__(self, items): - self._items = iter(items) - - def __aiter__(self): - return self - - async def __anext__(self): - try: - return next(self._items) - except StopIteration: - raise StopAsyncIteration - - class _MockAsyncPaginatedResponse: """Mock for AsyncNemoPaginatedResponse that exposes .items() as an async generator.""" @@ -71,48 +55,55 @@ def _make_jobs_client(jobs: list | None = None) -> MagicMock: return jobs_client -def _make_sdk( - jobs: list | None = None, - deployments: list | None = None, - filesets: list | None = None, -) -> tuple[MagicMock, AsyncMock]: - """Build a MagicMock SDK with async mocks wired to the correct paths. +def _make_models_client(deployments: list | None = None) -> MagicMock: + """Build a mock typed AsyncModelsClient. - Returns (sdk, mock_files_client) so tests can assert on files client calls. - Only deployments/filesets stay on the ``sdk.*`` / files-client accessors; jobs - are handled via the typed jobs client patched onto ``client_from_platform`` - (see ``_patch_jobs_client``). + Production routes deployment cleanup through ``client_from_platform(sdk, + AsyncModelsClient)`` and iterates ``(await models_client.list_deployments(...)).items()``, + then calls ``delete_deployment(name=..., workspace=...)`` per deployment. """ - sdk = MagicMock() - sdk.inference.deployments.list = AsyncMock(return_value=_AsyncIterator(deployments or [])) - sdk.inference.deployments.delete = AsyncMock() - mock_files = _make_mock_files_client(filesets) - return sdk, mock_files + models_client = MagicMock() + models_client.list_deployments = AsyncMock(return_value=_MockAsyncPaginatedResponse(deployments or [])) + models_client.delete_deployment = AsyncMock() + return models_client _CLIENT_FROM_PLATFORM_PATCH = "nmp.core.entities.controllers.workspace_cleanup.client_from_platform" def _patch_jobs_client(jobs_client: MagicMock): - """Patch ``client_from_platform`` in the workspace_cleanup module to return *jobs_client*.""" - return patch(_CLIENT_FROM_PLATFORM_PATCH, return_value=jobs_client) + """Patch ``client_from_platform`` to dispatch by requested client class. + + Returns *jobs_client* for ``AsyncJobsClient`` and safe empty mocks for the + deployment/fileset clients. Dispatching by class (rather than returning + *jobs_client* for every ``client_from_platform`` call) keeps ``_async_step`` + tests correct even when ``_cleanup_jobs`` succeeds and execution proceeds to + ``_cleanup_deployments`` and ``_cleanup_filesets``. + """ + return _patch_clients(jobs_client, _make_mock_files_client([]), _make_models_client([])) -def _patch_clients(jobs_client: MagicMock, files_client: MagicMock): +def _patch_clients(jobs_client: MagicMock, files_client: MagicMock, models_client: MagicMock | None = None): """Patch ``client_from_platform`` to dispatch by requested client class. - ``_async_step`` cleans up both jobs and filesets, so it calls - ``client_from_platform(sdk, AsyncJobsClient)`` and + ``_async_step`` cleans up jobs, deployments, and filesets, so it calls + ``client_from_platform(sdk, AsyncJobsClient)``, + ``client_from_platform(sdk, AsyncModelsClient)``, and ``client_from_platform(sdk, AsyncFilesClient)`` — return the matching mock. """ from nemo_platform_plugin.files.client import AsyncFilesClient from nemo_platform_plugin.jobs.client import AsyncJobsClient + from nemo_platform_plugin.models.client import AsyncModelsClient + + models = models_client if models_client is not None else _make_models_client([]) def _dispatch(_sdk, client_cls): if client_cls is AsyncFilesClient: return files_client if client_cls is AsyncJobsClient: return jobs_client + if client_cls is AsyncModelsClient: + return models raise AssertionError(f"unexpected client class: {client_cls!r}") return patch(_CLIENT_FROM_PLATFORM_PATCH, side_effect=_dispatch) @@ -213,12 +204,10 @@ async def test_successful_workspace_deletion(self): repo.mark_workspace_for_deletion.return_value = True sdk = MagicMock() - sdk.inference.deployments.list = AsyncMock(return_value=_AsyncIterator([])) - mock_files = _make_mock_files_client([]) controller = _make_controller(workspace_repo=repo, nmp_sdk=sdk) - with _patch_clients(_make_jobs_client([]), mock_files): + with _patch_clients(_make_jobs_client([]), mock_files, _make_models_client([])): await controller._async_step() repo.mark_workspace_for_deletion.assert_any_call( @@ -356,14 +345,12 @@ async def test_deletes_deployments(self): deployment = MagicMock() deployment.name = "test-deployment" - sdk = MagicMock() - sdk.inference.deployments.list = AsyncMock(return_value=_AsyncIterator([deployment])) - sdk.inference.deployments.delete = AsyncMock() - - controller = _make_controller(nmp_sdk=sdk) - await controller._cleanup_deployments(workspace) + models_client = _make_models_client([deployment]) + controller = _make_controller() + with patch(_CLIENT_FROM_PLATFORM_PATCH, return_value=models_client): + await controller._cleanup_deployments(workspace) - sdk.inference.deployments.delete.assert_awaited_once_with( + models_client.delete_deployment.assert_awaited_once_with( name="test-deployment", workspace="test-workspace", ) @@ -376,14 +363,14 @@ async def test_continues_on_individual_deployment_failure(self): dep2 = MagicMock() dep2.name = "dep2" - sdk = MagicMock() - sdk.inference.deployments.list = AsyncMock(return_value=_AsyncIterator([dep1, dep2])) - sdk.inference.deployments.delete = AsyncMock(side_effect=[Exception("fail"), None]) + models_client = _make_models_client([dep1, dep2]) + models_client.delete_deployment = AsyncMock(side_effect=[Exception("fail"), None]) - controller = _make_controller(nmp_sdk=sdk) - await controller._cleanup_deployments(workspace) + controller = _make_controller() + with patch(_CLIENT_FROM_PLATFORM_PATCH, return_value=models_client): + await controller._cleanup_deployments(workspace) - assert sdk.inference.deployments.delete.await_count == 2 + assert models_client.delete_deployment.await_count == 2 class TestWorkspaceCleanupFilesets: diff --git a/services/core/models/src/nmp/core/models/app/utils.py b/services/core/models/src/nmp/core/models/app/utils.py index b062636735..876f714fed 100644 --- a/services/core/models/src/nmp/core/models/app/utils.py +++ b/services/core/models/src/nmp/core/models/app/utils.py @@ -9,10 +9,6 @@ from logging import getLogger from typing import Generic, List, Optional, TypeVar -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_deployment_config import ModelDeploymentConfig -from nemo_platform.types.inference.model_provider import ModelProvider -from nemo_platform.types.models import ModelEntity from nemo_platform_plugin.k8s_naming import ( DNS_LABEL_MAX_LENGTH, DNS_SUBDOMAIN_MAX_LENGTH, @@ -20,6 +16,7 @@ k8s_safe_name, workspace_name_identity, ) +from nemo_platform_plugin.models.types import ModelDeployment, ModelDeploymentConfig, ModelEntity, ModelProvider from nmp.common.api.common import PaginationData from nmp.common.entities.constants import NAME_PATTERN as ENTITY_NAME_PATTERN from pydantic import BaseModel diff --git a/services/core/models/src/nmp/core/models/controllers/backends/backends.py b/services/core/models/src/nmp/core/models/controllers/backends/backends.py index 44d4dc5fe3..e1946403f6 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/backends.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/backends.py @@ -7,7 +7,7 @@ from typing import Any, Dict from nemo_platform import AsyncNeMoPlatform -from nemo_platform.types.inference import ModelDeploymentStatus +from nemo_platform_plugin.models.types import ModelDeploymentStatus from nmp.core.models.controllers.context import ModelContext from pydantic import BaseModel @@ -18,7 +18,7 @@ class DeploymentStatusUpdate(BaseModel): This is the message that the service backend returns to the controller for every operation. """ - status: ModelDeploymentStatus + status: ModelDeploymentStatus | str status_message: str = "" error_details: Dict[str, Any] | None = None host_url: str | None = None diff --git a/services/core/models/src/nmp/core/models/controllers/backends/common.py b/services/core/models/src/nmp/core/models/controllers/backends/common.py index fd9a903604..4424975514 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/common.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/common.py @@ -8,8 +8,8 @@ from typing import Any, Dict, List, Optional, Protocol from nemo_platform.types.inference.k8s_nim_operator_config import K8sNIMOperatorConfig -from nemo_platform.types.inference.model_deployment import ModelDeployment from nemo_platform.types.shared.tool_call_config import ToolCallConfig +from nemo_platform_plugin.models.types import ModelDeployment LOG_TAIL_LINES = 80 LOG_MAX_CHARS = 2048 diff --git a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/resolve.py b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/resolve.py index faa301bfd5..689918ca88 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/resolve.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/resolve.py @@ -6,9 +6,7 @@ from dataclasses import dataclass from urllib.parse import urljoin -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_deployment_config import ModelDeploymentConfig -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.models.types import ModelDeployment, ModelDeploymentConfig, ModelEntity from nmp.common.config import Runtime, get_platform_config from nmp.core.models.app import ModelWeightsType, get_model_weights_type, parse_model_name_revision from nmp.core.models.controllers.backends.common import DeploymentConfigView, deployment_config_view diff --git a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/status.py b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/status.py index 4a3fbce966..99e03995bf 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/status.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/status.py @@ -7,7 +7,7 @@ from nemo_deployments_plugin.entities import Deployment, Volume from nemo_deployments_plugin.types import Endpoint -from nemo_platform.types.inference import ModelDeploymentStatus +from nemo_platform_plugin.models.types import ModelDeploymentStatus from nmp.core.models.controllers.backends.backends import DeploymentStatusUpdate from nmp.core.models.controllers.backends.common import format_duration diff --git a/services/core/models/src/nmp/core/models/controllers/backends/vllm_compiler.py b/services/core/models/src/nmp/core/models/controllers/backends/vllm_compiler.py index 26b42acaea..eb3968fbd0 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/vllm_compiler.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/vllm_compiler.py @@ -16,7 +16,7 @@ from logging import getLogger from typing import Optional -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.models.types import ModelEntity from nmp.core.models.controllers.backends.common import DeploymentConfigView logger = getLogger(__name__) diff --git a/services/core/models/src/nmp/core/models/controllers/context.py b/services/core/models/src/nmp/core/models/controllers/context.py index 2df5a2c56f..67f09c6b84 100644 --- a/services/core/models/src/nmp/core/models/controllers/context.py +++ b/services/core/models/src/nmp/core/models/controllers/context.py @@ -6,11 +6,13 @@ from dataclasses import dataclass from typing import Optional -from nemo_platform.types.inference import ServedModelMapping -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_deployment_config import ModelDeploymentConfig -from nemo_platform.types.inference.model_provider import ModelProvider -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.models.types import ( + ModelDeployment, + ModelDeploymentConfig, + ModelEntity, + ModelProvider, + ServedModelMapping, +) @dataclass diff --git a/services/core/models/src/nmp/core/models/controllers/deployment_reconciler.py b/services/core/models/src/nmp/core/models/controllers/deployment_reconciler.py index a916362d6e..6873916af6 100644 --- a/services/core/models/src/nmp/core/models/controllers/deployment_reconciler.py +++ b/services/core/models/src/nmp/core/models/controllers/deployment_reconciler.py @@ -10,10 +10,19 @@ from logging import getLogger from typing import Awaitable, Callable, Optional -from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import ConflictError, NotFoundError -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_provider import ModelProvider +from nemo_platform_plugin.client.errors import ConflictError, NotFoundError +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ( + CreateModelProviderRequest, + ModelDeployment, + ModelDeploymentStatus, + ModelProvider, + UpdateModelDeploymentStatusRequest, + UpsertModelProviderRequest, +) +from nemo_platform_plugin.models.types import ( + ModelProviderStatus as PluginModelProviderStatus, +) from nmp.common.entities.utils import parse_entity_ref from nmp.core.models.config import ControllerConfig from nmp.core.models.controllers.backends.backends import DeploymentStatusUpdate, ServiceBackend @@ -142,7 +151,7 @@ class ModelDeploymentReconciler: def __init__( self, - models_sdk: AsyncNeMoPlatform, + models_client: AsyncModelsClient, backend_registry: BackendRegistry, controller_config: ControllerConfig, entity_cache: ModelEntityCache, @@ -151,14 +160,14 @@ def __init__( """Initialize the deployment reconciler. Args: - models_sdk: SDK client for Models API interactions + models_client: Typed client for Models API interactions backend_registry: Registry of available service backends controller_config: Controller configuration containing deployment settings entity_cache: Model Entity reads and staged writes for the current phase emit_heartbeat: Called as each unit of work finishes so a long pass is distinguishable from a stalled one """ - self._models_sdk = models_sdk + self._models_client = models_client self._backend_registry = backend_registry self._controller_config = controller_config self._entity_cache = entity_cache @@ -169,6 +178,30 @@ def __init__( max_delay_seconds=controller_config.drift_recovery_max_delay_seconds, ) + async def _update_deployment_status( + self, + *, + name: str, + workspace: str, + status: ModelDeploymentStatus | str, + version: int, + status_message: str = "", + model_provider_id: str | None = None, + ) -> None: + update_fields: dict[str, object] = { + "status": ModelDeploymentStatus(status), + "status_message": status_message, + } + if model_provider_id is not None: + update_fields["model_provider_id"] = model_provider_id + + await self._models_client.update_deployment_status( + name=name, + workspace=workspace, + body=UpdateModelDeploymentStatusRequest.model_validate(update_fields), + query_params={"version": str(version)}, + ) + def get_service_backend(self) -> ServiceBackend: """Get the service backend. @@ -189,6 +222,9 @@ async def reconcile_deployments(self, deployment_contexts: list[ModelContext]) - """ for ctx in deployment_contexts: deployment = ctx.model_deployment + if deployment is None: + logger.warning("Skipping deployment reconciliation for context with no model_deployment") + continue model_deployment_id = f"{deployment.workspace}/{deployment.name}" try: backend = self.get_service_backend() @@ -374,7 +410,7 @@ async def gc_error_deployments(self, error_deployments: list[ModelDeployment]) - if original_message: gc_message = f"{gc_message} Original error: {original_message}" - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="DELETING", @@ -442,7 +478,7 @@ async def _reconcile_individual_deployment( model_provider_id = await self._reconcile_model_provider(deployment, status_update, existing_provider) - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status=status_update.status, @@ -457,7 +493,7 @@ async def _reconcile_individual_deployment( except Exception as e: logger.exception(f"Failed to {action_description} deployment {model_deployment_id}: {e}") try: - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="ERROR", @@ -499,7 +535,7 @@ async def _handle_drift_recovery( attempts = cache.get_attempts(model_deployment_id) logger.error(f"Drift recovery failed for {model_deployment_id} after {attempts} attempts") try: - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="ERROR", @@ -538,7 +574,7 @@ async def _handle_drift_recovery( f"{status_update.status_message}" ) - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status=status_update.status, @@ -558,7 +594,7 @@ async def _handle_drift_recovery( # Update status to PENDING with error info for visibility, but don't set ERROR # The next cycle will retry (respecting backoff) and can detect if recovery succeeded try: - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="PENDING", @@ -596,7 +632,7 @@ async def _handle_unknown_status( attempts = cache.get_attempts(model_deployment_id) logger.error(f"Backend communication failed for {model_deployment_id} after {attempts} attempts") try: - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="ERROR", @@ -628,7 +664,7 @@ async def _handle_unknown_status( ) try: - await self._models_sdk.inference.deployments.update_status( + await self._update_deployment_status( name=deployment.name, workspace=deployment.workspace, status="UNKNOWN", @@ -704,23 +740,36 @@ async def _ensure_model_provider( provider_workspace, provider_name = _provider_ref.workspace, _provider_ref.name if not existing_provider: - existing_provider = await self._models_sdk.inference.providers.retrieve( - name=provider_name, - workspace=provider_workspace, - ) + existing_provider = ( + await self._models_client.get_provider( + name=provider_name, + workspace=provider_workspace, + ) + ).data() if existing_provider.host_url != host_url: logger.info( f"ModelProvider {deployment.model_provider_id} host_url changed from " f"{existing_provider.host_url} to {host_url}, updating provider" ) - await self._models_sdk.inference.providers.update( + await self._models_client.upsert_provider( name=provider_name, workspace=provider_workspace, - host_url=host_url, - description=existing_provider.description, - enabled_models=existing_provider.enabled_models, - status="READY", + body=UpsertModelProviderRequest( + project=existing_provider.project, + description=existing_provider.description, + host_url=host_url, + api_key_secret_name=existing_provider.api_key_secret_name, + enabled_models=existing_provider.enabled_models, + default_extra_body=existing_provider.default_extra_body, + default_extra_headers=existing_provider.default_extra_headers, + required_extra_body=existing_provider.required_extra_body, + required_extra_headers=existing_provider.required_extra_headers, + model_deployment_id=existing_provider.model_deployment_id, + status=PluginModelProviderStatus.READY, + status_message=None, + auth_header_format=existing_provider.auth_header_format, + ), ) else: logger.debug( @@ -743,7 +792,7 @@ async def _ensure_model_provider( provider_workspace = deployment.workspace try: - await self._models_sdk.inference.providers.retrieve( + await self._models_client.get_provider( name=provider_name, workspace=provider_workspace, ) @@ -756,14 +805,16 @@ async def _ensure_model_provider( except NotFoundError: logger.debug(f"Creating ModelProvider {provider_workspace}/{provider_name} for deployment") - await self._models_sdk.inference.providers.create( + await self._models_client.create_provider( workspace=provider_workspace, - name=provider_name, - host_url=host_url, - description=f"Auto-created provider for deployment {deployment.name}", - project=deployment.project, - model_deployment_id=model_deployment_id, - status="READY", + body=CreateModelProviderRequest( + name=provider_name, + host_url=host_url, + description=f"Auto-created provider for deployment {deployment.name}", + project=deployment.project, + model_deployment_id=model_deployment_id, + status=PluginModelProviderStatus.READY, + ), ) model_provider_id = f"{provider_workspace}/{provider_name}" @@ -783,10 +834,12 @@ async def _cleanup_model_entities_for_provider( """ try: # Get the provider to see what models it was serving - provider = await self._models_sdk.inference.providers.retrieve( - name=provider_name, - workspace=provider_workspace, - ) + provider = ( + await self._models_client.get_provider( + name=provider_name, + workspace=provider_workspace, + ) + ).data() if not provider.served_models: logger.debug(f"Provider {provider_id} has no served_models, no cleanup needed") @@ -868,7 +921,7 @@ async def _delete_model_provider(self, deployment: ModelDeployment) -> None: try: logger.info(f"Deleting ModelProvider {model_provider_id} for deployment {model_deployment_id}") - await self._models_sdk.inference.providers.delete( + await self._models_client.delete_provider( name=provider_name, workspace=provider_workspace, ) @@ -903,10 +956,10 @@ async def _handle_deleted_deployment(self, deployment: ModelDeployment) -> None: ) try: # Hard-delete this specific version by calling the delete API again on a DELETED deployment - await self._models_sdk.inference.deployments.versions.delete( - name=str(deployment.entity_version), # version number - workspace=deployment.workspace, # workspace - deployment=deployment.name, # deployment name + await self._models_client.delete_deployment_version( + name=str(deployment.entity_version), + workspace=deployment.workspace, + deployment=deployment.name, ) logger.info( f"Successfully hard-deleted deployment {model_deployment_id} version {deployment.entity_version}" diff --git a/services/core/models/src/nmp/core/models/controllers/entity_cache.py b/services/core/models/src/nmp/core/models/controllers/entity_cache.py index 8b4bd07489..06c620be2d 100644 --- a/services/core/models/src/nmp/core/models/controllers/entity_cache.py +++ b/services/core/models/src/nmp/core/models/controllers/entity_cache.py @@ -19,9 +19,13 @@ from logging import getLogger from typing import Callable -from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import ConflictError, NotFoundError -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.client.errors import ConflictError, NotFoundError +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ( + CreateModelEntityRequest, + ModelEntity, + UpdateModelEntityRequest, +) logger = getLogger(__name__) @@ -60,16 +64,16 @@ class ModelEntityCache: reads the same entity within a phase observes its own write. """ - def __init__(self, models_sdk: AsyncNeMoPlatform, emit_heartbeat: Callable[[], None]) -> None: + def __init__(self, models_client: AsyncModelsClient, emit_heartbeat: Callable[[], None]) -> None: """Initialize the cache. Args: - models_sdk: SDK client for Models API interactions + models_client: Typed client for Models API interactions emit_heartbeat: Called as each entity is read or written. Reading and writing are both proportional to the number of entities, so they have to report progress or a large batch looks like a stall. """ - self._models_sdk = models_sdk + self._models_client = models_client self._emit_heartbeat = emit_heartbeat self._entities: dict[tuple[str, str], ModelEntity] = {} self._pending: dict[tuple[str, str], _PendingEntity] = {} @@ -94,7 +98,8 @@ async def refresh(self) -> None: ) entities: dict[tuple[str, str], ModelEntity] = {} - async for entity in self._models_sdk.models.list(workspace="-", page_size=_PAGE_SIZE): + resp = await self._models_client.list_models(workspace="-", query_params={"page_size": _PAGE_SIZE}) + async for entity in resp.items(): entities[(entity.workspace, entity.name)] = entity self._emit_heartbeat() @@ -220,12 +225,17 @@ async def _create(self, workspace: str, name: str, staged: _PendingEntity) -> No create_kwargs["model_providers"] = list(staged.link_providers) try: - created = await self._models_sdk.models.create(workspace=workspace, name=name, **create_kwargs) + created = ( + await self._models_client.create_model( + workspace=workspace, + body=CreateModelEntityRequest.model_validate({"name": name, **create_kwargs}), + ) + ).data() except ConflictError: # Created concurrently; adopt it and apply the staged changes instead. logger.debug("Model Entity %s/%s already exists, applying staged changes", workspace, name) try: - existing = await self._models_sdk.models.retrieve(workspace=workspace, name=name) + existing = (await self._models_client.get_model(workspace=workspace, name=name)).data() except NotFoundError: return self._entities[(workspace, name)] = existing @@ -251,7 +261,13 @@ async def _update(self, workspace: str, name: str, staged: _PendingEntity, exist logger.debug("Model Entity %s/%s already matches desired state", workspace, name) return - updated = await self._models_sdk.models.update(workspace=workspace, name=name, **update_params) + updated = ( + await self._models_client.update_model( + workspace=workspace, + name=name, + body=UpdateModelEntityRequest.model_validate(update_params), + ) + ).data() if updated is not None: self._entities[(workspace, name)] = updated logger.debug("Updated Model Entity %s/%s: %s", workspace, name, sorted(update_params)) diff --git a/services/core/models/src/nmp/core/models/controllers/models_controller.py b/services/core/models/src/nmp/core/models/controllers/models_controller.py index a9d70e0081..b9b8ceaa75 100644 --- a/services/core/models/src/nmp/core/models/controllers/models_controller.py +++ b/services/core/models/src/nmp/core/models/controllers/models_controller.py @@ -2,16 +2,21 @@ # SPDX-License-Identifier: Apache-2.0 import asyncio +import json import threading from logging import getLogger from typing import Optional from nemo_platform import DefaultAsyncHttpxClient -from nemo_platform._exceptions import NotFoundError -from nemo_platform.types.inference import ModelDeploymentStatus -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_deployment_config import ModelDeploymentConfig -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import NotFoundError +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ( + ModelDeployment, + ModelDeploymentConfig, + ModelDeploymentStatus, + ModelEntity, +) from nmp.common.controller import Controller, HeartbeatMixin from nmp.common.entities.utils import parse_entity_ref from nmp.common.sdk_factory import get_async_platform_sdk @@ -27,11 +32,11 @@ logger = getLogger(__name__) NON_TERMINAL_STATES: list[ModelDeploymentStatus] = [ - "CREATED", - "PENDING", - "READY", - "DELETING", - "DELETED", # Poll DELETED deployments to clean them up after grace period + ModelDeploymentStatus.CREATED, + ModelDeploymentStatus.PENDING, + ModelDeploymentStatus.READY, + ModelDeploymentStatus.DELETING, + ModelDeploymentStatus.DELETED, # Poll DELETED deployments to clean them up after grace period ] @@ -59,27 +64,29 @@ def __init__( self._current_task: asyncio.Task | None = None # Use service principal for controller - runs in background thread without user context - self._models_sdk = get_async_platform_sdk( + self._platform_sdk = get_async_platform_sdk( as_service="models", internal=True, http_client=DefaultAsyncHttpxClient(), ) + self._models_client = client_from_platform(self._platform_sdk, AsyncModelsClient) self._service_backends = backend_registry.list_backends() # Shared by both reconcilers; re-read at the start of each phase that # writes entities so neither phase works from state the other has changed. - self._entity_cache = ModelEntityCache(models_sdk=self._models_sdk, emit_heartbeat=self.emit_heartbeat) + self._entity_cache = ModelEntityCache(models_client=self._models_client, emit_heartbeat=self.emit_heartbeat) # Initialize reconcilers self._deployment_reconciler = ModelDeploymentReconciler( - models_sdk=self._models_sdk, + models_client=self._models_client, backend_registry=backend_registry, controller_config=models_config.controller, entity_cache=self._entity_cache, emit_heartbeat=self.emit_heartbeat, ) self._provider_reconciler = ModelProviderReconciler( - models_sdk=self._models_sdk, + models_client=self._models_client, + platform_sdk=self._platform_sdk, controller_config=models_config.controller, entity_cache=self._entity_cache, emit_heartbeat=self.emit_heartbeat, @@ -96,7 +103,7 @@ def shutdown(self) -> None: """ if self._loop is not None and not self._loop.is_closed(): try: - self._loop.run_until_complete(self._models_sdk.close()) + self._loop.run_until_complete(self._platform_sdk.close()) except Exception as e: logger.warning(f"Error closing event loop: {e}") finally: @@ -124,7 +131,7 @@ def get_service_backend(self) -> ServiceBackend: return self._backend_registry.get_backend() async def _retrieve_deployment_config( - self, config_ref: str, config_version: str, deployment_workspace: str + self, config_ref: str, config_version: int | str, deployment_workspace: str ) -> ModelDeploymentConfig: """Retrieve the ModelDeploymentConfig from the API. @@ -141,12 +148,12 @@ async def _retrieve_deployment_config( workspace, name = ref.workspace, ref.name logger.debug(f"Fetching ModelDeploymentConfig {workspace}/{name}@{config_version}") - config = await self._models_sdk.inference.deployment_configs.versions.retrieve( - name=str(config_version), # version number - workspace=workspace, # workspace - config=name, # config name + response = await self._models_client.get_deployment_config_version( + name=str(config_version), + workspace=workspace, + config=name, ) - return config + return response.data() except Exception as e: logger.error(f"Failed to fetch ModelDeploymentConfig {config_ref}@{config_version}: {e}") raise @@ -236,10 +243,12 @@ async def _retrieve_model_entity( if revision or not self._entity_cache.loaded: # A revision resolves server-side and does not correspond to an # cache key, so it has to be fetched directly. - model_entity = await self._models_sdk.models.retrieve( - name=full_model_name, - workspace=workspace, - ) + model_entity = ( + await self._models_client.get_model( + name=full_model_name, + workspace=workspace, + ) + ).data() else: model_entity = self._entity_cache.get(workspace, model_name) if model_entity is None: @@ -284,16 +293,18 @@ async def retrieve_non_terminal_deployments(self) -> list[ModelContext]: try: logger.debug(f"Querying ModelDeployments with status: {status} across all workspaces") # SDK returns AsyncPaginator - iterate through all pages - resp = self._models_sdk.inference.deployments.list( + resp = await self._models_client.list_deployments( workspace="-", # Cross-workspace query - filter={"status": status}, - all_versions=True, - page_size=1000, + query_params={ + "filter": json.dumps({"status": status}), + "all_versions": True, + "page_size": 1000, + }, ) logger.debug(f"Got paginator response for status {status}, iterating...") # Collect all deployments from paginator - deployments = [deployment async for deployment in resp] + deployments = [deployment async for deployment in resp.items()] logger.debug(f"Iteration complete for status {status}, got {len(deployments)} deployment(s)") if deployments: @@ -323,10 +334,12 @@ async def retrieve_non_terminal_deployments(self) -> list[ModelContext]: try: _prov_ref = parse_entity_ref(deployment.model_provider_id) provider_workspace, provider_name = _prov_ref.workspace, _prov_ref.name - provider = await self._models_sdk.inference.providers.retrieve( - name=provider_name, - workspace=provider_workspace, - ) + provider = ( + await self._models_client.get_provider( + name=provider_name, + workspace=provider_workspace, + ) + ).data() except Exception as e: logger.warning( f"Failed to fetch provider for deployment {deployment.workspace}/{deployment.name}: {e}" @@ -364,11 +377,11 @@ async def retrieve_model_providers(self) -> list[ModelContext] | None: provider_contexts: list[ModelContext] = [] try: - providers = self._models_sdk.inference.providers.list( + providers = await self._models_client.list_providers( workspace="-", # Cross-workspace query ) - async for provider in providers: + async for provider in providers.items(): deployment = None config = None entity = None @@ -378,10 +391,12 @@ async def retrieve_model_providers(self) -> list[ModelContext] | None: try: _depl_ref = parse_entity_ref(provider.model_deployment_id) deployment_workspace, deployment_name = _depl_ref.workspace, _depl_ref.name - deployment = await self._models_sdk.inference.deployments.retrieve( - deployment_name, - workspace=deployment_workspace, - ) + deployment = ( + await self._models_client.get_deployment( + name=deployment_name, + workspace=deployment_workspace, + ) + ).data() # Fetch config if deployment has config reference if deployment and deployment.config and deployment.config_version: @@ -424,13 +439,15 @@ async def retrieve_error_deployments(self) -> list[ModelDeployment]: since GC only needs the deployment itself and its timestamps. """ try: - resp = self._models_sdk.inference.deployments.list( + resp = await self._models_client.list_deployments( workspace="-", - filter={"status": "ERROR"}, - all_versions=True, - page_size=1000, + query_params={ + "filter": json.dumps({"status": "ERROR"}), + "all_versions": True, + "page_size": 1000, + }, ) - return [deployment async for deployment in resp] + return [deployment async for deployment in resp.items()] except Exception: logger.warning("Error querying ERROR deployments for GC", exc_info=True) return [] @@ -467,7 +484,9 @@ async def async_controller_step(self) -> None: await self._deployment_reconciler.reconcile_deployments(deployment_contexts) known_deployment_ids = { - f"{ctx.model_deployment.workspace}/{ctx.model_deployment.name}" for ctx in deployment_contexts + f"{deployment.workspace}/{deployment.name}" + for ctx in deployment_contexts + if (deployment := ctx.model_deployment) is not None } await self._deployment_reconciler.reconcile_orphans(known_deployment_ids) self.emit_heartbeat() diff --git a/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py b/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py index 120a720697..2989e8afe7 100644 --- a/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py +++ b/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py @@ -10,13 +10,26 @@ from typing import Callable, TypedDict from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import APIStatusError, ConflictError, NotFoundError -from nemo_platform.types.inference import ServedModelMapping -from nemo_platform.types.inference.model_deployment import ModelDeployment -from nemo_platform.types.inference.model_deployment_config import ModelDeploymentConfig -from nemo_platform.types.inference.model_provider import ModelProvider + +# IGW discovery and VirtualModel operations still go through the umbrella SDK, so +# their exception types are the Stainless ones and are named to say so. +from nemo_platform._exceptions import APIError, APIStatusError +from nemo_platform._exceptions import ConflictError as StainlessConflictError +from nemo_platform._exceptions import NotFoundError as StainlessNotFoundError from nemo_platform.types.inference.virtual_model import VirtualModel -from nemo_platform.types.models.model_entity import ModelEntity +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import ( + APIEndpointData, + ModelDeployment, + ModelDeploymentConfig, + ModelEntity, + ModelProvider, + ServedModelMapping, + UpdateModelProviderStatusRequest, +) +from nemo_platform_plugin.models.types import ( + ModelProviderStatus as PluginModelProviderStatus, +) from nmp.common.datetime_utils import ensure_utc from nmp.common.entities.constants import NAME_PATTERN from nmp.common.entities.utils import parse_entity_ref @@ -30,6 +43,7 @@ from nmp.core.models.controllers.context import ModelContext from nmp.core.models.controllers.entity_cache import ModelEntityCache from nmp.core.models.schemas import BackendFormat, ModelProviderStatus +from pydantic import AnyUrl logger = getLogger(__name__) @@ -94,14 +108,6 @@ class DiscoveredModel(TypedDict, total=False): parent: str | None -class ApiEndpointDict(TypedDict): - """Shape of :attr:`ArtifactDetails.api_endpoint`, stored on Model Entities.""" - - url: str - model_id: str - format: str - - class DiscoverySuccess(DiscoveryResult): """Provider responded with valid OpenAI model list. @@ -143,7 +149,7 @@ class ArtifactDetails: """Artifact and API endpoint details resolved for a model entity.""" fileset_url: str | None = None - api_endpoint: ApiEndpointDict | None = field(default=None) + api_endpoint: APIEndpointData | None = field(default=None) # --------------------------------------------------------------------------- @@ -294,7 +300,8 @@ class ModelProviderReconciler: def __init__( self, - models_sdk: AsyncNeMoPlatform, + models_client: AsyncModelsClient, + platform_sdk: AsyncNeMoPlatform, controller_config: ControllerConfig, entity_cache: ModelEntityCache, emit_heartbeat: Callable[[], None], @@ -302,20 +309,45 @@ def __init__( """Initialize the provider reconciler. Args: - models_sdk: SDK client for Models API interactions + models_client: Typed client for Models API interactions + platform_sdk: Umbrella SDK retained for IGW discovery and VirtualModel operations controller_config: Models controller configuration (discovery timeout/retry policy) entity_cache: Model Entity reads and staged writes for the current phase emit_heartbeat: Called as each unit of work finishes so a long pass is distinguishable from a stalled one """ - self._models_sdk = models_sdk + self._models_client = models_client + self._platform_sdk = platform_sdk self._controller_config = controller_config self._entity_cache = entity_cache self._emit_heartbeat = emit_heartbeat - self._discovery_sdk = models_sdk.with_options( + self._discovery_sdk = platform_sdk.with_options( max_retries=controller_config.provider_discovery_max_retries, ) + async def _update_provider_status( + self, + provider: ModelProvider, + *, + served_models: list[ServedModelMapping] | None = None, + status: str | None = None, + status_message: str | None = None, + ) -> ModelProvider: + update_fields: dict[str, object] = {} + if served_models is not None: + update_fields["served_models"] = served_models + if status is not None: + update_fields["status"] = PluginModelProviderStatus(status) + if status_message is not None: + update_fields["status_message"] = status_message + + response = await self._models_client.update_provider_status( + name=provider.name, + workspace=provider.workspace, + body=UpdateModelProviderStatusRequest.model_validate(update_fields), + ) + return response.data() + # ------------------------------------------------------------------------- # Public entry point # ------------------------------------------------------------------------- @@ -385,15 +417,20 @@ async def _load_virtual_models(self) -> tuple[list[VirtualModel], set[tuple[str, Returns ``None`` if the listing failed, which callers read as "the VirtualModel state is unknown this pass" and skip VirtualModel work rather than acting on a partial view. + + Only API failures are treated that way. A bug in this method is not an + unreachable service, and catching everything here once turned an + ``AttributeError`` into a routine "listing failed" warning that silently + disabled orphan cleanup for every pass while the tests stayed green. """ vm_snapshot: list[VirtualModel] = [] try: - async for virtual_model in self._models_sdk.inference.virtual_models.list( + async for virtual_model in self._platform_sdk.inference.virtual_models.list( workspace="-", page_size=_VIRTUAL_MODEL_PAGE_SIZE ): vm_snapshot.append(virtual_model) self._emit_heartbeat() - except Exception: + except APIError: logger.warning("Failed to list VirtualModels for provider reconciliation", exc_info=True) return None existing_vm_names = {(vm.workspace, vm.name) for vm in vm_snapshot} @@ -442,9 +479,8 @@ async def _reconcile_single_provider( ) ctx.served_models = [] try: - ctx.model_provider = await self._models_sdk.inference.providers.update_status( - name=provider.name, - workspace=provider.workspace, + ctx.model_provider = await self._update_provider_status( + provider, served_models=[], status="READY", status_message="Non-OpenAI compliant endpoint, model entity routing disabled", @@ -495,9 +531,8 @@ async def _reconcile_single_provider( logger.debug(f"Provider {provider_id}: serving {len(served_models)} model(s)") try: - ctx.model_provider = await self._models_sdk.inference.providers.update_status( - name=provider.name, - workspace=provider.workspace, + ctx.model_provider = await self._update_provider_status( + provider, served_models=served_models, status="READY", ) @@ -549,9 +584,8 @@ def _is_past_lost_threshold(self, provider: ModelProvider, provider_id: str, now async def _mark_lost(self, ctx: ModelContext, provider: ModelProvider, provider_id: str) -> None: """Write LOST status for a provider that has permanently failed discovery.""" try: - ctx.model_provider = await self._models_sdk.inference.providers.update_status( - name=provider.name, - workspace=provider.workspace, + ctx.model_provider = await self._update_provider_status( + provider, status="LOST", status_message="Provider discovery permanently failed. Delete and recreate to retry.", ) @@ -585,9 +619,8 @@ async def _on_transient_failure( updated_at = ensure_utc(provider.updated_at) if updated_at and (now - updated_at).total_seconds() > PROVIDER_ERROR_THRESHOLD_SECONDS: try: - await self._models_sdk.inference.providers.update_status( - name=provider.name, - workspace=provider.workspace, + await self._update_provider_status( + provider, status="ERROR", status_message=f"Provider discovery failed: {err.message}" if err.message @@ -613,9 +646,8 @@ async def _on_transient_failure( elif provider.status == ModelProviderStatus.ERROR: # Bump updated_at to pace the next retry try: - await self._models_sdk.inference.providers.update_status( - name=provider.name, - workspace=provider.workspace, + await self._update_provider_status( + provider, status="ERROR", status_message=f"Discovery retry failed: {err.message}" if err.message @@ -669,26 +701,32 @@ async def _discover_models(self, provider: ModelProvider) -> DiscoveryResult: timeout=self._controller_config.provider_discovery_timeout_seconds, ) - if not isinstance(models_response, dict) or "data" not in models_response: + if not isinstance(models_response, dict): logger.warning(f"Non-OpenAI compliant response format from {provider_id}") return DiscoveryNonCompliant() - discovered_models = models_response["data"] + discovered_models = models_response.get("data") if not isinstance(discovered_models, list): logger.warning(f"Non-OpenAI compliant data field from {provider_id}") return DiscoveryNonCompliant() - models = [] + models: list[DiscoveredModel] = [] for model in discovered_models: - if not isinstance(model, dict) or not isinstance(model.get("id"), str): + if not isinstance(model, dict): + logger.warning(f"Skipping invalid model entry in {provider_id}: {model}") + continue + model_id = model.get("id") + if not isinstance(model_id, str): logger.warning(f"Skipping invalid model entry in {provider_id}: {model}") continue + root = model.get("root") + parent = model.get("parent") models.append( - { - "id": model["id"], - "root": model.get("root"), - "parent": model.get("parent"), - } + DiscoveredModel( + id=model_id, + root=root if isinstance(root, str) else None, + parent=parent if isinstance(parent, str) else None, + ) ) return DiscoverySuccess(models) @@ -969,10 +1007,15 @@ async def _ensure_model_entity_for_provider( # entity are merged into a single write. existing_model_entity = self._entity_cache.get(model_workspace, model_name) + provider = ctx.model_provider + if provider is None: + logger.warning("Cannot ensure Model Entity %s/%s without a model provider", model_workspace, model_name) + return + details = await self._build_artifact_details( model_name, provider_id, - ctx.model_provider, + provider, existing_model_entity, ctx.model_deployment, ctx.model_deployment_config, @@ -1053,7 +1096,7 @@ async def _ensure_passthrough_virtual_model( return try: - await self._models_sdk.inference.virtual_models.create( + await self._platform_sdk.inference.virtual_models.create( workspace=workspace, name=model_name, default_model_entity=f"{workspace}/{model_name}", @@ -1064,7 +1107,7 @@ async def _ensure_passthrough_virtual_model( workspace, model_name, ) - except ConflictError: + except StainlessConflictError: pass # Already exists — nothing to do except Exception: logger.warning( @@ -1111,6 +1154,10 @@ async def _cleanup_orphaned_virtual_models( ): continue + if virtual_model.name is None: + logger.warning("Skipping autoprovisioned VirtualModel with no name") + continue + expected_db_version = _get_virtual_model_db_version(virtual_model) if expected_db_version is None: logger.warning( @@ -1121,7 +1168,7 @@ async def _cleanup_orphaned_virtual_models( continue try: - await self._models_sdk.inference.virtual_models.delete( + await self._platform_sdk.inference.virtual_models.delete( name=virtual_model.name, workspace=virtual_model.workspace, expected_db_version=expected_db_version, @@ -1131,7 +1178,7 @@ async def _cleanup_orphaned_virtual_models( virtual_model.workspace, virtual_model.name, ) - except NotFoundError: + except StainlessNotFoundError: # Already gone; the snapshot predates the deletion. logger.debug( "Orphaned autoprovisioned VirtualModel %s/%s no longer exists", @@ -1181,11 +1228,11 @@ async def _build_artifact_details( ) if weights_type == ModelWeightsType.EXTERNAL_PROVIDER: - details.api_endpoint = { - "url": provider.host_url, - "model_id": model_name, - "format": "openai", - } + details.api_endpoint = APIEndpointData( + url=AnyUrl(provider.host_url), + model_id=model_name, + format="openai", + ) logger.debug(f"Built api_endpoint for external provider: {provider.host_url}") elif weights_type == ModelWeightsType.HUGGINGFACE and config: diff --git a/services/core/models/src/nmp/core/models/sidecars/adapters/main.py b/services/core/models/src/nmp/core/models/sidecars/adapters/main.py index 43503c0808..45d6e30095 100644 --- a/services/core/models/src/nmp/core/models/sidecars/adapters/main.py +++ b/services/core/models/src/nmp/core/models/sidecars/adapters/main.py @@ -13,9 +13,11 @@ import urllib.error import urllib.request -from nemo_platform import NeMoPlatform, NotFoundError -from nemo_platform.types.models import ModelEntity -from nemo_platform.types.models.adapter import Adapter +from nemo_platform import NeMoPlatform +from nemo_platform import NotFoundError as StainlessNotFoundError +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.models.client import ModelsClient +from nemo_platform_plugin.models.types import Adapter, ModelEntity from nmp.common.config import get_platform_config from nmp.common.controller import ( Controller, @@ -92,6 +94,7 @@ def __init__(self, stop_signal: threading.Event | None = None): as_service="models", internal=True, ) + self._models_client = client_from_platform(self._sdk, ModelsClient) def download_fileset(self, dest_dir: str, workspace: str, name: str) -> bool: try: @@ -113,7 +116,7 @@ def download_fileset(self, dest_dir: str, workspace: str, name: str) -> bool: ) return True - except NotFoundError: + except StainlessNotFoundError: return False def step(self): @@ -156,13 +159,11 @@ def _update_prompt_tuned_models(self, dirs_to_keep: set[str]): # (AALGO-129): they remain single-workspace for now and continue to # use the bare model_entity.name as their on-disk directory. logger.info(f"Fetching prompt data for {self.workspace}/{self.model_name}") - model_entities: list[ModelEntity] = self._sdk.models.list( + model_entities = self._models_client.list_models( workspace=self.workspace, - filter={ - "base_model": self.model_name, - }, + query_params={"filter": json.dumps({"base_model": self.model_name})}, ) - for model_entity in model_entities: + for model_entity in model_entities.items(): if model_entity.prompt: dirs_to_keep.add(model_entity.name) prompt_tuned_model_dir = f"{self.nim_peft_source}/{model_entity.name}" @@ -297,22 +298,9 @@ def _download_adapter(self, adapter_dir: str, adapter: Adapter, adapter_workspac raise @staticmethod - def _resolve_adapter_workspace(adapter: Adapter, base_model_workspace: str) -> str: - """Return the workspace where ``adapter`` lives. - - The flat ``{adapter_ws}--{adapter_name}`` directory layout needs each - adapter's own workspace to disambiguate cross-workspace name - collisions. The internal :class:`Adapter` entity carries ``workspace``, - but the public API/SDK schema does not yet expose it. Until it does, - fall back to the base model's workspace: the on-disk encoding stays in - the ``{ws}--{name}`` form so every directory the reconciler sees is - decodable — just collapsed onto the base model's workspace, matching - the legacy single-workspace assumption. - - TODO: drop the ``getattr`` fallback once ``Adapter`` exposes - ``workspace`` in the SDK schema; the read becomes ``adapter.workspace``. - """ - return getattr(adapter, "workspace", None) or base_model_workspace + def _resolve_adapter_workspace(adapter: Adapter) -> str: + """Return the adapter's workspace from the typed Models API contract.""" + return adapter.workspace def _vllm_api_call(self, route: str, payload: dict) -> tuple[int, str]: """POST ``payload`` to the local vLLM server at ``route``. @@ -416,7 +404,10 @@ def _update_lora_adapters(self, dirs_to_keep: set[str]): """ logger.info(f"Fetching adapters for {self.workspace}/{self.model_name}") - model_entity: ModelEntity = self._sdk.models.retrieve(name=self.model_name, workspace=self.workspace) + model_entity: ModelEntity = self._models_client.get_model( + name=self.model_name, + workspace=self.workspace, + ).data() if not model_entity.adapters: return @@ -427,7 +418,7 @@ def _update_lora_adapters(self, dirs_to_keep: set[str]): logger.warning(f"Adapter {adapter.name} has no fileset, skipping") continue - adapter_workspace = self._resolve_adapter_workspace(adapter, model_entity.workspace) + adapter_workspace = self._resolve_adapter_workspace(adapter) dir_name = f"{adapter_workspace}--{adapter.name}" dirs_to_keep.add(dir_name) adapter_dir = f"{self.nim_peft_source}/{dir_name}" diff --git a/services/core/models/src/nmp/core/models/tasks/model_spec/run.py b/services/core/models/src/nmp/core/models/tasks/model_spec/run.py index 5cff0d3eba..5793c40ca7 100644 --- a/services/core/models/src/nmp/core/models/tasks/model_spec/run.py +++ b/services/core/models/src/nmp/core/models/tasks/model_spec/run.py @@ -23,14 +23,27 @@ APITimeoutError, InternalServerError, NeMoPlatform, - NeMoPlatformError, - NotFoundError, ) -from nemo_platform.types.models import ModelEntity from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import ( + InternalServerError as ClientInternalServerError, +) +from nemo_platform_plugin.client.errors import ( + NemoClientError, + NemoTransportError, + NotFoundError, +) from nemo_platform_plugin.files.client import FilesClient from nemo_platform_plugin.files.storage_config import HuggingfaceStorageConfig, LocalStorageConfig, NGCStorageConfig from nemo_platform_plugin.files.types import FilesetOutput +from nemo_platform_plugin.models.client import ModelsClient +from nemo_platform_plugin.models.types import ( + ModelEntity, + UpdateModelEntityRequest, +) +from nemo_platform_plugin.models.types import ( + ModelSpec as ClientModelSpec, +) from nmp.common.entities.utils import parse_entity_ref from nmp.common.model_utils import is_embedding_model from nmp.common.sdk_factory import get_platform_sdk @@ -75,6 +88,7 @@ class ModelSpecRunner: def __init__(self, sdk: NeMoPlatform, job_ctx: NMPJobContext): self.sdk = sdk + self.models_client = client_from_platform(sdk, ModelsClient) self.job_ctx = job_ctx @staticmethod @@ -157,7 +171,15 @@ def _merge_existing_spec(me: ModelEntity, model_spec: ModelSpec) -> None: @retry( stop=stop_after_attempt(MAX_RETRIES), wait=wait_exponential(multiplier=2, min=INITIAL_BACKOFF_SECONDS, max=MAX_BACKOFF_SECONDS), - retry=retry_if_exception_type((InternalServerError, APITimeoutError, APIConnectionError)), + retry=retry_if_exception_type( + ( + InternalServerError, + APITimeoutError, + APIConnectionError, + ClientInternalServerError, + NemoTransportError, + ) + ), reraise=True, ) def analyze_checkpoint(self, config: ModelSpecTaskConfig) -> ModelEntity: @@ -178,12 +200,18 @@ def analyze_checkpoint(self, config: ModelSpecTaskConfig) -> ModelEntity: logger.info(f"Fetching model entity: {config.workspace}/{config.name}") try: - me = self.sdk.models.retrieve(config.name, workspace=config.workspace, verbose=True) + me = self.models_client.get_model( + name=config.name, + workspace=config.workspace, + query_params={"verbose": True}, + ).data() except NotFoundError as err: raise ModelSpecCreationError( f"Failed to create model spec: model entity {config.workspace}/{config.name} does not exist" ) from err - except NeMoPlatformError as err: + except (ClientInternalServerError, NemoTransportError): + raise + except NemoClientError as err: raise ModelSpecCreationError( f"Failed to create model spec: model entity {config.workspace}/{config.name} unable to be fetched" ) from err @@ -241,14 +269,14 @@ def analyze_checkpoint(self, config: ModelSpecTaskConfig) -> ModelEntity: self.sdk.files.download( remote_path=non_tensor_files, - local_path=dest_dir, + local_path=str(dest_dir), fileset=fs.name, workspace=fs.workspace, ) logger.info(os.listdir(dest_dir)) is_trusted = me.trust_remote_code if me.trust_remote_code is not None else False - model_spec = infer_model_cfg_from_hf(dest_dir, is_trusted=is_trusted, file_listing=all_file_paths) + model_spec = infer_model_cfg_from_hf(str(dest_dir), is_trusted=is_trusted, file_listing=all_file_paths) # Embedding if model name or storage path contains "embed"; use or to avoid # overwriting a correct True from model name when storage path lacks "embed" model_spec.is_embedding_model = is_embedding_model(me.name) @@ -286,14 +314,21 @@ def analyze_checkpoint(self, config: ModelSpecTaskConfig) -> ModelEntity: self._merge_existing_spec(me, model_spec) try: - me: ModelEntity = self.sdk.models.update( - name=config.name, workspace=config.workspace, spec=model_spec, verbose=True - ) + me = self.models_client.update_model( + name=config.name, + workspace=config.workspace, + body=UpdateModelEntityRequest( + spec=ClientModelSpec.model_validate(model_spec.model_dump()), + ), + query_params={"verbose": True}, + ).data() except NotFoundError as err: raise ModelSpecCreationError( f"Failed to update model spec: model entity {config.workspace}/{config.name} does not exist" ) from err - except NeMoPlatformError as err: + except (ClientInternalServerError, NemoTransportError): + raise + except NemoClientError as err: raise ModelSpecCreationError( f"Failed to update model spec: model entity {config.workspace}/{config.name} unable to be fetched" ) from err diff --git a/services/core/models/src/nmp/core/models/tasks/model_spec/schemas.py b/services/core/models/src/nmp/core/models/tasks/model_spec/schemas.py index 029c39ad1a..208de5d6f3 100644 --- a/services/core/models/src/nmp/core/models/tasks/model_spec/schemas.py +++ b/services/core/models/src/nmp/core/models/tasks/model_spec/schemas.py @@ -44,8 +44,8 @@ class NMPJobContext: files_url: str | None models_url: str | None - storage_path: str | None - config_path: str | None + storage_path: Path + config_path: Path @classmethod def from_env(cls) -> Self: diff --git a/services/core/models/tests/unit/controllers/conftest.py b/services/core/models/tests/unit/controllers/conftest.py index 4825683a78..65f9d067b9 100644 --- a/services/core/models/tests/unit/controllers/conftest.py +++ b/services/core/models/tests/unit/controllers/conftest.py @@ -3,11 +3,20 @@ """Test fixtures for Models Controller tests.""" +import inspect +from typing import Generic, TypeVar, get_origin from unittest.mock import AsyncMock, MagicMock, patch +import nemo_platform_plugin.client.types as _client_types +import nemo_platform_plugin.models.types as _models_types import pytest +from nemo_platform_plugin.client.method import EndpointMethod +from nemo_platform_plugin.client.types import Paginated +from nemo_platform_plugin.models.client import AsyncModelsClient from nmp.common.config import PlatformConfig +T = TypeVar("T") + def platform_config( *, @@ -54,6 +63,16 @@ def mock_get_config_patch(mock_platform_config): yield +@pytest.fixture(autouse=True) +def mock_models_client_bridge(): + """Use the injected platform mock as the typed Models client in controller tests.""" + with patch( + "nmp.core.models.controllers.models_controller.client_from_platform", + side_effect=lambda platform, _client_cls: platform, + ): + yield + + @pytest.fixture def mock_sdk_class_patch(): """Patch get_async_platform_sdk factory function.""" @@ -73,16 +92,13 @@ def mock_asyncio_run_patch(): @pytest.fixture def mock_models_sdk(): - """Create a mock AsyncNeMoPlatform SDK for testing.""" - mock_sdk = MagicMock() - - # Set up the nested structure for v2.inference.deployments - mock_sdk.v2 = MagicMock() - mock_sdk.v2.inference = MagicMock() - mock_sdk.v2.inference.deployments = MagicMock() - mock_sdk.v2.inference.deployments.list = MagicMock() + """Mock platform SDK for controller tests. - return mock_sdk + The autouse ``mock_models_client_bridge`` makes ``client_from_platform`` the + identity, so this one mock stands in for both the umbrella SDK and the typed + Models client the controller builds from it. + """ + return make_models_client(strict=False) @pytest.fixture @@ -144,7 +160,8 @@ def _assert_controller_has_required_attributes(controller): """Assert that controller has all required attributes.""" assert hasattr(controller, "_is_healthy") assert hasattr(controller, "_backend_registry") - assert hasattr(controller, "_models_sdk") + assert hasattr(controller, "_platform_sdk") + assert hasattr(controller, "_models_client") def _assert_controller_healthy(controller, is_healthy=True): @@ -168,8 +185,8 @@ def _assert_asyncio_run_called_once(mock_asyncio_run_patch): def _assert_sdk_list_called_for_all_statuses(mock_models_sdk, non_terminal_states_count): - """Assert that SDK list method was called for each non-terminal status.""" - assert mock_models_sdk.inference.deployments.list.call_count == non_terminal_states_count + """Assert that the typed client's deployment list was called for each non-terminal status.""" + assert mock_models_sdk.list_deployments.call_count == non_terminal_states_count def _assert_deployments_count(deployments, expected_count): @@ -215,6 +232,88 @@ async def __anext__(self): return self._items.pop(0) +class PaginatedResponse: + """Stand-in for AsyncNemoPaginatedResponse: ``items()`` yields across pages.""" + + def __init__(self, items): + self._items = list(items) + + def items(self): + return AsyncPaginator(self._items) + + +def paginated(items=()): + """An AsyncMock return value for a typed-client ``list_*`` call.""" + return PaginatedResponse(items) + + +class Response(Generic[T]): + """Stand-in for NemoResponse: the payload is behind ``data()``.""" + + def __init__(self, data: T) -> None: + self._data = data + + def data(self) -> T: + return self._data + + +def response(data: T) -> Response[T]: + return Response(data) + + +def _endpoint_returns_paginated(descriptor: EndpointMethod) -> bool: + """Whether an endpoint's declared return type is ``Paginated[...]``. + + Read from the endpoint's own annotation rather than guessed from its name: + ``list_deployment_versions`` and ``list_deployment_config_versions`` return + plain lists, so a name-prefix rule hands them the wrong response shape. + """ + annotation = inspect.signature(inspect.unwrap(descriptor)).return_annotation + if isinstance(annotation, str): + try: + annotation = eval(annotation, {**vars(_models_types), **vars(_client_types)}) # noqa: S307 + except Exception: + return False + return annotation is Paginated or get_origin(annotation) is Paginated + + +# Derived from the real client: only the method() descriptors are endpoints, so +# the URL builders and the wait_for_* helpers keep the sync/return shapes they +# actually have, and a new endpoint is picked up here without an edit. +_MODELS_CLIENT_ENDPOINTS = { + name: _endpoint_returns_paginated(descriptor) + for name in dir(AsyncModelsClient) + if not name.startswith("_") + and isinstance(descriptor := inspect.getattr_static(AsyncModelsClient, name), EndpointMethod) +} + + +def make_models_client(*, strict: bool = True) -> MagicMock: + """A mock AsyncModelsClient whose endpoint calls are awaitable and typed-shaped. + + Paginated endpoints yield an empty page; the rest resolve to a ``Response`` + wrapping a MagicMock. Tests override the specific calls they care about. + + Only ``method()`` endpoints are wired. Everything else on the client keeps + whatever ``spec`` gives it, which is already right: the OpenAI route builders + are sync and return ``str``, and ``wait_for_*`` returns ``bool``, so forcing + them into the endpoint shape would hand tests a coroutine where production + has a string. + + ``strict=False`` drops the spec so the same mock can also stand in for the + umbrella platform SDK, which the models controller passes to the provider + reconciler for IGW and VirtualModel work. + """ + client = MagicMock(spec=AsyncModelsClient) if strict else MagicMock() + for name, paginated_return in _MODELS_CLIENT_ENDPOINTS.items(): + result = PaginatedResponse([]) if paginated_return else Response(MagicMock()) + setattr(client, name, AsyncMock(return_value=result)) + if not strict: + # No spec to supply it, and the provider reconciler calls it on the SDK. + client.send = AsyncMock(return_value=Response(MagicMock())) + return client + + _ENTITY_FIELDS = ("model_providers", "fileset", "api_endpoint", "backend_format") @@ -238,7 +337,7 @@ def _copy(update): return entity -async def seed_entity_cache(mock_models_sdk, entity_cache, entities=()): - """Load the cache from the mock SDK so lookups resolve to ``entities``.""" - mock_models_sdk.models.list = MagicMock(return_value=AsyncPaginator(list(entities))) +async def seed_entity_cache(mock_models_client, entity_cache, entities=()): + """Load the cache from the mock Models client so lookups resolve to ``entities``.""" + mock_models_client.list_models = AsyncMock(return_value=PaginatedResponse(list(entities))) await entity_cache.refresh() diff --git a/services/core/models/tests/unit/controllers/test_deployment_reconciler.py b/services/core/models/tests/unit/controllers/test_deployment_reconciler.py index 2c2dfa7b8b..04ae1bb837 100644 --- a/services/core/models/tests/unit/controllers/test_deployment_reconciler.py +++ b/services/core/models/tests/unit/controllers/test_deployment_reconciler.py @@ -5,11 +5,19 @@ import logging from datetime import datetime, timedelta, timezone +from typing import Generic, TypeVar from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest -from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import ConflictError, NotFoundError +from nemo_platform_plugin.client.errors import ConflictError, NemoHTTPError, NotFoundError +from nemo_platform_plugin.models.types import ( + AuthContext, + ModelProvider, + ModelProviderStatus, + ServedModelMapping, + UpsertModelProviderRequest, +) from nmp.core.models.config import ControllerConfig from nmp.core.models.controllers.backends.backends import DeploymentStatusUpdate from nmp.core.models.controllers.backends.registry import BackendRegistry @@ -18,11 +26,29 @@ from nmp.core.models.controllers.entity_cache import ModelEntityCache from nmp.core.models.schemas import ModelDeployment -from .conftest import AsyncPaginator, make_entity, seed_entity_cache +from .conftest import AsyncPaginator, make_entity, make_models_client, seed_entity_cache + +T = TypeVar("T") _AsyncPaginator = AsyncPaginator +class _Response(Generic[T]): + def __init__(self, data: T) -> None: + self._data = data + + def data(self) -> T: + return self._data + + +def _response(data: T) -> _Response[T]: + return _Response(data) + + +def _client_error(error_type: type[NemoHTTPError], status_code: int) -> NemoHTTPError: + return error_type(httpx.Response(status_code, request=httpx.Request("GET", "http://test"))) + + def _entity(workspace, name, model_providers): """Model Entity stand-in addressable by the cache.""" return make_entity(workspace, name, model_providers=model_providers) @@ -30,10 +56,8 @@ def _entity(workspace, name, model_providers): @pytest.fixture def mock_models_sdk(): - """Create a mock AsyncNeMoPlatform SDK.""" - sdk = MagicMock(spec=AsyncNeMoPlatform) - sdk.models.list = MagicMock(return_value=_AsyncPaginator([])) - return sdk + """Create a mock typed Models client.""" + return make_models_client() @pytest.fixture @@ -50,8 +74,8 @@ def controller_config(): @pytest.fixture def entity_cache(mock_models_sdk): - """Model Entity cache backed by the mock SDK.""" - return ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None) + """Model Entity cache backed by the mock Models client.""" + return ModelEntityCache(models_client=mock_models_sdk, emit_heartbeat=lambda: None) @pytest.fixture @@ -64,7 +88,7 @@ def heartbeat_calls(): def reconciler(mock_models_sdk, mock_backend_registry, controller_config, entity_cache, heartbeat_calls): """Create a ModelDeploymentReconciler instance.""" return ModelDeploymentReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_sdk, backend_registry=mock_backend_registry, controller_config=controller_config, entity_cache=entity_cache, @@ -116,7 +140,7 @@ async def test_handle_created_deployment_success(reconciler, mock_backend_regist mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK update_status method - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Call the handler with the backend function await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") @@ -124,15 +148,16 @@ async def test_handle_created_deployment_success(reconciler, mock_backend_regist # Verify backend was called mock_backend.create_model_deployment.assert_called_once_with(deployment) - # Verify SDK update was called - reconciler._models_sdk.inference.deployments.update_status.assert_called_once_with( - name="test-deployment", - workspace="default", - status="PENDING", - version="v1", - status_message="Deployment created", - model_provider_id=None, # No provider created for PENDING status - ) + # Verify typed status request was sent. + reconciler._models_client.update_deployment_status.assert_called_once() + call = reconciler._models_client.update_deployment_status.call_args + assert call.kwargs["name"] == "test-deployment" + assert call.kwargs["workspace"] == "default" + assert call.kwargs["query_params"] == {"version": "v1"} + assert call.kwargs["body"].status.value == "PENDING" + assert call.kwargs["body"].status_message == "Deployment created" + assert call.kwargs["body"].model_provider_id is None + assert "model_provider_id" not in call.kwargs["body"].model_fields_set @pytest.mark.asyncio @@ -146,7 +171,7 @@ async def test_handle_created_deployment_backend_failure(reconciler, mock_backen mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK update_status method - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Call the handler - should not raise exception await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") @@ -155,13 +180,13 @@ async def test_handle_created_deployment_backend_failure(reconciler, mock_backen mock_backend.create_model_deployment.assert_called_once_with(deployment) # Verify SDK update was called with ERROR status - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs assert call_kwargs["name"] == "test-deployment" assert call_kwargs["workspace"] == "default" - assert call_kwargs["version"] == "v1" - assert call_kwargs["status"] == "ERROR" - assert "Failed to create deployment default/test-deployment" in call_kwargs["status_message"] + assert call_kwargs["query_params"]["version"] == "v1" + assert call_kwargs["body"].status.value == "ERROR" + assert "Failed to create deployment default/test-deployment" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -169,7 +194,7 @@ async def test_reconcile_individual_deployment_monitor_ready_no_message_logs_deb """Routine monitor + READY with no status message should log at DEBUG, not INFO.""" deployment = make_deployment(status="READY") status_update = DeploymentStatusUpdate(status="READY", status_message="", host_url=None) - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() with caplog.at_level(logging.DEBUG, logger="nmp.core.models.controllers.deployment_reconciler"): await reconciler._reconcile_individual_deployment( @@ -191,7 +216,7 @@ async def test_reconcile_individual_deployment_monitor_ready_with_message_logs_i """Monitor + READY with a status message stays at INFO.""" deployment = make_deployment(status="READY") status_update = DeploymentStatusUpdate(status="READY", status_message="NIM loading", host_url=None) - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() with caplog.at_level(logging.INFO, logger="nmp.core.models.controllers.deployment_reconciler"): await reconciler._reconcile_individual_deployment( @@ -222,16 +247,14 @@ async def test_reconcile_individual_deployment_conflict_is_noop(reconciler, mock mock_backend_registry.get_backend.return_value = mock_backend # Main status update conflicts (deployment was marked DELETING server-side) - reconciler._models_sdk.inference.deployments.update_status = AsyncMock( - side_effect=ConflictError("Conflict", response=MagicMock(), body=None) - ) + reconciler._models_client.update_deployment_status = AsyncMock(side_effect=_client_error(ConflictError, 409)) # Should not raise and should not attempt ERROR update await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "PENDING" + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "PENDING" @pytest.mark.asyncio @@ -246,16 +269,14 @@ async def test_reconcile_individual_deployment_error_fallback_conflict_is_noop( mock_backend.create_model_deployment = AsyncMock(side_effect=Exception("Backend error")) mock_backend_registry.get_backend.return_value = mock_backend - reconciler._models_sdk.inference.deployments.update_status = AsyncMock( - side_effect=ConflictError("Conflict", response=MagicMock(), body=None) - ) + reconciler._models_client.update_deployment_status = AsyncMock(side_effect=_client_error(ConflictError, 409)) # Should not raise if fallback ERROR update hits 409 conflict await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "ERROR" @pytest.mark.asyncio @@ -270,13 +291,13 @@ async def test_reconcile_created_backend_error_persisted(reconciler, mock_backen ) ) mock_backend_registry.get_backend.return_value = mock_backend - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock(return_value=_response(MagicMock())) await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" - assert call_kwargs["status_message"] == "Backend create failed for some reason" + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status == "ERROR" + assert call_kwargs["body"].status_message == "Backend create failed for some reason" @pytest.mark.asyncio @@ -314,7 +335,7 @@ async def test_reconcile_deployments_with_created_status(reconciler, mock_backen mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployments (now passing contexts with pre-fetched data) await reconciler.reconcile_deployments([created_context, pending_context]) @@ -328,7 +349,7 @@ async def test_reconcile_deployments_with_created_status(reconciler, mock_backen mock_backend.get_model_deployment_status.assert_called_once_with(pending_context) # Verify SDK update was called twice (once for each deployment) - assert reconciler._models_sdk.inference.deployments.update_status.call_count == 2 + assert reconciler._models_client.update_deployment_status.call_count == 2 # ============================================================================ @@ -340,10 +361,8 @@ async def test_reconcile_deployments_with_created_status(reconciler, mock_backen async def test_ensure_model_provider_creates_when_not_exists(reconciler, make_deployment): """Test that ensure_model_provider creates provider when it doesn't exist.""" # Mock provider doesn't exist (retrieve raises NotFoundError) - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", project="test-project") @@ -354,21 +373,21 @@ async def test_ensure_model_provider_creates_when_not_exists(reconciler, make_de assert model_provider_id == "test-ns/test-deployment" # Verify retrieve was called to check existence - reconciler._models_sdk.inference.providers.retrieve.assert_called_once_with( + reconciler._models_client.get_provider.assert_called_once_with( name="test-deployment", workspace="test-ns", ) - # Verify create was called with correct parameters including model_deployment_id and status - reconciler._models_sdk.inference.providers.create.assert_called_once_with( - workspace="test-ns", - name="test-deployment", - host_url="http://test-ns/test-deployment", - description="Auto-created provider for deployment test-deployment", - project="test-project", - model_deployment_id="test-ns/test-deployment", - status="READY", - ) + # Verify typed create request including model_deployment_id and status. + reconciler._models_client.create_provider.assert_called_once() + call = reconciler._models_client.create_provider.call_args + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].name == "test-deployment" + assert call.kwargs["body"].host_url == "http://test-ns/test-deployment" + assert call.kwargs["body"].description == "Auto-created provider for deployment test-deployment" + assert call.kwargs["body"].project == "test-project" + assert call.kwargs["body"].model_deployment_id == "test-ns/test-deployment" + assert call.kwargs["body"].status.value == "READY" @pytest.mark.asyncio @@ -382,8 +401,8 @@ async def test_ensure_model_provider_handles_name_collision(mock_uuid, reconcile # Mock provider exists (retrieve succeeds on first call - collision) mock_provider = MagicMock() - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) + reconciler._models_client.create_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", project="test-project") @@ -394,18 +413,16 @@ async def test_ensure_model_provider_handles_name_collision(mock_uuid, reconcile assert model_provider_id == "test-ns/test-deployment_abcdef12" # Verify retrieve was called to check existence - reconciler._models_sdk.inference.providers.retrieve.assert_called_once() + reconciler._models_client.get_provider.assert_called_once() - # Verify create was called with UUID-suffixed name and status - reconciler._models_sdk.inference.providers.create.assert_called_once_with( - workspace="test-ns", - name="test-deployment_abcdef12", - host_url="http://test-ns/test-deployment", - description="Auto-created provider for deployment test-deployment", - project="test-project", - model_deployment_id="test-ns/test-deployment", - status="READY", - ) + # Verify typed create request has the UUID-suffixed name. + reconciler._models_client.create_provider.assert_called_once() + call = reconciler._models_client.create_provider.call_args + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].name == "test-deployment_abcdef12" + assert call.kwargs["body"].host_url == "http://test-ns/test-deployment" + assert call.kwargs["body"].model_deployment_id == "test-ns/test-deployment" + assert call.kwargs["body"].status.value == "READY" @pytest.mark.asyncio @@ -416,9 +433,9 @@ async def test_ensure_model_provider_reuses_existing_when_already_set(reconciler mock_provider.host_url = "http://test-ns/test-deployment" mock_provider.description = "Existing provider" mock_provider.enabled_models = None - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) - reconciler._models_sdk.inference.providers.create = AsyncMock() - reconciler._models_sdk.inference.providers.update = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) + reconciler._models_client.create_provider = AsyncMock() + reconciler._models_client.upsert_provider = AsyncMock() deployment = make_deployment( workspace="test-ns", project="test-project", model_provider_id="test-ns/existing-provider" @@ -431,29 +448,47 @@ async def test_ensure_model_provider_reuses_existing_when_already_set(reconciler assert model_provider_id == "test-ns/existing-provider" # Verify retrieve was called to check the existing provider exists - reconciler._models_sdk.inference.providers.retrieve.assert_called_once_with( + reconciler._models_client.get_provider.assert_called_once_with( name="existing-provider", workspace="test-ns", ) # Verify create was NOT called since we're reusing existing provider - reconciler._models_sdk.inference.providers.create.assert_not_called() + reconciler._models_client.create_provider.assert_not_called() # Verify update was NOT called since host_url matches - reconciler._models_sdk.inference.providers.update.assert_not_called() + reconciler._models_client.upsert_provider.assert_not_called() @pytest.mark.asyncio -async def test_ensure_model_provider_updates_when_host_url_changes(reconciler, make_deployment): - """Test that ensure_model_provider updates existing provider when host_url changes.""" - # Mock provider exists (retrieve succeeds) with different host_url - mock_provider = MagicMock() - mock_provider.host_url = "http://old-host/test-deployment" - mock_provider.description = "Existing provider" - mock_provider.enabled_models = ["model1", "model2"] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) - reconciler._models_sdk.inference.providers.create = AsyncMock() - reconciler._models_sdk.inference.providers.update = AsyncMock() +async def test_ensure_model_provider_updates_host_without_losing_provider_metadata(reconciler, make_deployment): + """A host-change PUT preserves metadata and retains the legacy READY transition.""" + now = datetime.now(timezone.utc) + existing_provider = ModelProvider( + id="provider-id", + name="existing-provider", + workspace="test-ns", + project="projects/test-project", + created_at=now, + updated_at=now, + description="Existing provider", + host_url="http://old-host/test-deployment", + api_key_secret_name="provider-api-key", + served_models=[ServedModelMapping(model_entity_id="test-ns/model1", served_model_name="model1")], + enabled_models=["model1", "model2"], + status=ModelProviderStatus.ERROR, + status_message="Existing status message", + default_extra_body={"temperature": 0.25}, + default_extra_headers={"X-Default": "default"}, + required_extra_body={"seed": 7}, + required_extra_headers={"X-Required": "required"}, + model_deployment_id="test-ns/test-deployment", + auth_context=AuthContext(principal_id="original-user", principal_groups=["model-users"]), + auth_header_format="X-Api-Key: {{ auth_secret }}", + ) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(existing_provider)) + reconciler._models_client.create_provider = AsyncMock() + reconciler._models_client.upsert_provider = AsyncMock() deployment = make_deployment( workspace="test-ns", project="test-project", model_provider_id="test-ns/existing-provider" @@ -462,37 +497,49 @@ async def test_ensure_model_provider_updates_when_host_url_changes(reconciler, m new_host_url = "http://new-host/test-deployment" model_provider_id = await reconciler._ensure_model_provider(deployment, new_host_url) - # Verify the existing provider ID is returned assert model_provider_id == "test-ns/existing-provider" - - # Verify retrieve was called to check the existing provider - reconciler._models_sdk.inference.providers.retrieve.assert_called_once_with( - name="existing-provider", - workspace="test-ns", - ) - - # Verify update was called with new host_url, existing metadata, and status - reconciler._models_sdk.inference.providers.update.assert_called_once_with( + reconciler._models_client.get_provider.assert_called_once_with( name="existing-provider", workspace="test-ns", - host_url=new_host_url, - description="Existing provider", - enabled_models=["model1", "model2"], - status="READY", ) - # Verify create was NOT called since we're updating existing provider - reconciler._models_sdk.inference.providers.create.assert_not_called() + reconciler._models_client.upsert_provider.assert_called_once() + call = reconciler._models_client.upsert_provider.call_args + assert call.kwargs["name"] == "existing-provider" + assert call.kwargs["workspace"] == "test-ns" + body = call.kwargs["body"] + assert isinstance(body, UpsertModelProviderRequest) + assert body.model_fields_set == set(UpsertModelProviderRequest.model_fields) + assert body.model_dump(mode="json") == { + "project": "projects/test-project", + "description": "Existing provider", + "host_url": new_host_url, + "api_key_secret_name": "provider-api-key", + "enabled_models": ["model1", "model2"], + "default_extra_body": {"temperature": 0.25}, + "default_extra_headers": {"X-Default": "default"}, + "required_extra_body": {"seed": 7}, + "required_extra_headers": {"X-Required": "required"}, + "model_deployment_id": "test-ns/test-deployment", + "status": "READY", + "status_message": None, + "auth_header_format": "X-Api-Key: {{ auth_secret }}", + } + + # served_models and auth_context are not writable through this DTO. The server + # preserves served_models and captures the authenticated caller's auth context. + assert existing_provider.served_models == [ + ServedModelMapping(model_entity_id="test-ns/model1", served_model_name="model1") + ] + reconciler._models_client.create_provider.assert_not_called() @pytest.mark.asyncio async def test_ensure_model_provider_creates_new_when_existing_not_found(reconciler, make_deployment): """Test that ensure_model_provider creates new provider when existing provider_id points to non-existent provider.""" # First retrieve (checking existing provider) fails, second retrieve (checking name collision) fails too - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() deployment = make_deployment( workspace="test-ns", project="test-project", model_provider_id="test-ns/missing-provider" @@ -505,25 +552,23 @@ async def test_ensure_model_provider_creates_new_when_existing_not_found(reconci assert model_provider_id == "test-ns/test-deployment" # Verify retrieve was called twice (once for existing, once for name collision check) - assert reconciler._models_sdk.inference.providers.retrieve.call_count == 2 + assert reconciler._models_client.get_provider.call_count == 2 - # Verify create was called to create new provider with status - reconciler._models_sdk.inference.providers.create.assert_called_once_with( - workspace="test-ns", - name="test-deployment", - host_url="http://test-ns/test-deployment", - description="Auto-created provider for deployment test-deployment", - project="test-project", - model_deployment_id="test-ns/test-deployment", - status="READY", - ) + # Verify typed create request creates a replacement provider. + reconciler._models_client.create_provider.assert_called_once() + call = reconciler._models_client.create_provider.call_args + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].name == "test-deployment" + assert call.kwargs["body"].host_url == "http://test-ns/test-deployment" + assert call.kwargs["body"].model_deployment_id == "test-ns/test-deployment" + assert call.kwargs["body"].status.value == "READY" @pytest.mark.asyncio async def test_delete_model_provider_deletes_when_exists(reconciler, make_deployment): """Test that delete_model_provider deletes provider when it exists.""" # Mock provider exists and delete succeeds - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() # Mock the cleanup method to track if it's called reconciler._cleanup_model_entities_for_provider = AsyncMock() @@ -538,7 +583,7 @@ async def test_delete_model_provider_deletes_when_exists(reconciler, make_deploy ) # Verify delete was called with correct parameters - reconciler._models_sdk.inference.providers.delete.assert_called_once_with( + reconciler._models_client.delete_provider.assert_called_once_with( name="test-deployment", workspace="test-ns", ) @@ -548,9 +593,7 @@ async def test_delete_model_provider_deletes_when_exists(reconciler, make_deploy async def test_delete_model_provider_handles_not_found(reconciler, make_deployment): """Test that delete_model_provider handles NotFoundError gracefully.""" # Mock provider doesn't exist (delete raises NotFoundError) - reconciler._models_sdk.inference.providers.delete = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) + reconciler._models_client.delete_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) # Mock the cleanup method to track if it's called reconciler._cleanup_model_entities_for_provider = AsyncMock() @@ -566,13 +609,13 @@ async def test_delete_model_provider_handles_not_found(reconciler, make_deployme ) # Verify delete was called - reconciler._models_sdk.inference.providers.delete.assert_called_once() + reconciler._models_client.delete_provider.assert_called_once() @pytest.mark.asyncio async def test_delete_model_provider_skips_when_no_provider_id(reconciler, make_deployment): """Test that delete_model_provider skips deletion when model_provider_id is not set.""" - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() # Mock the cleanup method to track if it's called reconciler._cleanup_model_entities_for_provider = AsyncMock() @@ -585,13 +628,13 @@ async def test_delete_model_provider_skips_when_no_provider_id(reconciler, make_ reconciler._cleanup_model_entities_for_provider.assert_not_called() # Verify delete was NOT called - reconciler._models_sdk.inference.providers.delete.assert_not_called() + reconciler._models_client.delete_provider.assert_not_called() @pytest.mark.asyncio async def test_delete_model_provider_handles_uuid_suffix(reconciler, make_deployment): """Test that delete_model_provider correctly parses provider ID with UUID suffix.""" - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() # Mock the cleanup method to track if it's called reconciler._cleanup_model_entities_for_provider = AsyncMock() @@ -606,7 +649,7 @@ async def test_delete_model_provider_handles_uuid_suffix(reconciler, make_deploy ) # Verify delete was called with UUID-suffixed name - reconciler._models_sdk.inference.providers.delete.assert_called_once_with( + reconciler._models_client.delete_provider.assert_called_once_with( name="test-deployment_abcdef12", workspace="test-ns", ) @@ -616,10 +659,8 @@ async def test_delete_model_provider_handles_uuid_suffix(reconciler, make_deploy async def test_reconcile_model_provider_creates_for_ready_status(reconciler, make_deployment): """Test that reconcile_model_provider creates provider when status is READY.""" # Mock provider doesn't exist - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", project="test-project") @@ -632,14 +673,14 @@ async def test_reconcile_model_provider_creates_for_ready_status(reconciler, mak assert model_provider_id == "test-ns/test-deployment" # Verify create was called - reconciler._models_sdk.inference.providers.create.assert_called_once() + reconciler._models_client.create_provider.assert_called_once() @pytest.mark.asyncio async def test_reconcile_model_provider_deletes_for_deleted_status(reconciler, make_deployment): """Test that reconcile_model_provider deletes provider when status is DELETED or DELETING.""" # Mock provider exists - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", model_provider_id="test-ns/test-deployment") @@ -650,14 +691,14 @@ async def test_reconcile_model_provider_deletes_for_deleted_status(reconciler, m assert model_provider_id is None # Verify delete was called - reconciler._models_sdk.inference.providers.delete.assert_called_once() + reconciler._models_client.delete_provider.assert_called_once() @pytest.mark.asyncio async def test_reconcile_model_provider_deletes_for_deleting_status(reconciler, make_deployment): """Test that reconcile_model_provider deletes provider when status is DELETING.""" # Mock provider exists - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", model_provider_id="test-ns/test-deployment") @@ -668,15 +709,15 @@ async def test_reconcile_model_provider_deletes_for_deleting_status(reconciler, assert model_provider_id is None # Verify delete was called - reconciler._models_sdk.inference.providers.delete.assert_called_once() + reconciler._models_client.delete_provider.assert_called_once() @pytest.mark.asyncio async def test_reconcile_model_provider_does_nothing_for_other_statuses(reconciler, make_deployment): """Test that reconcile_model_provider does nothing for statuses other than READY/DELETED/DELETING.""" - reconciler._models_sdk.inference.providers.retrieve = AsyncMock() - reconciler._models_sdk.inference.providers.create = AsyncMock() - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.get_provider = AsyncMock() + reconciler._models_client.create_provider = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() deployment = make_deployment(workspace="test-ns") @@ -696,19 +737,17 @@ async def test_reconcile_model_provider_does_nothing_for_other_statuses(reconcil assert result is None # Verify no provider operations were called - reconciler._models_sdk.inference.providers.retrieve.assert_not_called() - reconciler._models_sdk.inference.providers.create.assert_not_called() - reconciler._models_sdk.inference.providers.delete.assert_not_called() + reconciler._models_client.get_provider.assert_not_called() + reconciler._models_client.create_provider.assert_not_called() + reconciler._models_client.delete_provider.assert_not_called() @pytest.mark.asyncio async def test_reconcile_model_provider_handles_errors_gracefully(reconciler, make_deployment): """Test that reconcile_model_provider handles errors without failing deployment update.""" # Mock provider creation fails with unexpected error - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock(side_effect=Exception("API Error")) + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock(side_effect=Exception("API Error")) deployment = make_deployment(workspace="test-ns", project="test-project") @@ -722,7 +761,7 @@ async def test_reconcile_model_provider_handles_errors_gracefully(reconciler, ma assert result is None # Verify create was attempted - reconciler._models_sdk.inference.providers.create.assert_called_once() + reconciler._models_client.create_provider.assert_called_once() @pytest.mark.asyncio @@ -756,12 +795,10 @@ async def test_full_deployment_lifecycle_with_provider_management(reconciler, mo mock_backend.delete_model_deployment = AsyncMock(return_value=deleted_status) # Mock provider operations - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() - reconciler._models_sdk.inference.providers.delete = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() + reconciler._models_client.delete_provider = AsyncMock() deployment = make_deployment(workspace="test-ns", status="CREATED", project="test-project") @@ -769,7 +806,7 @@ async def test_full_deployment_lifecycle_with_provider_management(reconciler, mo await reconciler._reconcile_individual_deployment(deployment, mock_backend.create_model_deployment, "create") # Provider should NOT be created for PENDING status - reconciler._models_sdk.inference.providers.create.assert_not_called() + reconciler._models_client.create_provider.assert_not_called() # Step 2: Deployment becomes READY deployment.status = "PENDING" @@ -777,20 +814,18 @@ async def test_full_deployment_lifecycle_with_provider_management(reconciler, mo deployment, mock_backend.get_model_deployment_status, "check status" ) - # Provider SHOULD be created for READY status with backend-provided host_url and status - reconciler._models_sdk.inference.providers.create.assert_called_once_with( - workspace="test-ns", - name="test-deployment", - host_url="http://test-ns/test-deployment", # From backend's status update - description="Auto-created provider for deployment test-deployment", - project="test-project", - model_deployment_id="test-ns/test-deployment", # New field linking to deployment - status="READY", - ) + # Provider SHOULD be created for READY status with a typed request. + reconciler._models_client.create_provider.assert_called_once() + provider_call = reconciler._models_client.create_provider.call_args + assert provider_call.kwargs["workspace"] == "test-ns" + assert provider_call.kwargs["body"].name == "test-deployment" + assert provider_call.kwargs["body"].host_url == "http://test-ns/test-deployment" + assert provider_call.kwargs["body"].model_deployment_id == "test-ns/test-deployment" + assert provider_call.kwargs["body"].status.value == "READY" - # Verify the status update for READY included model_provider_id - ready_call = reconciler._models_sdk.inference.deployments.update_status.call_args_list[1] - assert ready_call.kwargs["model_provider_id"] == "test-ns/test-deployment" + # Verify the typed status request for READY included model_provider_id. + ready_call = reconciler._models_client.update_deployment_status.call_args_list[1] + assert ready_call.kwargs["body"].model_provider_id == "test-ns/test-deployment" # Step 3: Delete deployment (READY -> DELETED) # The deployment should now have the model_provider_id set from when it was READY @@ -801,13 +836,13 @@ async def test_full_deployment_lifecycle_with_provider_management(reconciler, mo ) # Provider SHOULD be deleted for DELETED status - reconciler._models_sdk.inference.providers.delete.assert_called_once_with( + reconciler._models_client.delete_provider.assert_called_once_with( name="test-deployment", workspace="test-ns", ) # Verify all deployment status updates were called - assert reconciler._models_sdk.inference.deployments.update_status.call_count == 3 + assert reconciler._models_client.update_deployment_status.call_count == 3 # ============================================================================ @@ -828,13 +863,13 @@ async def test_handle_deleted_deployment_cleanup_after_grace_period(reconciler, ) # Mock the SDK versions.delete method - reconciler._models_sdk.inference.deployments.versions.delete = AsyncMock() + reconciler._models_client.delete_deployment_version = AsyncMock() # Call handle_deleted_deployment await reconciler._handle_deleted_deployment(deployment) # Verify hard delete was called for the specific version - reconciler._models_sdk.inference.deployments.versions.delete.assert_called_once_with( + reconciler._models_client.delete_deployment_version.assert_called_once_with( name="1", workspace="test-workspace", deployment="test-deployment", @@ -854,13 +889,13 @@ async def test_handle_deleted_deployment_no_cleanup_within_grace_period(reconcil ) # Mock the SDK versions.delete method - reconciler._models_sdk.inference.deployments.versions.delete = AsyncMock() + reconciler._models_client.delete_deployment_version = AsyncMock() # Call handle_deleted_deployment await reconciler._handle_deleted_deployment(deployment) # Verify hard delete was NOT called - reconciler._models_sdk.inference.deployments.versions.delete.assert_not_called() + reconciler._models_client.delete_deployment_version.assert_not_called() @pytest.mark.asyncio @@ -879,13 +914,13 @@ async def test_handle_deleted_deployment_with_naive_datetime(reconciler, make_de ) # Mock the SDK versions.delete method - reconciler._models_sdk.inference.deployments.versions.delete = AsyncMock() + reconciler._models_client.delete_deployment_version = AsyncMock() # Call handle_deleted_deployment - should NOT raise TypeError await reconciler._handle_deleted_deployment(deployment) # Verify hard delete was called for the specific version (deployment is past grace period) - reconciler._models_sdk.inference.deployments.versions.delete.assert_called_once_with( + reconciler._models_client.delete_deployment_version.assert_called_once_with( name="1", workspace="test-workspace", deployment="test-deployment", @@ -914,13 +949,13 @@ async def test_reconcile_deployments_calls_handle_deleted(reconciler, make_deplo ) # Mock the SDK versions.delete method - reconciler._models_sdk.inference.deployments.versions.delete = AsyncMock() + reconciler._models_client.delete_deployment_version = AsyncMock() # Call reconcile_deployments with a list containing the DELETED deployment context await reconciler.reconcile_deployments([deleted_context]) # Verify hard delete was called for the specific version (since it's past grace period) - reconciler._models_sdk.inference.deployments.versions.delete.assert_called_once_with( + reconciler._models_client.delete_deployment_version.assert_called_once_with( name="1", workspace="test-workspace", deployment="deleted-deployment", @@ -942,39 +977,34 @@ async def test_cleanup_model_entities_removes_provider_from_entities(reconciler) MagicMock(model_entity_id="test-ns/model-2"), ] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) await seed_entity_cache( - reconciler._models_sdk, + reconciler._models_client, reconciler._entity_cache, [ _entity("test-ns", "model-1", ["test-ns/provider-1", "other-ns/other-provider"]), _entity("test-ns", "model-2", ["test-ns/provider-1"]), ], ) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.update_model = AsyncMock(return_value=_response(MagicMock())) # Call cleanup await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify provider was retrieved - reconciler._models_sdk.inference.providers.retrieve.assert_called_once_with( + reconciler._models_client.get_provider.assert_called_once_with( name="provider-1", workspace="test-ns", ) - # Verify model entities were updated with provider removed - assert reconciler._models_sdk.models.update.call_count == 2 - reconciler._models_sdk.models.update.assert_any_call( - name="model-1", - workspace="test-ns", - model_providers=["other-ns/other-provider"], - ) - reconciler._models_sdk.models.update.assert_any_call( - name="model-2", - workspace="test-ns", - model_providers=[], - ) + # Verify typed model update requests removed the provider. + assert reconciler._models_client.update_model.call_count == 2 + by_name = {c.kwargs["name"]: c for c in reconciler._models_client.update_model.call_args_list} + assert by_name["model-1"].kwargs["workspace"] == "test-ns" + assert by_name["model-1"].kwargs["body"].model_providers == ["other-ns/other-provider"] + assert by_name["model-2"].kwargs["workspace"] == "test-ns" + assert by_name["model-2"].kwargs["body"].model_providers == [] @pytest.mark.asyncio @@ -984,34 +1014,32 @@ async def test_cleanup_model_entities_no_served_models(reconciler): mock_provider = MagicMock() mock_provider.served_models = [] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) + reconciler._models_client.update_model = AsyncMock() # Call cleanup await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify provider was retrieved - reconciler._models_sdk.inference.providers.retrieve.assert_called_once() + reconciler._models_client.get_provider.assert_called_once() # Verify no model entity operations were performed - reconciler._models_sdk.models.update.assert_not_called() + reconciler._models_client.update_model.assert_not_called() @pytest.mark.asyncio async def test_cleanup_model_entities_provider_not_found(reconciler): """Test that cleanup handles NotFoundError gracefully when provider doesn't exist.""" - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Provider not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.update_model = AsyncMock() # Call cleanup - should not raise await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify no model entity operations were performed - reconciler._models_sdk.models.update.assert_not_called() + reconciler._models_client.update_model.assert_not_called() @pytest.mark.asyncio @@ -1023,20 +1051,20 @@ async def test_cleanup_model_entities_provider_not_in_list(reconciler): MagicMock(model_entity_id="test-ns/model-1"), ] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) await seed_entity_cache( - reconciler._models_sdk, + reconciler._models_client, reconciler._entity_cache, [_entity("test-ns", "model-1", ["other-ns/other-provider"])], ) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.update_model = AsyncMock(return_value=_response(MagicMock())) # Call cleanup await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify model entity was NOT updated (provider wasn't in the list) - reconciler._models_sdk.models.update.assert_not_called() + reconciler._models_client.update_model.assert_not_called() @pytest.mark.asyncio @@ -1049,25 +1077,25 @@ async def test_cleanup_model_entities_skips_missing_entity_and_continues(reconci MagicMock(model_entity_id="test-ns/model-2"), ] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) # Only model-2 exists. await seed_entity_cache( - reconciler._models_sdk, + reconciler._models_client, reconciler._entity_cache, [_entity("test-ns", "model-2", ["test-ns/provider-1"])], ) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.update_model = AsyncMock() # Call cleanup - should not raise await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() - # Verify only the existing model was updated - reconciler._models_sdk.models.update.assert_called_once_with( - workspace="test-ns", - name="model-2", - model_providers=[], - ) + # Verify only the existing model was updated, with a typed request. + reconciler._models_client.update_model.assert_called_once() + call = reconciler._models_client.update_model.call_args + assert call.kwargs["name"] == "model-2" + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].model_providers == [] @pytest.mark.asyncio @@ -1080,9 +1108,9 @@ async def test_cleanup_model_entities_handles_model_update_failure(reconciler): MagicMock(model_entity_id="test-ns/model-2"), ] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) await seed_entity_cache( - reconciler._models_sdk, + reconciler._models_client, reconciler._entity_cache, [ _entity("test-ns", "model-1", ["test-ns/provider-1"]), @@ -1090,14 +1118,14 @@ async def test_cleanup_model_entities_handles_model_update_failure(reconciler): ], ) # First update fails, second succeeds - reconciler._models_sdk.models.update = AsyncMock(side_effect=[Exception("Update failed"), None]) + reconciler._models_client.update_model = AsyncMock(side_effect=[Exception("Update failed"), None]) # Call cleanup - should not raise await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify both models were attempted to be updated - assert reconciler._models_sdk.models.update.call_count == 2 + assert reconciler._models_client.update_model.call_count == 2 @pytest.mark.asyncio @@ -1109,20 +1137,20 @@ async def test_cleanup_model_entities_with_null_model_providers(reconciler): MagicMock(model_entity_id="test-ns/model-1"), ] - reconciler._models_sdk.inference.providers.retrieve = AsyncMock(return_value=mock_provider) + reconciler._models_client.get_provider = AsyncMock(return_value=_response(mock_provider)) await seed_entity_cache( - reconciler._models_sdk, + reconciler._models_client, reconciler._entity_cache, [_entity("test-ns", "model-1", None)], ) - reconciler._models_sdk.models.update = AsyncMock() + reconciler._models_client.update_model = AsyncMock(return_value=_response(MagicMock())) # Call cleanup - should not raise await reconciler._cleanup_model_entities_for_provider("test-ns", "provider-1", "test-ns/provider-1") await reconciler._entity_cache.flush() # Verify model entity was NOT updated (provider wasn't in the empty/null list) - reconciler._models_sdk.models.update.assert_not_called() + reconciler._models_client.update_model.assert_not_called() @pytest.mark.asyncio @@ -1158,7 +1186,7 @@ async def test_lost_status_triggers_drift_recovery(reconciler, mock_backend_regi mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1167,11 +1195,11 @@ async def test_lost_status_triggers_drift_recovery(reconciler, mock_backend_regi mock_backend.create_model_deployment.assert_called_once_with(ctx) # Verify status was updated to PENDING with recovery message - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "PENDING" - assert "Recovering deployment" in call_kwargs["status_message"] - assert "attempt 1/" in call_kwargs["status_message"] + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "PENDING" + assert "Recovering deployment" in call_kwargs["body"].status_message + assert "attempt 1/" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1202,11 +1230,9 @@ async def test_successful_status_clears_drift_state(reconciler, mock_backend_reg mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1247,7 +1273,7 @@ async def test_pending_status_preserves_drift_state(reconciler, mock_backend_reg mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1292,7 +1318,7 @@ async def test_drift_recovery_max_retries_exceeded(reconciler, mock_backend_regi mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1301,10 +1327,10 @@ async def test_drift_recovery_max_retries_exceeded(reconciler, mock_backend_regi mock_backend.create_model_deployment.assert_not_called() # Verify status was updated to ERROR - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" - assert "failed after 3 attempts" in call_kwargs["status_message"] + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "ERROR" + assert "failed after 3 attempts" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1347,7 +1373,7 @@ async def test_drift_recovery_respects_backoff(reconciler, mock_backend_registry mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1356,7 +1382,7 @@ async def test_drift_recovery_respects_backoff(reconciler, mock_backend_registry mock_backend.create_model_deployment.assert_not_called() # Verify status was NOT updated (skipped this cycle) - reconciler._models_sdk.inference.deployments.update_status.assert_not_called() + reconciler._models_client.update_deployment_status.assert_not_called() @pytest.mark.asyncio @@ -1404,7 +1430,7 @@ async def test_drift_recovery_proceeds_after_backoff(reconciler, mock_backend_re mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1447,7 +1473,7 @@ async def test_drift_recovery_ready_deployment(reconciler, mock_backend_registry mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1456,8 +1482,8 @@ async def test_drift_recovery_ready_deployment(reconciler, mock_backend_registry mock_backend.create_model_deployment.assert_called_once() # Verify status message indicates recovery - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert "Recovering deployment" in call_kwargs["status_message"] + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert "Recovering deployment" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1483,17 +1509,17 @@ async def test_unknown_status_triggers_handler_and_updates_status(reconciler, mo mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) # Verify status was updated to UNKNOWN with attempt info - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "UNKNOWN" - assert "attempt 1/" in call_kwargs["status_message"] - assert "Unable to determine deployment status" in call_kwargs["status_message"] + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "UNKNOWN" + assert "attempt 1/" in call_kwargs["body"].status_message + assert "Unable to determine deployment status" in call_kwargs["body"].status_message # Verify attempt was tracked assert reconciler._drift_recovery_cache.get_attempts("default/test-deployment") == 1 @@ -1531,16 +1557,16 @@ async def test_unknown_status_max_retries_sets_error(reconciler, mock_backend_re mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) # Verify status was set to ERROR - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" - assert "Unable to communicate with backend after 3 attempts" in call_kwargs["status_message"] + reconciler._models_client.update_deployment_status.assert_called_once() + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "ERROR" + assert "Unable to communicate with backend after 3 attempts" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1578,13 +1604,13 @@ async def test_unknown_status_respects_backoff(reconciler, mock_backend_registry mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) # Verify status was NOT updated (in backoff period) - reconciler._models_sdk.inference.deployments.update_status.assert_not_called() + reconciler._models_client.update_deployment_status.assert_not_called() # Verify attempts was NOT incremented assert reconciler._drift_recovery_cache.get_attempts("default/test-deployment") == 1 @@ -1618,11 +1644,9 @@ async def test_unknown_status_clears_on_recovery(reconciler, mock_backend_regist mock_backend_registry.get_backend.return_value = mock_backend # Mock SDK - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() - reconciler._models_sdk.inference.providers.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.inference.providers.create = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() + reconciler._models_client.get_provider = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_provider = AsyncMock() # Process deployment await reconciler.reconcile_deployments([ctx]) @@ -1631,8 +1655,8 @@ async def test_unknown_status_clears_on_recovery(reconciler, mock_backend_regist assert "default/test-deployment" not in reconciler._drift_recovery_cache._states # Verify status was updated to READY - call_kwargs = reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kwargs["status"] == "READY" + call_kwargs = reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kwargs["body"].status.value == "READY" # ============================================================================ @@ -1727,10 +1751,10 @@ def gc_reconciler(mock_models_sdk, mock_backend_registry): """Create a reconciler with default ERROR GC TTL for GC tests.""" config = ControllerConfig() reconciler = ModelDeploymentReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_sdk, backend_registry=mock_backend_registry, controller_config=config, - entity_cache=ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None), + entity_cache=ModelEntityCache(models_client=mock_models_sdk, emit_heartbeat=lambda: None), emit_heartbeat=lambda: None, ) mock_backend = MagicMock() @@ -1738,7 +1762,7 @@ def gc_reconciler(mock_models_sdk, mock_backend_registry): return_value=DeploymentStatusUpdate(status="DELETED", status_message="") ) mock_backend_registry.get_backend.return_value = mock_backend - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() reconciler._delete_model_provider = AsyncMock() return reconciler @@ -1781,12 +1805,12 @@ async def test_gc_triggers_after_ttl(gc_reconciler, mock_backend_registry): mock_backend = mock_backend_registry.get_backend() mock_backend.delete_model_deployment.assert_called_once_with("default", "err-deploy") - gc_reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kw = gc_reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kw["status"] == "DELETING" + gc_reconciler._models_client.update_deployment_status.assert_called_once() + call_kw = gc_reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kw["body"].status.value == "DELETING" assert call_kw["name"] == "err-deploy" assert call_kw["workspace"] == "default" - assert "garbage collected" in call_kw["status_message"] + assert "garbage collected" in call_kw["body"].status_message @pytest.mark.asyncio @@ -1810,7 +1834,7 @@ async def test_gc_skips_deployment_with_no_updated_at(gc_reconciler, mock_backen mock_backend = mock_backend_registry.get_backend() mock_backend.delete_model_deployment.assert_not_called() - gc_reconciler._models_sdk.inference.deployments.update_status.assert_not_called() + gc_reconciler._models_client.update_deployment_status.assert_not_called() @pytest.mark.asyncio @@ -1823,9 +1847,9 @@ async def test_gc_backend_delete_failure_still_transitions(gc_reconciler, mock_b await gc_reconciler.gc_error_deployments([dep]) - gc_reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kw = gc_reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kw["status"] == "DELETING" + gc_reconciler._models_client.update_deployment_status.assert_called_once() + call_kw = gc_reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kw["body"].status.value == "DELETING" @pytest.mark.asyncio @@ -1834,9 +1858,7 @@ async def test_gc_status_update_failure_does_not_block_others(gc_reconciler, moc dep1 = _make_error_deployment(name="dep-1") dep2 = _make_error_deployment(name="dep-2") - gc_reconciler._models_sdk.inference.deployments.update_status = AsyncMock( - side_effect=[Exception("version conflict"), None] - ) + gc_reconciler._models_client.update_deployment_status = AsyncMock(side_effect=[Exception("version conflict"), None]) await gc_reconciler.gc_error_deployments([dep1, dep2]) @@ -1860,7 +1882,7 @@ async def test_gc_mixed_ttl_only_expired_cleaned(gc_reconciler, mock_backend_reg mock_backend = mock_backend_registry.get_backend() mock_backend.delete_model_deployment.assert_called_once_with("default", "old-deploy") - gc_reconciler._models_sdk.inference.deployments.update_status.assert_called_once() + gc_reconciler._models_client.update_deployment_status.assert_called_once() @pytest.mark.asyncio @@ -1884,10 +1906,10 @@ async def test_gc_provider_cleanup_failure_is_non_fatal(gc_reconciler, mock_back mock_backend = mock_backend_registry.get_backend() mock_backend.delete_model_deployment.assert_called_once() - gc_reconciler._models_sdk.inference.deployments.update_status.assert_called_once() - call_kw = gc_reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert call_kw["status"] == "DELETING" - assert "Provider cleanup failed" in call_kw["status_message"] + gc_reconciler._models_client.update_deployment_status.assert_called_once() + call_kw = gc_reconciler._models_client.update_deployment_status.call_args.kwargs + assert call_kw["body"].status.value == "DELETING" + assert "Provider cleanup failed" in call_kw["body"].status_message @pytest.mark.asyncio @@ -1897,7 +1919,7 @@ async def test_gc_empty_list_no_ops(gc_reconciler, mock_backend_registry): mock_backend = mock_backend_registry.get_backend() mock_backend.delete_model_deployment.assert_not_called() - gc_reconciler._models_sdk.inference.deployments.update_status.assert_not_called() + gc_reconciler._models_client.update_deployment_status.assert_not_called() @pytest.mark.asyncio @@ -1906,10 +1928,10 @@ async def test_gc_custom_ttl_respected(mock_models_sdk, mock_backend_registry): custom_ttl = 3600 # 1 hour config = ControllerConfig(error_deployment_ttl_seconds=custom_ttl) reconciler = ModelDeploymentReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_sdk, backend_registry=mock_backend_registry, controller_config=config, - entity_cache=ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None), + entity_cache=ModelEntityCache(models_client=mock_models_sdk, emit_heartbeat=lambda: None), emit_heartbeat=lambda: None, ) @@ -1918,7 +1940,7 @@ async def test_gc_custom_ttl_respected(mock_models_sdk, mock_backend_registry): return_value=DeploymentStatusUpdate(status="DELETED", status_message="") ) mock_backend_registry.get_backend.return_value = mock_backend - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() reconciler._delete_model_provider = AsyncMock() within_default_but_past_custom = _make_error_deployment( @@ -1937,10 +1959,10 @@ async def test_gc_status_message_includes_original_error(gc_reconciler, mock_bac await gc_reconciler.gc_error_deployments([dep]) - call_kw = gc_reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert "garbage collected" in call_kw["status_message"] - assert "NIM health check timed out after 7200s" in call_kw["status_message"] - assert "Original error:" in call_kw["status_message"] + call_kw = gc_reconciler._models_client.update_deployment_status.call_args.kwargs + assert "garbage collected" in call_kw["body"].status_message + assert "NIM health check timed out after 7200s" in call_kw["body"].status_message + assert "Original error:" in call_kw["body"].status_message @pytest.mark.asyncio @@ -1950,18 +1972,16 @@ async def test_gc_status_message_without_original_error(gc_reconciler, mock_back await gc_reconciler.gc_error_deployments([dep]) - call_kw = gc_reconciler._models_sdk.inference.deployments.update_status.call_args.kwargs - assert "garbage collected" in call_kw["status_message"] - assert "Original error:" not in call_kw["status_message"] + call_kw = gc_reconciler._models_client.update_deployment_status.call_args.kwargs + assert "garbage collected" in call_kw["body"].status_message + assert "Original error:" not in call_kw["body"].status_message @pytest.mark.asyncio async def test_gc_not_found_on_status_update_handled(gc_reconciler, mock_backend_registry): """NotFoundError on status update (deployment deleted between query and GC) is handled.""" dep = _make_error_deployment() - gc_reconciler._models_sdk.inference.deployments.update_status = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) + gc_reconciler._models_client.update_deployment_status = AsyncMock(side_effect=_client_error(NotFoundError, 404)) # Should not raise await gc_reconciler.gc_error_deployments([dep]) @@ -1995,10 +2015,10 @@ async def test_gc_ttl_boundary_parametrized(mock_models_sdk, mock_backend_regist now = datetime.now(timezone.utc) config = ControllerConfig(error_deployment_ttl_seconds=ttl) reconciler = ModelDeploymentReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_sdk, backend_registry=mock_backend_registry, controller_config=config, - entity_cache=ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None), + entity_cache=ModelEntityCache(models_client=mock_models_sdk, emit_heartbeat=lambda: None), emit_heartbeat=lambda: None, ) @@ -2007,7 +2027,7 @@ async def test_gc_ttl_boundary_parametrized(mock_models_sdk, mock_backend_regist return_value=DeploymentStatusUpdate(status="DELETED", status_message="") ) mock_backend_registry.get_backend.return_value = mock_backend - reconciler._models_sdk.inference.deployments.update_status = AsyncMock() + reconciler._models_client.update_deployment_status = AsyncMock() reconciler._delete_model_provider = AsyncMock() dep = _make_error_deployment( @@ -2020,7 +2040,7 @@ async def test_gc_ttl_boundary_parametrized(mock_models_sdk, mock_backend_regist if should_gc: mock_backend.delete_model_deployment.assert_called_once() - reconciler._models_sdk.inference.deployments.update_status.assert_called_once() + reconciler._models_client.update_deployment_status.assert_called_once() else: mock_backend.delete_model_deployment.assert_not_called() - reconciler._models_sdk.inference.deployments.update_status.assert_not_called() + reconciler._models_client.update_deployment_status.assert_not_called() diff --git a/services/core/models/tests/unit/controllers/test_entity_cache.py b/services/core/models/tests/unit/controllers/test_entity_cache.py index 343b054181..fceba0bab3 100644 --- a/services/core/models/tests/unit/controllers/test_entity_cache.py +++ b/services/core/models/tests/unit/controllers/test_entity_cache.py @@ -3,28 +3,51 @@ """Unit tests for ModelEntityCache.""" -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock +import httpx import pytest -from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import ConflictError, NotFoundError +from nemo_platform_plugin.client.errors import ConflictError, NemoHTTPError, NotFoundError from nmp.core.models.controllers.entity_cache import ModelEntityCache, UnflushedMutationsError -from .conftest import AsyncPaginator, make_entity, seed_entity_cache +from .conftest import PaginatedResponse, make_entity, make_models_client, response, seed_entity_cache def _entity(workspace="ws", name="model", model_providers=None, **attrs): return make_entity(workspace, name, model_providers=model_providers, **attrs) +def _client_error(error_type: type[NemoHTTPError], status_code: int) -> NemoHTTPError: + return error_type(httpx.Response(status_code, request=httpx.Request("GET", "http://test"))) + + +def _assert_updated(mock_models_client, *, workspace, name, **fields): + """Assert exactly one update carrying exactly ``fields`` in the request body. + + ``exclude_unset`` keeps this honest: a field the cache writes but the test does + not name fails here rather than passing silently. + """ + mock_models_client.update_model.assert_awaited_once() + call = mock_models_client.update_model.call_args + assert call.kwargs["workspace"] == workspace + assert call.kwargs["name"] == name + assert call.kwargs["body"].model_dump(mode="json", exclude_unset=True) == fields + + +def _assert_created(mock_models_client, *, workspace, **fields): + """Assert exactly one create carrying exactly ``fields`` in the request body.""" + mock_models_client.create_model.assert_awaited_once() + call = mock_models_client.create_model.call_args + assert call.kwargs["workspace"] == workspace + assert call.kwargs["body"].model_dump(mode="json", exclude_unset=True) == fields + + @pytest.fixture -def mock_models_sdk(): - sdk = MagicMock(spec=AsyncNeMoPlatform) - sdk.models.list = MagicMock(return_value=AsyncPaginator([])) - sdk.models.create = AsyncMock(return_value=None) - sdk.models.update = AsyncMock(return_value=None) - sdk.models.retrieve = AsyncMock() - return sdk +def mock_models_client(): + client = make_models_client() + client.create_model = AsyncMock(return_value=response(None)) + client.update_model = AsyncMock(return_value=response(None)) + return client @pytest.fixture @@ -34,17 +57,28 @@ def heartbeat_calls(): @pytest.fixture -def cache(mock_models_sdk, heartbeat_calls): - return ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: heartbeat_calls.append(1)) +def cache(mock_models_client, heartbeat_calls): + return ModelEntityCache(models_client=mock_models_client, emit_heartbeat=lambda: heartbeat_calls.append(1)) + +async def _load(mock_models_client, cache, entities=()): + await seed_entity_cache(mock_models_client, cache, entities) + + +@pytest.mark.asyncio +async def test_refresh_reads_every_workspace_in_one_paginated_call(mock_models_client, cache): + """The snapshot is one cross-workspace read, not a call per workspace.""" + await _load(mock_models_client, cache, [_entity("ws-a", "m1")]) -async def _load(mock_models_sdk, cache, entities=()): - await seed_entity_cache(mock_models_sdk, cache, entities) + mock_models_client.list_models.assert_awaited_once() + call = mock_models_client.list_models.call_args + assert call.kwargs["workspace"] == "-" + assert call.kwargs["query_params"]["page_size"] > 1 @pytest.mark.asyncio -async def test_refresh_loads_entities_keyed_by_workspace_and_name(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws-a", "m1"), _entity("ws-b", "m1")]) +async def test_refresh_loads_entities_keyed_by_workspace_and_name(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws-a", "m1"), _entity("ws-b", "m1")]) assert cache.loaded assert cache.get("ws-a", "m1") is not None @@ -53,8 +87,8 @@ async def test_refresh_loads_entities_keyed_by_workspace_and_name(mock_models_sd @pytest.mark.asyncio -async def test_refresh_rejects_unflushed_mutations(mock_models_sdk, cache): - await _load(mock_models_sdk, cache) +async def test_refresh_rejects_unflushed_mutations(mock_models_client, cache): + await _load(mock_models_client, cache) cache.stage_provider_link("ws", "model", "ws/p1") with pytest.raises(UnflushedMutationsError): @@ -62,9 +96,9 @@ async def test_refresh_rejects_unflushed_mutations(mock_models_sdk, cache): @pytest.mark.asyncio -async def test_get_reflects_staged_provider_link_within_a_phase(mock_models_sdk, cache): +async def test_get_reflects_staged_provider_link_within_a_phase(mock_models_client, cache): """A read after a stage in the same phase must observe the staged change.""" - await _load(mock_models_sdk, cache, [_entity("ws", "model", ["ws/p1"])]) + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1"])]) cache.stage_provider_link("ws", "model", "ws/p2") @@ -72,8 +106,8 @@ async def test_get_reflects_staged_provider_link_within_a_phase(mock_models_sdk, @pytest.mark.asyncio -async def test_get_reflects_staged_provider_unlink_within_a_phase(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws", "model", ["ws/p1", "ws/p2"])]) +async def test_get_reflects_staged_provider_unlink_within_a_phase(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1", "ws/p2"])]) cache.stage_provider_unlink("ws", "model", "ws/p1") @@ -81,9 +115,9 @@ async def test_get_reflects_staged_provider_unlink_within_a_phase(mock_models_sd @pytest.mark.asyncio -async def test_entity_staged_for_creation_still_reads_as_absent(mock_models_sdk, cache): +async def test_entity_staged_for_creation_still_reads_as_absent(mock_models_client, cache): """A staged creation is not fabricated, so callers keep treating it as new.""" - await _load(mock_models_sdk, cache) + await _load(mock_models_client, cache) cache.stage_create("ws", "new-model", description="d", backend_format="OPENAI_CHAT") cache.stage_provider_link("ws", "new-model", "ws/p1") @@ -92,34 +126,32 @@ async def test_entity_staged_for_creation_still_reads_as_absent(mock_models_sdk, @pytest.mark.asyncio -async def test_multiple_providers_produce_a_single_update(mock_models_sdk, cache): +async def test_multiple_providers_produce_a_single_update(mock_models_client, cache): """An entity linked by several providers is written once, not once per provider.""" - await _load(mock_models_sdk, cache, [_entity("ws", "model", [])]) + await _load(mock_models_client, cache, [_entity("ws", "model", [])]) cache.stage_provider_link("ws", "model", "ws/p1") cache.stage_provider_link("ws", "model", "ws/p2") await cache.flush() - mock_models_sdk.models.update.assert_awaited_once_with( - workspace="ws", name="model", model_providers=["ws/p1", "ws/p2"] - ) + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/p1", "ws/p2"]) @pytest.mark.asyncio -async def test_no_write_when_already_converged(mock_models_sdk, cache): +async def test_no_write_when_already_converged(mock_models_client, cache): """Staging state that already matches the entity performs no write.""" - await _load(mock_models_sdk, cache, [_entity("ws", "model", ["ws/p1"])]) + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1"])]) cache.stage_provider_link("ws", "model", "ws/p1") await cache.flush() - mock_models_sdk.models.update.assert_not_awaited() - mock_models_sdk.models.create.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() + mock_models_client.create_model.assert_not_awaited() @pytest.mark.asyncio -async def test_two_providers_creating_the_same_entity_collapse_to_one_create(mock_models_sdk, cache): - await _load(mock_models_sdk, cache) +async def test_two_providers_creating_the_same_entity_collapse_to_one_create(mock_models_client, cache): + await _load(mock_models_client, cache) cache.stage_create("ws", "model", description="from p1", backend_format="OPENAI_CHAT") cache.stage_provider_link("ws", "model", "ws/p1") @@ -127,7 +159,8 @@ async def test_two_providers_creating_the_same_entity_collapse_to_one_create(moc cache.stage_provider_link("ws", "model", "ws/p2") await cache.flush() - mock_models_sdk.models.create.assert_awaited_once_with( + _assert_created( + mock_models_client, workspace="ws", name="model", description="from p1", @@ -137,55 +170,53 @@ async def test_two_providers_creating_the_same_entity_collapse_to_one_create(moc @pytest.mark.asyncio -async def test_create_conflict_falls_back_to_updating_the_existing_entity(mock_models_sdk, cache): +async def test_create_conflict_falls_back_to_updating_the_existing_entity(mock_models_client, cache): """An entity created concurrently is adopted rather than reported as an error.""" - await _load(mock_models_sdk, cache) - mock_models_sdk.models.create = AsyncMock(side_effect=ConflictError("exists", response=MagicMock(), body=None)) - mock_models_sdk.models.retrieve = AsyncMock(return_value=_entity("ws", "model", ["ws/other"])) + await _load(mock_models_client, cache) + mock_models_client.create_model = AsyncMock(side_effect=_client_error(ConflictError, 409)) + mock_models_client.get_model = AsyncMock(return_value=response(_entity("ws", "model", ["ws/other"]))) cache.stage_create("ws", "model", description="d", backend_format="OPENAI_CHAT") cache.stage_provider_link("ws", "model", "ws/p1") await cache.flush() - mock_models_sdk.models.update.assert_awaited_once_with( - workspace="ws", name="model", model_providers=["ws/other", "ws/p1"] - ) + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/other", "ws/p1"]) @pytest.mark.asyncio -async def test_create_conflict_with_vanished_entity_is_ignored(mock_models_sdk, cache): - await _load(mock_models_sdk, cache) - mock_models_sdk.models.create = AsyncMock(side_effect=ConflictError("exists", response=MagicMock(), body=None)) - mock_models_sdk.models.retrieve = AsyncMock(side_effect=NotFoundError("gone", response=MagicMock(), body=None)) +async def test_create_conflict_with_vanished_entity_is_ignored(mock_models_client, cache): + await _load(mock_models_client, cache) + mock_models_client.create_model = AsyncMock(side_effect=_client_error(ConflictError, 409)) + mock_models_client.get_model = AsyncMock(side_effect=_client_error(NotFoundError, 404)) cache.stage_create("ws", "model", description="d") await cache.flush() - mock_models_sdk.models.update.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() @pytest.mark.asyncio -async def test_one_failing_entity_does_not_stop_the_others(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws", "m1", []), _entity("ws", "m2", [])]) - mock_models_sdk.models.update = AsyncMock(side_effect=[Exception("boom"), None]) +async def test_one_failing_entity_does_not_stop_the_others(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "m1", []), _entity("ws", "m2", [])]) + mock_models_client.update_model = AsyncMock(side_effect=[Exception("boom"), response(None)]) cache.stage_provider_link("ws", "m1", "ws/p1") cache.stage_provider_link("ws", "m2", "ws/p1") await cache.flush() - assert mock_models_sdk.models.update.await_count == 2 + assert mock_models_client.update_model.await_count == 2 @pytest.mark.asyncio -async def test_failed_write_is_kept_for_retry_and_succeeds_later(mock_models_sdk, cache): +async def test_failed_write_is_kept_for_retry_and_succeeds_later(mock_models_client, cache): """A write that fails must not be lost. Some staged changes cannot be recomputed by a later pass -- unlinking a provider that is being deleted is derived from that provider -- so a dropped failure would leave the entity permanently inconsistent. """ - await _load(mock_models_sdk, cache, [_entity("ws", "m1", ["ws/p1"]), _entity("ws", "m2", ["ws/p1"])]) - mock_models_sdk.models.update = AsyncMock(side_effect=[Exception("boom"), None]) + await _load(mock_models_client, cache, [_entity("ws", "m1", ["ws/p1"]), _entity("ws", "m2", ["ws/p1"])]) + mock_models_client.update_model = AsyncMock(side_effect=[Exception("boom"), response(None)]) cache.stage_provider_unlink("ws", "m1", "ws/p1") cache.stage_provider_unlink("ws", "m2", "ws/p1") @@ -197,25 +228,25 @@ async def test_failed_write_is_kept_for_retry_and_succeeds_later(mock_models_sdk assert ("ws", "m2") not in cache._pending # A later flush retries it, and this time it lands. - mock_models_sdk.models.update = AsyncMock(return_value=None) + mock_models_client.update_model = AsyncMock(return_value=response(None)) await cache.flush() - mock_models_sdk.models.update.assert_awaited_once_with(workspace="ws", name="m1", model_providers=[]) + _assert_updated(mock_models_client, workspace="ws", name="m1", model_providers=[]) assert cache._pending == {} @pytest.mark.asyncio -async def test_refresh_allows_retained_failures_but_still_rejects_unflushed_work(mock_models_sdk, cache): +async def test_refresh_allows_retained_failures_but_still_rejects_unflushed_work(mock_models_client, cache): """Refresh distinguishes "flushed and failed" from "staged and forgotten".""" - await _load(mock_models_sdk, cache, [_entity("ws", "m1", ["ws/p1"])]) - mock_models_sdk.models.update = AsyncMock(side_effect=Exception("boom")) + await _load(mock_models_client, cache, [_entity("ws", "m1", ["ws/p1"])]) + mock_models_client.update_model = AsyncMock(side_effect=Exception("boom")) cache.stage_provider_unlink("ws", "m1", "ws/p1") await cache.flush() assert ("ws", "m1") in cache._pending # A retained failure does not block the next phase from re-reading. - await _load(mock_models_sdk, cache, [_entity("ws", "m1", ["ws/p1"])]) + await _load(mock_models_client, cache, [_entity("ws", "m1", ["ws/p1"])]) # Work that no flush has attempted still does. cache.stage_provider_link("ws", "m2", "ws/p2") @@ -224,73 +255,72 @@ async def test_refresh_allows_retained_failures_but_still_rejects_unflushed_work @pytest.mark.asyncio -async def test_retained_failure_replays_against_a_newer_snapshot(mock_models_sdk, cache): +async def test_retained_failure_replays_against_a_newer_snapshot(mock_models_client, cache): """Staged changes are differences, so replaying them after a refresh stays correct.""" - await _load(mock_models_sdk, cache, [_entity("ws", "m1", ["ws/p1", "ws/p2"])]) - mock_models_sdk.models.update = AsyncMock(side_effect=Exception("boom")) + await _load(mock_models_client, cache, [_entity("ws", "m1", ["ws/p1", "ws/p2"])]) + mock_models_client.update_model = AsyncMock(side_effect=Exception("boom")) cache.stage_provider_unlink("ws", "m1", "ws/p1") await cache.flush() # Snapshot moves on: another writer added a third provider meanwhile. - await _load(mock_models_sdk, cache, [_entity("ws", "m1", ["ws/p1", "ws/p2", "ws/p3"])]) - mock_models_sdk.models.update = AsyncMock(return_value=None) + await _load(mock_models_client, cache, [_entity("ws", "m1", ["ws/p1", "ws/p2", "ws/p3"])]) + mock_models_client.update_model = AsyncMock(return_value=response(None)) await cache.flush() # The unlink applies to the newer state rather than reinstating the old list. - mock_models_sdk.models.update.assert_awaited_once_with( - workspace="ws", name="m1", model_providers=["ws/p2", "ws/p3"] - ) + _assert_updated(mock_models_client, workspace="ws", name="m1", model_providers=["ws/p2", "ws/p3"]) @pytest.mark.asyncio -async def test_flush_clears_staged_changes(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws", "model", [])]) +async def test_flush_clears_staged_changes(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "model", [])]) cache.stage_provider_link("ws", "model", "ws/p1") await cache.flush() - mock_models_sdk.models.update.reset_mock() + mock_models_client.update_model.reset_mock() # Nothing left staged, so a second flush writes nothing and a refresh is allowed. await cache.flush() - mock_models_sdk.models.update.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() await cache.refresh() @pytest.mark.asyncio -async def test_link_then_unlink_for_the_same_provider_cancels_out(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws", "model", ["ws/p1"])]) +async def test_link_then_unlink_for_the_same_provider_cancels_out(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1"])]) cache.stage_provider_unlink("ws", "model", "ws/p1") cache.stage_provider_link("ws", "model", "ws/p1") await cache.flush() - mock_models_sdk.models.update.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() @pytest.mark.asyncio -async def test_field_updates_are_written_as_staged(mock_models_sdk, cache): - await _load(mock_models_sdk, cache, [_entity("ws", "model", ["ws/p1"])]) +async def test_field_updates_are_written_as_staged(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1"])]) cache.stage_field_updates("ws", "model", fileset="hub/model", api_endpoint=None) await cache.flush() - mock_models_sdk.models.update.assert_awaited_once_with(workspace="ws", name="model", fileset="hub/model") + # api_endpoint was None, so it is dropped rather than written as a null. + _assert_updated(mock_models_client, workspace="ws", name="model", fileset="hub/model") @pytest.mark.asyncio -async def test_staged_change_for_missing_entity_without_create_is_skipped(mock_models_sdk, cache): - await _load(mock_models_sdk, cache) +async def test_staged_change_for_missing_entity_without_create_is_skipped(mock_models_client, cache): + await _load(mock_models_client, cache) cache.stage_provider_unlink("ws", "ghost", "ws/p1") await cache.flush() - mock_models_sdk.models.create.assert_not_awaited() - mock_models_sdk.models.update.assert_not_awaited() + mock_models_client.create_model.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() @pytest.mark.asyncio -async def test_refresh_after_flush_does_not_reapply_earlier_state(mock_models_sdk, cache): +async def test_refresh_after_flush_does_not_reapply_earlier_state(mock_models_client, cache): """A link removed in one phase is not reinstated by the next phase. The second phase must decide from a snapshot that already reflects the first @@ -298,16 +328,17 @@ async def test_refresh_after_flush_does_not_reapply_earlier_state(mock_models_sd """ store = {("ws", "model"): _entity("ws", "model", ["ws/p1"])} - def _list(**_kwargs): - return AsyncPaginator(list(store.values())) + async def _list(**_kwargs): + return PaginatedResponse(list(store.values())) - async def _update(*, workspace, name, **params): + async def _update(*, workspace, name, body): current = store[(workspace, name)] - store[(workspace, name)] = _entity(workspace, name, params.get("model_providers", current.model_providers)) - return store[(workspace, name)] + providers = body.model_providers if "model_providers" in body.model_fields_set else current.model_providers + store[(workspace, name)] = _entity(workspace, name, providers) + return response(store[(workspace, name)]) - mock_models_sdk.models.list = MagicMock(side_effect=_list) - mock_models_sdk.models.update = AsyncMock(side_effect=_update) + mock_models_client.list_models = AsyncMock(side_effect=_list) + mock_models_client.update_model = AsyncMock(side_effect=_update) # Phase one removes the provider link and applies it. await cache.refresh() @@ -321,36 +352,36 @@ async def _update(*, workspace, name, **params): @pytest.mark.asyncio -async def test_refresh_reports_progress_per_entity_read(mock_models_sdk, cache, heartbeat_calls): +async def test_refresh_reports_progress_per_entity_read(mock_models_client, cache, heartbeat_calls): """Reading a large batch has to report progress as it goes.""" - await _load(mock_models_sdk, cache, [_entity("ws", f"m{i}") for i in range(25)]) + await _load(mock_models_client, cache, [_entity("ws", f"m{i}") for i in range(25)]) assert len(heartbeat_calls) == 25 @pytest.mark.asyncio -async def test_flush_reports_progress_per_entity_written(mock_models_sdk, cache, heartbeat_calls): +async def test_flush_reports_progress_per_entity_written(mock_models_client, cache, heartbeat_calls): """Writing a large batch has to report progress as it goes. Writes are the slowest part of a pass, so a flush that reported nothing would make a long but advancing pass indistinguishable from a stalled one. """ - await _load(mock_models_sdk, cache, [_entity("ws", f"m{i}", []) for i in range(25)]) + await _load(mock_models_client, cache, [_entity("ws", f"m{i}", []) for i in range(25)]) heartbeat_calls.clear() for i in range(25): cache.stage_provider_link("ws", f"m{i}", "ws/p1") await cache.flush() - assert mock_models_sdk.models.update.await_count == 25 + assert mock_models_client.update_model.await_count == 25 assert len(heartbeat_calls) == 25 @pytest.mark.asyncio -async def test_flush_reports_progress_even_when_an_entity_write_fails(mock_models_sdk, cache, heartbeat_calls): +async def test_flush_reports_progress_even_when_an_entity_write_fails(mock_models_client, cache, heartbeat_calls): """Moving past a failed entity is still progress.""" - await _load(mock_models_sdk, cache, [_entity("ws", "m1", []), _entity("ws", "m2", [])]) - mock_models_sdk.models.update = AsyncMock(side_effect=[Exception("boom"), None]) + await _load(mock_models_client, cache, [_entity("ws", "m1", []), _entity("ws", "m2", [])]) + mock_models_client.update_model = AsyncMock(side_effect=[Exception("boom"), response(None)]) heartbeat_calls.clear() cache.stage_provider_link("ws", "m1", "ws/p1") @@ -360,18 +391,160 @@ async def test_flush_reports_progress_even_when_an_entity_write_fails(mock_model assert len(heartbeat_calls) == 2 +# --------------------------------------------------------------------------- +# The typed client returns responses, not entities. Dropping .data() would put a +# NemoResponse into the snapshot, and dropping the write-back would make a second +# change to the same entity in the same phase take a create/409 round trip. +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_created_entity_is_unwrapped_into_the_snapshot(mock_models_client, cache): + """A create seeds the snapshot with the entity, so the phase can build on it.""" + await _load(mock_models_client, cache) + server_entity = _entity("ws", "model", ["ws/p1"]) + mock_models_client.create_model = AsyncMock(return_value=response(server_entity)) + + cache.stage_create("ws", "model", description="d") + cache.stage_provider_link("ws", "model", "ws/p1") + await cache.flush() + + # The server's entity is in the snapshot, not the response wrapping it. + assert cache.get("ws", "model") is server_entity + + # A second link in the same phase is therefore an update, not another create. + cache.stage_provider_link("ws", "model", "ws/p2") + await cache.flush() + + mock_models_client.create_model.assert_awaited_once() + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/p1", "ws/p2"]) + + +@pytest.mark.asyncio +async def test_updated_entity_is_unwrapped_into_the_snapshot(mock_models_client, cache): + """An update replaces the snapshot entry with what the server returned.""" + await _load(mock_models_client, cache, [_entity("ws", "model", [])]) + server_entity = _entity("ws", "model", ["ws/p1"], fileset="hub/model") + mock_models_client.update_model = AsyncMock(return_value=response(server_entity)) + + cache.stage_provider_link("ws", "model", "ws/p1") + await cache.flush() + + assert cache.get("ws", "model") is server_entity + # Reading back a server-side value the request never carried proves the + # snapshot came from the response rather than from the pre-write entity. + assert cache.get("ws", "model").fileset == "hub/model" + + +@pytest.mark.asyncio +async def test_create_body_carries_staged_field_updates(mock_models_client, cache): + """stage_create + stage_field_updates is the only create shape in production. + + provider_reconciler stages the description and backend_format, then stages + fileset and api_endpoint separately. Both have to reach the same POST body. + """ + await _load(mock_models_client, cache) + + cache.stage_create("ws", "model", description="auto-discovered", backend_format="OPENAI_CHAT") + cache.stage_provider_link("ws", "model", "ws/p1") + cache.stage_field_updates("ws", "model", fileset="hub/model", api_endpoint=None) + await cache.flush() + + _assert_created( + mock_models_client, + workspace="ws", + name="model", + description="auto-discovered", + backend_format="OPENAI_CHAT", + fileset="hub/model", + model_providers=["ws/p1"], + ) + + +@pytest.mark.asyncio +async def test_field_updates_win_over_create_attributes(mock_models_client, cache): + """A staged field update overrides the same attribute staged for creation.""" + await _load(mock_models_client, cache) + + cache.stage_create("ws", "model", backend_format="OPENAI_CHAT") + cache.stage_field_updates("ws", "model", backend_format="ANTHROPIC_MESSAGES") + await cache.flush() + + _assert_created(mock_models_client, workspace="ws", name="model", backend_format="ANTHROPIC_MESSAGES") + + +@pytest.mark.asyncio +async def test_create_addresses_the_entity_by_its_cache_key(mock_models_client, cache): + """The name and workspace on the wire come from the cache key. + + ``stage_create`` takes both positionally, so an attribute of either name + cannot reach ``create_kwargs`` and the filter that drops them in ``_create`` + is unreachable. This pins the property that filter was defending. + """ + await _load(mock_models_client, cache) + + cache.stage_create("ws", "real", description="d") + await cache.flush() + + _assert_created(mock_models_client, workspace="ws", name="real", description="d") + + +# --------------------------------------------------------------------------- +# Staging is idempotent and order-independent +# --------------------------------------------------------------------------- + + @pytest.mark.asyncio -async def test_conflict_adoption_does_not_overwrite_the_existing_entity_attributes(mock_models_sdk, cache): +async def test_repeated_staging_of_the_same_provider_does_not_duplicate_it(mock_models_client, cache): + """Several providers converging on one entity must not produce duplicate links.""" + await _load(mock_models_client, cache, [_entity("ws", "model", [])]) + + cache.stage_provider_link("ws", "model", "ws/p1") + cache.stage_provider_link("ws", "model", "ws/p1") + await cache.flush() + + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/p1"]) + + +@pytest.mark.asyncio +async def test_repeated_unlink_of_the_same_provider_is_staged_once(mock_models_client, cache): + await _load(mock_models_client, cache, [_entity("ws", "model", ["ws/p1", "ws/p2"])]) + + cache.stage_provider_unlink("ws", "model", "ws/p1") + cache.stage_provider_unlink("ws", "model", "ws/p1") + await cache.flush() + + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/p2"]) + + +@pytest.mark.asyncio +async def test_unlink_after_link_for_the_same_provider_cancels_out(mock_models_client, cache): + """The mirror of the link-after-unlink case. + + Reachable when a provider is discovered and then deleted within one phase. + """ + await _load(mock_models_client, cache, [_entity("ws", "model", [])]) + + cache.stage_provider_link("ws", "model", "ws/p1") + cache.stage_provider_unlink("ws", "model", "ws/p1") + await cache.flush() + + assert cache.get("ws", "model").model_providers == [] + mock_models_client.update_model.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_conflict_adoption_does_not_overwrite_the_existing_entity_attributes(mock_models_client, cache): """Adopting a concurrently-created entity leaves its own attributes alone. Attributes supplied for creation describe an entity we expected to create. When another writer got there first, theirs win; the owning reconciler re-evaluates what is still missing on a later pass. """ - await _load(mock_models_sdk, cache) - mock_models_sdk.models.create = AsyncMock(side_effect=ConflictError("exists", response=MagicMock(), body=None)) - mock_models_sdk.models.retrieve = AsyncMock( - return_value=_entity("ws", "model", ["ws/other"], backend_format="ANTHROPIC_MESSAGES") + await _load(mock_models_client, cache) + mock_models_client.create_model = AsyncMock(side_effect=_client_error(ConflictError, 409)) + mock_models_client.get_model = AsyncMock( + return_value=response(_entity("ws", "model", ["ws/other"], backend_format="ANTHROPIC_MESSAGES")) ) cache.stage_create("ws", "model", description="ours", backend_format="OPENAI_CHAT") @@ -379,6 +552,4 @@ async def test_conflict_adoption_does_not_overwrite_the_existing_entity_attribut await cache.flush() # Only the provider link is written; description/backend_format are not forced. - mock_models_sdk.models.update.assert_awaited_once_with( - workspace="ws", name="model", model_providers=["ws/other", "ws/p1"] - ) + _assert_updated(mock_models_client, workspace="ws", name="model", model_providers=["ws/other", "ws/p1"]) diff --git a/services/core/models/tests/unit/controllers/test_models_controller_unit.py b/services/core/models/tests/unit/controllers/test_models_controller_unit.py index 1a233fdc5c..86f914eaf0 100644 --- a/services/core/models/tests/unit/controllers/test_models_controller_unit.py +++ b/services/core/models/tests/unit/controllers/test_models_controller_unit.py @@ -4,28 +4,53 @@ """Unit tests for ModelsController.""" import asyncio +import json import threading from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from nemo_platform_plugin.client.errors import NotFoundError from nmp.core.models.config import config as models_config from nmp.core.models.controllers.context import ModelContext from nmp.core.models.controllers.models_controller import NON_TERMINAL_STATES, ModelsController +def _requested_status(kwargs): + return json.loads(kwargs["query_params"]["filter"])["status"] + + +def _close_mocked_step_coroutine(mock_run_until_complete: MagicMock) -> None: + """Close the coroutine handed to a mocked event loop.""" + mock_run_until_complete.call_args.args[0].close() + + +class MockResponse: + """Minimal typed response shape used by controller tests.""" + + def __init__(self, data): + self._data = data + + def data(self): + return self._data + + class MockAsyncPaginator: - """Mock async paginator to simulate SDK's paginated response.""" + """Minimal async paginated response shape used by controller tests.""" def __init__(self, items): - self.items = items + self._items = list(items) + + def items(self): + return self def __aiter__(self): return self async def __anext__(self): - if not self.items: + if not self._items: raise StopAsyncIteration - return self.items.pop(0) + return self._items.pop(0) def test_controller_initialization(mock_sdk_class_patch, mock_get_config_patch, mock_backend_registry, assert_helpers): @@ -50,6 +75,7 @@ def test_step_with_no_deployments( # Run step controller.step() + _close_mocked_step_coroutine(mock_asyncio_run_patch) # Verify state assert_helpers.assert_controller_healthy(controller) @@ -73,6 +99,7 @@ def test_step_with_deployments( # Run step controller.step() + _close_mocked_step_coroutine(mock_asyncio_run_patch) # Verify state assert_helpers.assert_controller_healthy(controller) @@ -91,6 +118,7 @@ def test_step_with_exception( # Run step and expect exception with pytest.raises(Exception, match="Test error"): controller.step() + _close_mocked_step_coroutine(mock_asyncio_run_patch) # Verify controller is not healthy assert_helpers.assert_controller_healthy(controller, is_healthy=False) @@ -106,9 +134,7 @@ async def test_get_non_terminal_deployments_calls_sdk( """Test that retrieve_non_terminal_deployments calls SDK with correct statuses.""" # Setup SDK mock responses - SDK returns AsyncPaginator for each call # Use MagicMock (not AsyncMock) because .list() returns an async iterator, not a coroutine - mock_models_sdk.inference.deployments.list = MagicMock( - side_effect=lambda **kwargs: MockAsyncPaginator([sample_deployment]) - ) + mock_models_sdk.list_deployments = AsyncMock(side_effect=lambda **kwargs: MockAsyncPaginator([sample_deployment])) # Create controller and inject mock SDK with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): @@ -118,7 +144,7 @@ async def test_get_non_terminal_deployments_calls_sdk( deployment_contexts = await controller.retrieve_non_terminal_deployments() # Verify SDK was called for each non-terminal status - assert mock_models_sdk.inference.deployments.list.call_count == len(NON_TERMINAL_STATES) + assert mock_models_sdk.list_deployments.call_count == len(NON_TERMINAL_STATES) # Verify we got ModelContext objects back assert len(deployment_contexts) > 0 @@ -135,14 +161,13 @@ async def test_get_non_terminal_deployments_handles_sdk_errors( # Setup SDK mock to raise exception on first call, succeed on others def side_effect(**kwargs): - filter_dict = kwargs.get("filter", {}) - status = filter_dict.get("status") + status = _requested_status(kwargs) if status == "CREATED": raise Exception("API Error") return MockAsyncPaginator([]) # Use MagicMock (not AsyncMock) because .list() returns an async iterator, not a coroutine - mock_models_sdk.inference.deployments.list = MagicMock(side_effect=side_effect) + mock_models_sdk.list_deployments = AsyncMock(side_effect=side_effect) # Create controller and inject mock SDK with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): @@ -171,8 +196,7 @@ async def test_get_non_terminal_deployments_with_multiple_deployments( # Setup SDK mock responses - return different deployments for each status def list_side_effect(**kwargs): - filter_dict = kwargs.get("filter", {}) - status = filter_dict.get("status") + status = _requested_status(kwargs) if status == "CREATED": return MockAsyncPaginator([sample_deployment]) elif status == "READY": @@ -180,7 +204,7 @@ def list_side_effect(**kwargs): return MockAsyncPaginator([]) # Use MagicMock (not AsyncMock) because .list() returns an async iterator, not a coroutine - mock_models_sdk.inference.deployments.list = MagicMock(side_effect=list_side_effect) + mock_models_sdk.list_deployments = AsyncMock(side_effect=list_side_effect) # Create controller and inject mock SDK with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): @@ -208,7 +232,7 @@ async def test_get_model_providers_calls_sdk(mock_get_config_patch, mock_models_ mock_provider.model_deployment_id = None # Use MagicMock (not AsyncMock) because .list() returns an async iterator, not a coroutine - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([mock_provider])) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([mock_provider])) # Create controller and inject mock SDK with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): @@ -218,7 +242,7 @@ async def test_get_model_providers_calls_sdk(mock_get_config_patch, mock_models_ provider_contexts = await controller.retrieve_model_providers() # Verify SDK was called - mock_models_sdk.inference.providers.list.assert_called_once() + mock_models_sdk.list_providers.assert_called_once() # Verify we got ModelContext objects back assert provider_contexts is not None @@ -242,15 +266,14 @@ async def test_async_controller_step_calls_reconcilers(mock_get_config_patch, mo # Mock SDK to return deployment only for CREATED status def list_deployments_side_effect(**kwargs): - filter_dict = kwargs.get("filter", {}) - status = filter_dict.get("status") + status = _requested_status(kwargs) if status == "CREATED": return MockAsyncPaginator([mock_deployment]) return MockAsyncPaginator([]) # Use MagicMock (not AsyncMock) because .list() returns an async iterator, not a coroutine - mock_models_sdk.inference.deployments.list = MagicMock(side_effect=list_deployments_side_effect) - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([mock_provider])) + mock_models_sdk.list_deployments = AsyncMock(side_effect=list_deployments_side_effect) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([mock_provider])) # Create controller and inject mock SDK with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): @@ -279,13 +302,47 @@ def list_deployments_side_effect(**kwargs): assert provider_contexts[0].model_provider == mock_provider +@pytest.mark.asyncio +async def test_async_controller_step_skips_contexts_without_a_deployment( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """Orphan detection derives its known set only from contexts that carry one. + + ``ModelContext.model_deployment`` is Optional. Reading it unguarded raises + inside the step, which aborts provider reconciliation and VirtualModel cleanup + for that tick and marks the controller unhealthy, every tick, until fixed. + """ + deployment = MagicMock() + deployment.workspace = "default" + deployment.name = "has-deployment" + contexts = [ + ModelContext(model_deployment=deployment, model_deployment_config=None, model_entity=None), + ModelContext(model_deployment=None, model_deployment_config=None, model_entity=None), + ] + + with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): + controller = ModelsController(backend_registry=mock_backend_registry) + + controller.retrieve_non_terminal_deployments = AsyncMock(return_value=contexts) + controller.retrieve_error_deployments = AsyncMock(return_value=[]) + controller.retrieve_model_providers = AsyncMock(return_value=[]) + controller._deployment_reconciler.reconcile_deployments = AsyncMock() + controller._deployment_reconciler.reconcile_orphans = AsyncMock() + controller._deployment_reconciler.gc_error_deployments = AsyncMock() + controller._provider_reconciler.reconcile_model_providers = AsyncMock() + + await controller.async_controller_step() + + controller._deployment_reconciler.reconcile_orphans.assert_awaited_once_with({"default/has-deployment"}) + + @pytest.mark.asyncio async def test_async_controller_step_runs_provider_reconciler_with_no_providers( mock_get_config_patch, mock_models_sdk, mock_backend_registry ): """The provider reconciler still runs with an empty list so VM orphan cleanup can execute.""" - mock_models_sdk.inference.deployments.list = MagicMock(return_value=MockAsyncPaginator([])) - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_deployments = AsyncMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([])) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -305,8 +362,8 @@ async def test_async_controller_step_skips_provider_reconciler_when_provider_lis mock_get_config_patch, mock_models_sdk, mock_backend_registry ): """A provider list failure must not look like a successful empty list to cleanup.""" - mock_models_sdk.inference.deployments.list = MagicMock(return_value=MockAsyncPaginator([])) - mock_models_sdk.inference.providers.list = MagicMock(side_effect=RuntimeError("providers unavailable")) + mock_models_sdk.list_deployments = AsyncMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_providers = AsyncMock(side_effect=RuntimeError("providers unavailable")) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -336,6 +393,7 @@ def test_step_handles_cancelled_error( # Should not raise -- CancelledError is expected during shutdown controller.step() + _close_mocked_step_coroutine(mock_asyncio_run_patch) # Controller should not be marked healthy (step didn't complete) assert controller._is_healthy is False @@ -445,7 +503,7 @@ async def test_retrieve_model_entity_for_config_uses_model_entity_id_when_set( mock_entity = MagicMock() mock_entity.workspace = "my-ws" mock_entity.name = "my-model" - mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(mock_entity)) config = MagicMock() config.model_entity_id = "my-ws/my-model" @@ -459,7 +517,7 @@ async def test_retrieve_model_entity_for_config_uses_model_entity_id_when_set( result = await controller._retrieve_model_entity_for_config(config) assert result is mock_entity - mock_models_sdk.models.retrieve.assert_called_once_with(name="my-model", workspace="my-ws") + mock_models_sdk.get_model.assert_called_once_with(name="my-model", workspace="my-ws") @pytest.mark.asyncio @@ -468,7 +526,7 @@ async def test_retrieve_model_entity_for_config_uses_model_entity_id_with_revisi ): """When config.model_entity_id includes @revision, revision is passed to retrieve.""" mock_entity = MagicMock() - mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(mock_entity)) config = MagicMock() config.model_entity_id = "my-ws/my-model@v2" @@ -479,8 +537,8 @@ async def test_retrieve_model_entity_for_config_uses_model_entity_id_with_revisi result = await controller._retrieve_model_entity_for_config(config) assert result is mock_entity - mock_models_sdk.models.retrieve.assert_called_once_with(name="my-model@v2", workspace="my-ws") - call_kw = mock_models_sdk.models.retrieve.call_args[1] + mock_models_sdk.get_model.assert_called_once_with(name="my-model@v2", workspace="my-ws") + call_kw = mock_models_sdk.get_model.call_args[1] assert call_kw["name"] == "my-model@v2" assert call_kw["workspace"] == "my-ws" @@ -491,7 +549,7 @@ async def test_retrieve_model_entity_for_config_falls_back_to_nim_deployment_whe ): """When config.model_entity_id is not set, entity is derived from nim_deployment.""" mock_entity = MagicMock() - mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(mock_entity)) config = MagicMock() config.model_entity_id = None @@ -505,7 +563,7 @@ async def test_retrieve_model_entity_for_config_falls_back_to_nim_deployment_whe result = await controller._retrieve_model_entity_for_config(config) assert result is mock_entity - mock_models_sdk.models.retrieve.assert_called_once_with(name="nim-model@v1", workspace="nim-ns") + mock_models_sdk.get_model.assert_called_once_with(name="nim-model@v1", workspace="nim-ns") @pytest.mark.asyncio @@ -522,7 +580,7 @@ async def test_retrieve_model_entity_for_config_returns_none_when_no_nim_deploym result = await controller._retrieve_model_entity_for_config(config) assert result is None - mock_models_sdk.models.retrieve.assert_not_called() + mock_models_sdk.get_model.assert_not_called() @pytest.mark.asyncio @@ -531,7 +589,7 @@ async def test_retrieve_model_entity_for_config_invalid_model_entity_id_falls_ba ): """When model_entity_id is set but unparseable (e.g. no slash), fall back to nim_deployment.""" mock_entity = MagicMock() - mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(mock_entity)) config = MagicMock() config.model_entity_id = "bogus" @@ -545,7 +603,126 @@ async def test_retrieve_model_entity_for_config_invalid_model_entity_id_falls_ba result = await controller._retrieve_model_entity_for_config(config) assert result is mock_entity - mock_models_sdk.models.retrieve.assert_called_once_with(name="fallback-model", workspace="fallback-ns") + mock_models_sdk.get_model.assert_called_once_with(name="fallback-model", workspace="fallback-ns") + + +# ============================================================================= +# _retrieve_model_entity: entity cache vs direct fetch +# +# The cache is refreshed at the start of every controller phase, so at runtime +# `loaded` is True and revision-less lookups are answered from it. The tests above +# reach the direct-fetch path only because they never run a phase. +# ============================================================================= + + +async def _controller_with_cache(mock_models_sdk, mock_backend_registry, entities): + """Build a controller whose entity cache holds ``entities``.""" + mock_models_sdk.list_models = AsyncMock(return_value=MockAsyncPaginator(entities)) + with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): + controller = ModelsController(backend_registry=mock_backend_registry) + await controller._entity_cache.refresh() + mock_models_sdk.get_model.reset_mock() + return controller + + +def _entity(workspace, name): + entity = MagicMock() + entity.workspace = workspace + entity.name = name + return entity + + +@pytest.mark.asyncio +async def test_retrieve_model_entity_reads_the_loaded_cache_instead_of_the_api( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """A revision-less hit is served from the snapshot, costing no round trip. + + Avoiding the per-model round trip is the entire point of the cache, so the + assertion that matters is that get_model was never called. + """ + entity = _entity("my-ws", "my-model") + controller = await _controller_with_cache(mock_models_sdk, mock_backend_registry, [entity]) + + result = await controller._retrieve_model_entity(workspace="my-ws", model_name="my-model") + + assert result is entity + mock_models_sdk.get_model.assert_not_called() + + +@pytest.mark.asyncio +async def test_retrieve_model_entity_returns_none_on_cache_miss_without_falling_back( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """A miss against a loaded cache is reported as absent, with no API fallback. + + This is a deliberate behaviour change: an entity created after the phase's + refresh reads as nonexistent until the next tick re-reads the snapshot. + """ + controller = await _controller_with_cache(mock_models_sdk, mock_backend_registry, []) + + result = await controller._retrieve_model_entity(workspace="my-ws", model_name="absent") + + assert result is None + mock_models_sdk.get_model.assert_not_called() + + +@pytest.mark.asyncio +async def test_retrieve_model_entity_bypasses_the_cache_when_a_revision_is_given( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """A revision resolves server-side and has no cache key, so it must be fetched. + + The cache is keyed by bare name, so answering a revisioned lookup from it + would hand back whichever revision happened to be current. + """ + cached = _entity("my-ws", "my-model") + controller = await _controller_with_cache(mock_models_sdk, mock_backend_registry, [cached]) + fetched = _entity("my-ws", "my-model") + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(fetched)) + + result = await controller._retrieve_model_entity(workspace="my-ws", model_name="my-model", revision="v2") + + assert result is fetched + mock_models_sdk.get_model.assert_called_once_with(name="my-model@v2", workspace="my-ws") + + +@pytest.mark.asyncio +async def test_retrieve_model_entity_queries_the_api_while_the_cache_is_unloaded( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """Before the first refresh a miss is indistinguishable from absence.""" + entity = _entity("my-ws", "my-model") + mock_models_sdk.get_model = AsyncMock(return_value=MockResponse(entity)) + + with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): + controller = ModelsController(backend_registry=mock_backend_registry) + assert controller._entity_cache.loaded is False + result = await controller._retrieve_model_entity(workspace="my-ws", model_name="my-model") + + assert result is entity + mock_models_sdk.get_model.assert_called_once_with(name="my-model", workspace="my-ws") + + +@pytest.mark.asyncio +async def test_retrieve_model_entity_returns_none_when_the_api_reports_not_found( + mock_get_config_patch, mock_models_sdk, mock_backend_registry +): + """NIMs with baked-in weights have no registered entity; that is not an error. + + The clause catches the typed client's NotFoundError. Catching the umbrella + SDK's instead would let this escape into the generic handler and be logged as + an unexpected exception on every pass. + """ + mock_models_sdk.get_model = AsyncMock( + side_effect=NotFoundError(httpx.Response(404, request=httpx.Request("GET", "http://test"))) + ) + + with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): + controller = ModelsController(backend_registry=mock_backend_registry) + result = await controller._retrieve_model_entity(workspace="my-ws", model_name="missing") + + assert result is None # ============================================================================= @@ -559,7 +736,7 @@ async def test_retrieve_error_deployments_calls_sdk(mock_get_config_patch, mock_ mock_deployment = MagicMock() mock_deployment.status = "ERROR" - mock_models_sdk.inference.deployments.list = MagicMock(return_value=MockAsyncPaginator([mock_deployment])) + mock_models_sdk.list_deployments = AsyncMock(return_value=MockAsyncPaginator([mock_deployment])) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -568,11 +745,13 @@ async def test_retrieve_error_deployments_calls_sdk(mock_get_config_patch, mock_ assert len(result) == 1 assert result[0] == mock_deployment - mock_models_sdk.inference.deployments.list.assert_called_once_with( + mock_models_sdk.list_deployments.assert_called_once_with( workspace="-", - filter={"status": "ERROR"}, - all_versions=True, - page_size=1000, + query_params={ + "filter": json.dumps({"status": "ERROR"}), + "all_versions": True, + "page_size": 1000, + }, ) @@ -581,7 +760,7 @@ async def test_retrieve_error_deployments_handles_sdk_error( mock_get_config_patch, mock_models_sdk, mock_backend_registry ): """Test that retrieve_error_deployments returns empty list on SDK error.""" - mock_models_sdk.inference.deployments.list = MagicMock(side_effect=Exception("API Error")) + mock_models_sdk.list_deployments = AsyncMock(side_effect=Exception("API Error")) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -603,14 +782,13 @@ async def test_async_controller_step_calls_gc(mock_get_config_patch, mock_models def list_side_effect(**kwargs): nonlocal call_count - filter_dict = kwargs.get("filter", {}) - status = filter_dict.get("status") + status = _requested_status(kwargs) if status == "ERROR": return MockAsyncPaginator([mock_error_deployment]) return MockAsyncPaginator([]) - mock_models_sdk.inference.deployments.list = MagicMock(side_effect=list_side_effect) - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([mock_provider])) + mock_models_sdk.list_deployments = AsyncMock(side_effect=list_side_effect) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([mock_provider])) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -636,8 +814,8 @@ async def test_async_controller_step_skips_gc_when_no_error_deployments( mock_provider = MagicMock() mock_provider.model_deployment_id = None - mock_models_sdk.inference.deployments.list = MagicMock(return_value=MockAsyncPaginator([])) - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([mock_provider])) + mock_models_sdk.list_deployments = AsyncMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([mock_provider])) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry) @@ -659,8 +837,8 @@ async def test_async_controller_step_stop_signal_skips_gc( """Test that GC is skipped when stop signal is set before GC runs.""" stop_signal = threading.Event() - mock_models_sdk.inference.deployments.list = MagicMock(return_value=MockAsyncPaginator([])) - mock_models_sdk.inference.providers.list = MagicMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_deployments = AsyncMock(return_value=MockAsyncPaginator([])) + mock_models_sdk.list_providers = AsyncMock(return_value=MockAsyncPaginator([])) with patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk", return_value=mock_models_sdk): controller = ModelsController(backend_registry=mock_backend_registry, stop_signal=stop_signal) diff --git a/services/core/models/tests/unit/controllers/test_provider_reconciler.py b/services/core/models/tests/unit/controllers/test_provider_reconciler.py index eaf867846e..9a417fce31 100644 --- a/services/core/models/tests/unit/controllers/test_provider_reconciler.py +++ b/services/core/models/tests/unit/controllers/test_provider_reconciler.py @@ -5,21 +5,25 @@ import logging from datetime import datetime, timedelta, timezone +from typing import Generic, TypeVar from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest -from nemo_platform import AsyncNeMoPlatform -from nemo_platform._exceptions import APIStatusError, ConflictError, NotFoundError -from nemo_platform.types.inference import ServedModelMapping -from nemo_platform.types.inference.model_provider import ModelProvider +from nemo_platform._exceptions import APIStatusError, ConflictError +from nemo_platform._exceptions import NotFoundError as StainlessNotFoundError +from nemo_platform_plugin.client.errors import NemoHTTPError, NotFoundError +from nemo_platform_plugin.models.types import ModelProvider, ModelProviderStatus, ServedModelMapping from nmp.core.models.config import ControllerConfig from nmp.core.models.controllers.context import ModelContext from nmp.core.models.controllers.entity_cache import ModelEntityCache from nmp.core.models.controllers.provider_reconciler import ( + _VIRTUAL_MODEL_PAGE_SIZE, PROVIDER_ERROR_RETRY_INTERVAL_SECONDS, PROVIDER_ERROR_THRESHOLD_SECONDS, PROVIDER_LOST_THRESHOLD_SECONDS, ArtifactDetails, + DiscoveredModel, DiscoveryNonCompliant, DiscoverySuccess, DiscoveryTransientError, @@ -29,12 +33,36 @@ _is_valid_served_model_entity_id, _resolve_base_backend_model_id, ) -from nmp.core.models.schemas import ModelProviderStatus -from .conftest import AsyncPaginator, make_entity, seed_entity_cache +from .conftest import AsyncPaginator, make_entity, make_models_client, seed_entity_cache +T = TypeVar("T") -def _discovery_models_from_ids(ids: list[str]) -> list[dict]: + +class _Response(Generic[T]): + def __init__(self, data: T) -> None: + self._data = data + + def data(self) -> T: + return self._data + + +def _response(data: T) -> _Response[T]: + return _Response(data) + + +def _client_error(error_type: type[NemoHTTPError], status_code: int) -> NemoHTTPError: + return error_type(httpx.Response(status_code, request=httpx.Request("GET", "http://test"))) + + +def _api_status_error(status_code: int, detail: str = "upstream error") -> APIStatusError: + """An umbrella-SDK API failure, as opposed to a bug in our own code.""" + response = MagicMock() + response.status_code = status_code + return APIStatusError("api failure", response=response, body={"detail": detail}) + + +def _discovery_models_from_ids(ids: list[str]) -> list[DiscoveredModel]: """Build GET /v1/models data[] entries (id only; root/parent omitted in external-path tests).""" return [{"id": i, "root": None, "parent": None} for i in ids] @@ -97,7 +125,7 @@ def controller_config(): @pytest.fixture def mock_models_sdk(): """Create a mock AsyncNeMoPlatform SDK.""" - sdk = MagicMock(spec=AsyncNeMoPlatform) + sdk = MagicMock() # virtual_models.create must be an AsyncMock so tests that exercise the full # reconcile path don't fail when _ensure_passthrough_virtual_model awaits it. sdk.inference.virtual_models.create = AsyncMock(return_value=None) @@ -110,9 +138,15 @@ def mock_models_sdk(): @pytest.fixture -def entity_cache(mock_models_sdk): - """Model Entity cache backed by the mock SDK, pre-loaded and empty.""" - return ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None) +def mock_models_client(): + """Create a typed Models client mock distinct from the umbrella SDK mock.""" + return make_models_client() + + +@pytest.fixture +def entity_cache(mock_models_client): + """Model Entity cache backed by the mock Models client, pre-loaded and empty.""" + return ModelEntityCache(models_client=mock_models_client, emit_heartbeat=lambda: None) @pytest.fixture @@ -122,10 +156,11 @@ def heartbeat_calls(): @pytest.fixture -def reconciler(mock_models_sdk, controller_config, entity_cache, heartbeat_calls): +def reconciler(mock_models_client, mock_models_sdk, controller_config, entity_cache, heartbeat_calls): """Create a ModelProviderReconciler instance.""" return ModelProviderReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_client, + platform_sdk=mock_models_sdk, controller_config=controller_config, entity_cache=entity_cache, emit_heartbeat=lambda: heartbeat_calls.append(1), @@ -174,16 +209,17 @@ async def test_get_available_models_from_provider_success(reconciler, mock_model @pytest.mark.asyncio -async def test_discover_models_passes_configured_timeout(mock_models_sdk): +async def test_discover_models_passes_configured_timeout(mock_models_sdk, mock_models_client): """Discovery should honor controller_config.provider_discovery_timeout_seconds.""" mock_models_sdk.inference.gateway.provider.get = AsyncMock( return_value={"object": "list", "data": [{"id": "model-1"}]} ) config = ControllerConfig(provider_discovery_timeout_seconds=240) reconciler = ModelProviderReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_client, + platform_sdk=mock_models_sdk, controller_config=config, - entity_cache=ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None), + entity_cache=ModelEntityCache(models_client=mock_models_client, emit_heartbeat=lambda: None), emit_heartbeat=lambda: None, ) @@ -213,15 +249,16 @@ async def test_discover_models_passes_configured_timeout(mock_models_sdk): ) @pytest.mark.asyncio async def test_discover_models_uses_discovery_sdk_with_configured_retries( - mock_models_sdk, max_retries, expect_get_call_kwargs + mock_models_sdk, mock_models_client, max_retries, expect_get_call_kwargs ): """Discovery SDK should honor controller_config.provider_discovery_max_retries.""" discovery_sdk = _configure_discovery_sdk(mock_models_sdk) config = ControllerConfig(provider_discovery_max_retries=max_retries) reconciler = ModelProviderReconciler( - models_sdk=mock_models_sdk, + models_client=mock_models_client, + platform_sdk=mock_models_sdk, controller_config=config, - entity_cache=ModelEntityCache(models_sdk=mock_models_sdk, emit_heartbeat=lambda: None), + entity_cache=ModelEntityCache(models_client=mock_models_client, emit_heartbeat=lambda: None), emit_heartbeat=lambda: None, ) @@ -525,11 +562,10 @@ async def test_get_artifact_details_external_provider(reconciler): ) assert details.fileset_url is None - assert details.api_endpoint == { - "url": "https://external-api.com", - "model_id": "test-model", - "format": "openai", - } + assert details.api_endpoint is not None + assert str(details.api_endpoint.url) == "https://external-api.com/" + assert details.api_endpoint.model_id == "test-model" + assert details.api_endpoint.format == "openai" @pytest.mark.asyncio @@ -643,7 +679,7 @@ async def test_get_artifact_details_handles_exception(reconciler): async def test_ensure_model_entity_creates_new_entity(reconciler): """Test creating a new model entity when it doesn't exist.""" # Mock entity doesn't exist - reconciler._models_sdk.models.create = AsyncMock() + reconciler._models_client.create_model = AsyncMock(return_value=_response(MagicMock())) # Mock context ctx = ModelContext( @@ -664,31 +700,30 @@ async def test_ensure_model_entity_creates_new_entity(reconciler): ) await reconciler._entity_cache.flush() - # Verify entity creation was called - reconciler._models_sdk.models.create.assert_called_once_with( - workspace="test-ns", - name="test-model", - description="Auto-discovered model from provider test-ns/test-provider", - model_providers=["test-ns/test-provider"], - backend_format="OPENAI_CHAT", - fileset="test/model", - ) + # Verify typed entity creation request. + reconciler._models_client.create_model.assert_called_once() + call = reconciler._models_client.create_model.call_args + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].name == "test-model" + assert call.kwargs["body"].description == "Auto-discovered model from provider test-ns/test-provider" + assert call.kwargs["body"].model_providers == ["test-ns/test-provider"] + assert call.kwargs["body"].backend_format.value == "OPENAI_CHAT" + assert call.kwargs["body"].fileset == "test/model" @pytest.mark.asyncio async def test_ensure_model_entity_updates_existing_adds_provider(reconciler): - """Test updating existing entity to add provider to model_providers list.""" - # Mock existing entity + """A provider-link update leaves an existing non-default backend format unset.""" existing_entity = MagicMock() existing_entity.model_providers = ["other-ns/other-provider"] existing_entity.fileset = None existing_entity.api_endpoint = None + existing_entity.backend_format = "ANTHROPIC_MESSAGES" existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -706,13 +741,17 @@ async def test_ensure_model_entity_updates_existing_adds_provider(reconciler): ) await reconciler._entity_cache.flush() - # Verify update was called to add provider - reconciler._models_sdk.models.update.assert_called_once_with( - name="test-model", - workspace="test-ns", - model_providers=["other-ns/other-provider", "test-ns/test-provider"], - backend_format="OPENAI_CHAT", - ) + # Verify typed update request adds the provider. + reconciler._models_client.update_model.assert_called_once() + call = reconciler._models_client.update_model.call_args + assert call.kwargs["name"] == "test-model" + assert call.kwargs["workspace"] == "test-ns" + body = call.kwargs["body"] + assert body.model_providers == ["other-ns/other-provider", "test-ns/test-provider"] + assert body.model_fields_set == {"model_providers"} + assert body.model_dump(mode="json", exclude_unset=True) == { + "model_providers": ["other-ns/other-provider", "test-ns/test-provider"] + } @pytest.mark.asyncio @@ -725,9 +764,8 @@ async def test_ensure_model_entity_skips_if_provider_already_linked(reconciler): existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(), @@ -746,7 +784,7 @@ async def test_ensure_model_entity_skips_if_provider_already_linked(reconciler): await reconciler._entity_cache.flush() # Verify update was NOT called - reconciler._models_sdk.models.update.assert_not_called() + reconciler._models_client.update_model.assert_not_called() @pytest.mark.asyncio @@ -760,9 +798,8 @@ async def test_ensure_model_entity_backfills_missing_backend_format(reconciler): existing_entity.workspace = "test-ns" existing_entity.name = "anthropic.claude-3-5-sonnet" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(), @@ -780,11 +817,14 @@ async def test_ensure_model_entity_backfills_missing_backend_format(reconciler): ) await reconciler._entity_cache.flush() - reconciler._models_sdk.models.update.assert_called_once_with( - name="anthropic.claude-3-5-sonnet", - workspace="test-ns", - backend_format="ANTHROPIC_MESSAGES", - ) + reconciler._models_client.update_model.assert_called_once() + call = reconciler._models_client.update_model.call_args + assert call.kwargs["name"] == "anthropic.claude-3-5-sonnet" + assert call.kwargs["workspace"] == "test-ns" + body = call.kwargs["body"] + assert body.backend_format.value == "ANTHROPIC_MESSAGES" + assert body.model_fields_set == {"backend_format"} + assert body.model_dump(mode="json", exclude_unset=True) == {"backend_format": "ANTHROPIC_MESSAGES"} @pytest.mark.asyncio @@ -798,9 +838,8 @@ async def test_ensure_model_entity_adds_artifact_to_existing_without_artifact(re existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -820,14 +859,14 @@ async def test_ensure_model_entity_adds_artifact_to_existing_without_artifact(re ) await reconciler._entity_cache.flush() - # Verify update includes artifact - reconciler._models_sdk.models.update.assert_called_once_with( - name="test-model", - workspace="test-ns", - model_providers=["other-ns/other-provider", "test-ns/test-provider"], - backend_format="OPENAI_CHAT", - fileset="test/model", - ) + # Verify typed update request includes the artifact. + reconciler._models_client.update_model.assert_called_once() + call = reconciler._models_client.update_model.call_args + assert call.kwargs["name"] == "test-model" + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].model_providers == ["other-ns/other-provider", "test-ns/test-provider"] + assert call.kwargs["body"].backend_format.value == "OPENAI_CHAT" + assert call.kwargs["body"].fileset == "test/model" @pytest.mark.asyncio @@ -841,9 +880,8 @@ async def test_ensure_model_entity_doesnt_overwrite_existing_artifact(reconciler existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -863,9 +901,10 @@ async def test_ensure_model_entity_doesnt_overwrite_existing_artifact(reconciler ) await reconciler._entity_cache.flush() - # Verify update does NOT include fileset (since it already exists) - call_kwargs = reconciler._models_sdk.models.update.call_args.kwargs - assert "fileset" not in call_kwargs + # Verify typed update does not replace the existing fileset. + body = reconciler._models_client.update_model.call_args.kwargs["body"] + assert "fileset" not in body.model_fields_set + assert "fileset" not in body.model_dump(mode="json", exclude_unset=True) @pytest.mark.asyncio @@ -879,9 +918,8 @@ async def test_ensure_model_entity_doesnt_overwrite_existing_backend_format(reco existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -899,8 +937,9 @@ async def test_ensure_model_entity_doesnt_overwrite_existing_backend_format(reco ) await reconciler._entity_cache.flush() - call_kwargs = reconciler._models_sdk.models.update.call_args.kwargs - assert "backend_format" not in call_kwargs + body = reconciler._models_client.update_model.call_args.kwargs["body"] + assert "backend_format" not in body.model_fields_set + assert "backend_format" not in body.model_dump(mode="json", exclude_unset=True) @pytest.mark.asyncio @@ -914,9 +953,8 @@ async def test_ensure_model_entity_handles_null_model_providers(reconciler): existing_entity.workspace = "test-ns" existing_entity.name = "test-model" - reconciler._models_sdk.models.list = MagicMock(return_value=_AsyncPaginator([existing_entity])) - await reconciler._entity_cache.refresh() - reconciler._models_sdk.models.update = AsyncMock() + await seed_entity_cache(reconciler._models_client, reconciler._entity_cache, [existing_entity]) + reconciler._models_client.update_model = AsyncMock(return_value=_response(existing_entity)) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -934,22 +972,20 @@ async def test_ensure_model_entity_handles_null_model_providers(reconciler): ) await reconciler._entity_cache.flush() - # Verify update was called with provider as first in list - reconciler._models_sdk.models.update.assert_called_once_with( - name="test-model", - workspace="test-ns", - model_providers=["test-ns/test-provider"], - backend_format="OPENAI_CHAT", - ) + # Verify typed update request has the provider first. + reconciler._models_client.update_model.assert_called_once() + call = reconciler._models_client.update_model.call_args + assert call.kwargs["name"] == "test-model" + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["body"].model_providers == ["test-ns/test-provider"] + assert call.kwargs["body"].backend_format.value == "OPENAI_CHAT" @pytest.mark.asyncio async def test_ensure_model_entity_handles_create_exception(reconciler): """Test handling exception during entity creation.""" - reconciler._models_sdk.models.retrieve = AsyncMock( - side_effect=NotFoundError("Not found", response=MagicMock(), body=None) - ) - reconciler._models_sdk.models.create = AsyncMock(side_effect=Exception("Creation failed")) + reconciler._models_client.get_model = AsyncMock(side_effect=_client_error(NotFoundError, 404)) + reconciler._models_client.create_model = AsyncMock(side_effect=Exception("Creation failed")) ctx = ModelContext( model_provider=MagicMock(host_url="https://api.com"), @@ -970,7 +1006,7 @@ async def test_ensure_model_entity_handles_create_exception(reconciler): @pytest.mark.asyncio -async def test_entity_cache_load_failure_propagates_and_stages_nothing(reconciler, mock_models_sdk): +async def test_entity_cache_load_failure_propagates_and_stages_nothing(reconciler, mock_models_client): """An unreadable entity list surfaces to the caller with nothing staged. The controller loads the cache at the start of the phase, so this failure aborts @@ -978,30 +1014,30 @@ async def test_entity_cache_load_failure_propagates_and_stages_nothing(reconcile """ mock_response = MagicMock() mock_response.status_code = 503 - mock_models_sdk.models.list = MagicMock( + mock_models_client.list_models = AsyncMock( side_effect=APIStatusError( "Service unavailable", response=mock_response, body={"detail": "upstream error"}, ) ) - mock_models_sdk.models.create = AsyncMock() - mock_models_sdk.models.update = AsyncMock() + mock_models_client.create_model = AsyncMock() + mock_models_client.update_model = AsyncMock() with pytest.raises(APIStatusError): await reconciler._entity_cache.refresh() await reconciler._entity_cache.flush() - mock_models_sdk.models.create.assert_not_called() - mock_models_sdk.models.update.assert_not_called() + mock_models_client.create_model.assert_not_called() + mock_models_client.update_model.assert_not_called() @pytest.mark.asyncio async def test_virtual_model_listing_failure_does_not_abort_provider_reconciliation( - reconciler, mock_models_sdk, entity_cache + reconciler, mock_models_sdk, mock_models_client, entity_cache ): """A VirtualModel listing failure must not cost discovery, status, or entity work.""" - await seed_entity_cache(mock_models_sdk, entity_cache, []) + await seed_entity_cache(mock_models_client, entity_cache, []) provider = MagicMock() provider.workspace = "test-ns" provider.name = "test-provider" @@ -1015,9 +1051,9 @@ async def test_virtual_model_listing_failure_does_not_abort_provider_reconciliat model_entity=None, ) - mock_models_sdk.inference.virtual_models.list = MagicMock(side_effect=Exception("listing unavailable")) - mock_models_sdk.models.create = AsyncMock() - mock_models_sdk.inference.providers.update_status = AsyncMock() + mock_models_sdk.inference.virtual_models.list = MagicMock(side_effect=_api_status_error(503)) + mock_models_client.create_model = AsyncMock(return_value=_response(MagicMock())) + mock_models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1028,8 +1064,8 @@ async def test_virtual_model_listing_failure_does_not_abort_provider_reconciliat await entity_cache.flush() # Provider status and entity linking still happened. - mock_models_sdk.inference.providers.update_status.assert_awaited() - mock_models_sdk.models.create.assert_awaited_once() + mock_models_client.update_provider_status.assert_awaited() + mock_models_client.create_model.assert_awaited_once() # VirtualModel work was skipped rather than acted on with an unknown state. mock_models_sdk.inference.virtual_models.create.assert_not_awaited() mock_models_sdk.inference.virtual_models.delete.assert_not_awaited() @@ -1057,7 +1093,7 @@ async def test_update_model_providers_success(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1074,12 +1110,12 @@ async def test_update_model_providers_success(reconciler): assert mock_ensure.call_count == 2 # Verify provider was updated with served models - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs assert call_kwargs["name"] == "test-provider" assert call_kwargs["workspace"] == "test-ns" - assert call_kwargs["status"] == "READY" - assert len(call_kwargs["served_models"]) == 2 + assert call_kwargs["body"].status.value == "READY" + assert len(call_kwargs["body"].served_models) == 2 @pytest.mark.asyncio @@ -1099,7 +1135,7 @@ async def test_update_model_providers_filters_by_enabled_models(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1147,7 +1183,7 @@ async def test_ensure_external_entities_retries_after_transient_entity_failure(r model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1187,7 +1223,7 @@ async def test_update_model_providers_removes_no_longer_served_models(reconciler model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) # Now only serving model-1 and model-2 (model-3 removed) with patch.object( @@ -1199,9 +1235,11 @@ async def test_update_model_providers_removes_no_longer_served_models(reconciler await reconciler.reconcile_model_providers([ctx]) # Verify only model-1 and model-2 are in final served_models - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - served_model_names = {m.served_model_name for m in call_kwargs["served_models"]} + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + body = call_kwargs["body"] + served_model_names = {m.served_model_name for m in body.served_models} assert served_model_names == {"model-1", "model-2"} + assert body.model_fields_set == {"served_models", "status"} @pytest.mark.asyncio @@ -1219,7 +1257,7 @@ async def test_update_model_providers_handles_non_compliant_provider(reconciler) model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) # Provider returns DiscoveryNonCompliant (confirmed non-compliant) with patch.object(reconciler, "_discover_models", return_value=DiscoveryNonCompliant()): @@ -1230,11 +1268,11 @@ async def test_update_model_providers_handles_non_compliant_provider(reconciler) mock_ensure.assert_not_called() # Verify provider was updated with empty served_models and appropriate message - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["served_models"] == [] - assert call_kwargs["status"] == "READY" - assert "Non-OpenAI compliant" in call_kwargs["status_message"] + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].served_models == [] + assert call_kwargs["body"].status.value == "READY" + assert "Non-OpenAI compliant" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1254,7 +1292,7 @@ async def test_update_model_providers_handles_update_exception(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock(side_effect=Exception("Update failed")) + reconciler._models_client.update_provider_status = AsyncMock(side_effect=Exception("Update failed")) with patch.object( reconciler, @@ -1283,7 +1321,7 @@ async def test_update_model_providers_normalizes_model_names(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) # Model with special characters that need normalization with patch.object( @@ -1299,8 +1337,8 @@ async def test_update_model_providers_normalizes_model_names(reconciler): assert "model-with-colons" in str(mock_ensure.call_args) # Normalized # Verify served_models keeps original name - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - served_models = call_kwargs["served_models"] + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + served_models = call_kwargs["body"].served_models assert len(served_models) == 1 assert served_models[0].served_model_name == "model:with:colons" # Original assert "model-with-colons" in served_models[0].model_entity_id # Normalized @@ -1323,7 +1361,7 @@ async def test_update_model_providers_strips_same_workspace_prefix_from_model_id model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) # Backend reports model id as workspace/name (e.g. NIM_SERVED_MODEL_NAME set to workspace/name) with patch.object( @@ -1340,8 +1378,8 @@ async def test_update_model_providers_strips_same_workspace_prefix_from_model_id assert call_kwargs["model_name"] == "qwen-2-5-1-5b" # served_models should have model_entity_id = workspace/name (no duplicate prefix in name) - update_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - served_models = update_kwargs["served_models"] + update_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + served_models = update_kwargs["body"].served_models assert len(served_models) == 1 assert served_models[0].model_entity_id == "test-ns/qwen-2-5-1-5b" assert served_models[0].served_model_name == "test-ns/qwen-2-5-1-5b" @@ -1364,7 +1402,7 @@ async def test_update_model_providers_with_empty_discovery(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1378,9 +1416,9 @@ async def test_update_model_providers_with_empty_discovery(reconciler): mock_ensure.assert_not_called() # Verify provider was updated with empty served_models - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["served_models"] == [] - assert call_kwargs["status"] == "READY" + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].served_models == [] + assert call_kwargs["body"].status.value == "READY" @pytest.mark.asyncio @@ -1403,7 +1441,7 @@ async def test_update_model_providers_multiple_providers(reconciler): ctx1 = ModelContext(model_provider=provider1) ctx2 = ModelContext(model_provider=provider2) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) async def get_models_side_effect(model_provider: ModelProvider): if model_provider.workspace == "ns1": @@ -1416,7 +1454,7 @@ async def get_models_side_effect(model_provider: ModelProvider): # Verify both providers were processed assert mock_get_models.call_count == 2 - assert reconciler._models_sdk.inference.providers.update_status.call_count == 2 + assert reconciler._models_client.update_provider_status.call_count == 2 @pytest.mark.asyncio @@ -1437,14 +1475,14 @@ async def test_reconcile_preserves_served_models_on_transient_error(reconciler): model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoveryTransientError()): with patch.object(reconciler, "_ensure_model_entity_for_provider") as mock_ensure: await reconciler.reconcile_model_providers([ctx]) # Transient error must not trigger any status update — served_models are preserved implicitly - reconciler._models_sdk.inference.providers.update_status.assert_not_called() + reconciler._models_client.update_provider_status.assert_not_called() mock_ensure.assert_not_called() @@ -1469,7 +1507,7 @@ async def test_reconcile_preserves_served_models_when_deployment_base_id_unresol model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with ( patch.object(reconciler, "_discover_models", return_value=DiscoverySuccess([{"id": "test-ns/base"}])), @@ -1478,7 +1516,7 @@ async def test_reconcile_preserves_served_models_when_deployment_base_id_unresol ): await reconciler.reconcile_model_providers([ctx]) - reconciler._models_sdk.inference.providers.update_status.assert_not_called() + reconciler._models_client.update_provider_status.assert_not_called() mock_ensure.assert_not_called() # WARNING must surface the provider id so operators can correlate with # downstream "model not found" reports during a flaky prefetch tick. @@ -1506,7 +1544,7 @@ async def test_reconcile_clears_served_models_on_confirmed_non_compliant(reconci model_entity=None, ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoveryNonCompliant()): with patch.object(reconciler, "_ensure_model_entity_for_provider") as mock_ensure: @@ -1514,11 +1552,11 @@ async def test_reconcile_clears_served_models_on_confirmed_non_compliant(reconci # Non-compliant must clear served_models mock_ensure.assert_not_called() - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["served_models"] == [] - assert call_kwargs["status"] == "READY" - assert "Non-OpenAI compliant" in call_kwargs["status_message"] + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].served_models == [] + assert call_kwargs["body"].status.value == "READY" + assert "Non-OpenAI compliant" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1539,7 +1577,7 @@ async def test_reconcile_prunes_invalid_served_model_entity_ids_before_update_st ctx = ModelContext(model_provider=provider, model_deployment=None, model_deployment_config=None, model_entity=None) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) bad = ServedModelMapping(model_entity_id="ws/Bad.Name", served_model_name="Bad.Name") good = ServedModelMapping(model_entity_id="ws/model-a", served_model_name="model-a") @@ -1555,12 +1593,12 @@ async def test_reconcile_prunes_invalid_served_model_entity_ids_before_update_st ): await reconciler.reconcile_model_providers([ctx]) - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - emitted = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs["served_models"] + reconciler._models_client.update_provider_status.assert_called_once() + emitted = reconciler._models_client.update_provider_status.call_args.kwargs["body"].served_models assert [m.model_entity_id for m in emitted] == ["ws/model-a"] # Passthrough VirtualModel is attempted only for the surviving (non-LoRA) mapping. created_names = { - call.kwargs["name"] for call in reconciler._models_sdk.inference.virtual_models.create.call_args_list + call.kwargs["name"] for call in reconciler._platform_sdk.inference.virtual_models.create.call_args_list } assert created_names == {"model-a"} @@ -1584,7 +1622,7 @@ async def test_reconcile_keeps_valid_lora_composite_through_gate(reconciler): model_provider=provider, model_deployment=None, model_deployment_config=config, model_entity=None ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1598,12 +1636,12 @@ async def test_reconcile_keeps_valid_lora_composite_through_gate(reconciler): ): await reconciler.reconcile_model_providers([ctx]) - emitted = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs["served_models"] + emitted = reconciler._models_client.update_provider_status.call_args.kwargs["body"].served_models eids = {m.model_entity_id for m in emitted} assert eids == {"ws/base", "ws/base&adapters/ws/lora-1"} # Only the base entity gets a passthrough VirtualModel; LoRA is skipped by design. created_names = { - call.kwargs["name"] for call in reconciler._models_sdk.inference.virtual_models.create.call_args_list + call.kwargs["name"] for call in reconciler._platform_sdk.inference.virtual_models.create.call_args_list } assert created_names == {"base"} @@ -1656,7 +1694,7 @@ async def test_exception_in_one_provider_does_not_affect_others(reconciler): ctx_bad = ModelContext(model_provider=bad_provider) ctx_good = ModelContext(model_provider=good_provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) call_count = 0 @@ -1675,10 +1713,10 @@ async def query_side_effect(provider): # Both providers were attempted assert call_count == 2 # Good provider was still updated successfully - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs assert call_kwargs["workspace"] == "ns-good" - assert call_kwargs["status"] == "READY" + assert call_kwargs["body"].status.value == "READY" @pytest.mark.asyncio @@ -1707,7 +1745,7 @@ async def test_created_provider_escalated_to_error_after_threshold(reconciler, _ provider = _make_provider(status=ModelProviderStatus.CREATED, created_at=stale, updated_at=stale) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1716,10 +1754,10 @@ async def test_created_provider_escalated_to_error_after_threshold(reconciler, _ ): await reconciler.reconcile_model_providers([ctx]) - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" - assert "connection refused" in call_kwargs["status_message"] + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].status.value == "ERROR" + assert "connection refused" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1729,13 +1767,13 @@ async def test_created_provider_not_escalated_before_threshold(reconciler, _make provider = _make_provider(status=ModelProviderStatus.CREATED, created_at=recent, updated_at=recent) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoveryTransientError()): await reconciler.reconcile_model_providers([ctx]) # Should NOT update status — still within grace period - reconciler._models_sdk.inference.providers.update_status.assert_not_called() + reconciler._models_client.update_provider_status.assert_not_called() @pytest.mark.asyncio @@ -1767,7 +1805,7 @@ async def test_error_provider_retried_after_cooldown(reconciler, _make_provider) ) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1777,10 +1815,10 @@ async def test_error_provider_retried_after_cooldown(reconciler, _make_provider) await reconciler.reconcile_model_providers([ctx]) # Should update status to bump updated_at for next retry pacing - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["status"] == "ERROR" - assert "still down" in call_kwargs["status_message"] + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].status.value == "ERROR" + assert "still down" in call_kwargs["body"].status_message @pytest.mark.asyncio @@ -1799,17 +1837,17 @@ async def test_error_provider_transitions_to_lost(reconciler, _make_provider): updated_at=datetime.now(timezone.utc), ) - reconciler._models_sdk.inference.providers.update_status = AsyncMock(return_value=updated_provider) + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(updated_provider)) with patch.object(reconciler, "_discover_models") as mock_query: await reconciler.reconcile_model_providers([ctx]) # Should transition to LOST without attempting discovery mock_query.assert_not_called() - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["status"] == "LOST" - assert "permanently failed" in call_kwargs["status_message"] + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].status.value == "LOST" + assert "permanently failed" in call_kwargs["body"].status_message assert ctx.model_provider is updated_provider @@ -1824,7 +1862,7 @@ async def test_error_provider_recovers_to_ready(reconciler, _make_provider): ) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -1834,10 +1872,10 @@ async def test_error_provider_recovers_to_ready(reconciler, _make_provider): with patch.object(reconciler, "_ensure_model_entity_for_provider"): await reconciler.reconcile_model_providers([ctx]) - reconciler._models_sdk.inference.providers.update_status.assert_called_once() - call_kwargs = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs - assert call_kwargs["status"] == "READY" - assert len(call_kwargs["served_models"]) == 1 + reconciler._models_client.update_provider_status.assert_called_once() + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert call_kwargs["body"].status.value == "READY" + assert len(call_kwargs["body"].served_models) == 1 @pytest.mark.asyncio @@ -1846,13 +1884,13 @@ async def test_lost_provider_skipped_entirely(reconciler, _make_provider): provider = _make_provider(status=ModelProviderStatus.LOST) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models") as mock_query: await reconciler.reconcile_model_providers([ctx]) mock_query.assert_not_called() - reconciler._models_sdk.inference.providers.update_status.assert_not_called() + reconciler._models_client.update_provider_status.assert_not_called() @pytest.mark.asyncio @@ -1872,13 +1910,13 @@ async def test_ready_provider_preserves_served_models_on_transient_error(reconci ) ctx = ModelContext(model_provider=provider) - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoveryTransientError()): await reconciler.reconcile_model_providers([ctx]) # Should NOT update status — existing served_models preserved - reconciler._models_sdk.inference.providers.update_status.assert_not_called() + reconciler._models_client.update_provider_status.assert_not_called() @pytest.mark.asyncio @@ -1964,7 +2002,7 @@ async def test_reconcile_creates_passthrough_virtual_models_for_all_served_model model_entity=None, ) - mock_models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object( reconciler, @@ -2182,6 +2220,60 @@ async def test_cleanup_delete_failure_is_logged_and_non_fatal(reconciler, mock_m assert any("Failed to delete orphaned autoprovisioned VirtualModel ws/model-a" in r.message for r in caplog.records) +@pytest.mark.asyncio +async def test_cleanup_skips_orphaned_virtual_model_without_a_name(reconciler, mock_models_sdk, caplog): + """``VirtualModel.name`` is Optional on the wire, so an unnamed one must be skipped. + + Without the guard the delete goes out as ``name=None``, which is a malformed + request against IGW that the surrounding ``except Exception`` would downgrade + to a warning. + """ + mock_models_sdk.inference.virtual_models.list = MagicMock( + return_value=_AsyncPaginator( + [ + _virtual_model( + None, + default_model_entity="ws/orphan", + autoprovisioned=True, + db_version=1, + ) + ] + ) + ) + + with caplog.at_level(logging.WARNING): + vm_snapshot, _ = await reconciler._load_virtual_models() + await reconciler._cleanup_orphaned_virtual_models([], vm_snapshot) + + mock_models_sdk.inference.virtual_models.delete.assert_not_awaited() + assert any("Skipping autoprovisioned VirtualModel with no name" in r.message for r in caplog.records) + + +@pytest.mark.asyncio +async def test_load_virtual_models_does_not_swallow_programming_errors(reconciler, mock_models_sdk): + """Only API failures mean "state unknown"; a bug in our own code must surface. + + An earlier version caught every exception here, which turned an AttributeError + into a routine "listing failed" warning and silently disabled orphan cleanup + on every pass while the suite stayed green. + """ + mock_models_sdk.inference.virtual_models.list = MagicMock(side_effect=AttributeError("no such attribute")) + + with pytest.raises(AttributeError): + await reconciler._load_virtual_models() + + +@pytest.mark.asyncio +async def test_load_virtual_models_reports_unknown_state_on_api_failure(reconciler, mock_models_sdk): + """A genuine API failure still degrades to None rather than propagating.""" + mock_models_sdk.inference.virtual_models.list = MagicMock(side_effect=_api_status_error(503)) + + assert await reconciler._load_virtual_models() is None + mock_models_sdk.inference.virtual_models.list.assert_called_once_with( + workspace="-", page_size=_VIRTUAL_MODEL_PAGE_SIZE + ) + + @pytest.mark.asyncio async def test_cleanup_skips_orphaned_virtual_model_without_db_version(reconciler, mock_models_sdk, caplog): """Cleanup must not fall back to an unconditional delete when the listed VM has no version.""" @@ -2233,9 +2325,9 @@ async def test_deployment_backed_never_calls_ensure_model_entity(reconciler, moc model_entity=None, ) - mock_models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) - discovered = [ + discovered: list[DiscoveredModel] = [ {"id": "ws/base-entity", "root": "ws/base-entity", "parent": None}, {"id": "adapter-1", "root": "/scratch/loras/adapter-1", "parent": "ws/base-entity"}, ] @@ -2244,8 +2336,8 @@ async def test_deployment_backed_never_calls_ensure_model_entity(reconciler, moc await reconciler.reconcile_model_providers([ctx]) mock_ensure.assert_not_called() - call_kwargs = mock_models_sdk.inference.providers.update_status.call_args.kwargs - assert len(call_kwargs["served_models"]) == 2 + call_kwargs = reconciler._models_client.update_provider_status.call_args.kwargs + assert len(call_kwargs["body"].served_models) == 2 @pytest.mark.asyncio @@ -2269,7 +2361,7 @@ async def test_reconcile_creates_virtual_models_for_previously_served_models(rec model_entity=None, ) - mock_models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) # Discover old-model (already served) and new-model (new) with patch.object( @@ -2301,7 +2393,7 @@ async def test_reconcile_does_not_create_virtual_models_for_non_compliant_provid model_entity=None, ) - mock_models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoveryNonCompliant()): await reconciler.reconcile_model_providers([ctx]) @@ -2327,7 +2419,7 @@ async def test_reconcile_creates_virtual_models_even_when_update_status_fails(re ) # update_status raises — VirtualModel creation must still run - mock_models_sdk.inference.providers.update_status = AsyncMock(side_effect=Exception("service unavailable")) + reconciler._models_client.update_provider_status = AsyncMock(side_effect=Exception("service unavailable")) mock_models_sdk.inference.virtual_models.list = MagicMock(return_value=_AsyncPaginator([])) with patch.object( @@ -2372,17 +2464,17 @@ async def test_deployment_backed_served_models_base_lora_prompt_tuned(reconciler model_entity=None, ) - discovered = [ + discovered: list[DiscoveredModel] = [ {"id": "e2e-ws/qwen-lora-base", "root": "e2e-ws/qwen-lora-base", "parent": None}, {"id": "qwen-lora-base-lora-e2e-dataset-5c30", "root": "/scratch/loras/...", "parent": "e2e-ws/qwen-lora-base"}, {"id": "qwen-lora-prompt-tuned", "root": "e2e-ws/qwen-lora-base", "parent": None}, ] - reconciler._models_sdk.inference.providers.update_status = AsyncMock() + reconciler._models_client.update_provider_status = AsyncMock(return_value=_response(MagicMock())) with patch.object(reconciler, "_discover_models", return_value=DiscoverySuccess(discovered)): await reconciler.reconcile_model_providers([ctx]) - served = reconciler._models_sdk.inference.providers.update_status.call_args.kwargs["served_models"] + served = reconciler._models_client.update_provider_status.call_args.kwargs["body"].served_models by_entity_id = {m.model_entity_id: m.served_model_name for m in served} assert by_entity_id["e2e-ws/qwen-lora-base"] == "e2e-ws/qwen-lora-base" lora_entity_id = "e2e-ws/qwen-lora-base&adapters/e2e-ws/qwen-lora-base-lora-e2e-dataset-5c30" @@ -2390,7 +2482,7 @@ async def test_deployment_backed_served_models_base_lora_prompt_tuned(reconciler assert by_entity_id["e2e-ws/qwen-lora-prompt-tuned"] == "qwen-lora-prompt-tuned" assert len(served) == 3 - vm_create_calls = reconciler._models_sdk.inference.virtual_models.create.call_args_list + vm_create_calls = reconciler._platform_sdk.inference.virtual_models.create.call_args_list created_names = {call.kwargs["name"] for call in vm_create_calls} assert created_names == {"qwen-lora-base", "qwen-lora-prompt-tuned"} for call in vm_create_calls: @@ -2406,7 +2498,7 @@ def test_handle_model_deployment_provider_base_only(reconciler): provider = MagicMock() provider.workspace = "ws" provider.enabled_models = None - models = [{"id": "ws/base", "root": "ws/base", "parent": None}] + models: list[DiscoveredModel] = [{"id": "ws/base", "root": "ws/base", "parent": None}] result = DiscoverySuccess(models) ctx = ModelContext( model_provider=provider, model_deployment=None, model_deployment_config=config, model_entity=None @@ -2424,7 +2516,7 @@ def test_handle_model_deployment_provider_unmatched_skipped(reconciler): provider = MagicMock() provider.workspace = "ws" provider.enabled_models = None - models = [ + models: list[DiscoveredModel] = [ {"id": "ws/base", "root": "ws/base", "parent": None}, {"id": "other-model", "root": "other-root", "parent": None}, ] @@ -2461,7 +2553,7 @@ def test_handle_model_deployment_provider_uses_nim_when_no_model_entity_id(recon provider = MagicMock() provider.workspace = "ws" provider.enabled_models = None - models = [ + models: list[DiscoveredModel] = [ {"id": "ws/base", "root": "ws/base", "parent": None}, {"id": "adapter-x", "root": "/scratch/x", "parent": "ws/base"}, ] @@ -2881,8 +2973,10 @@ async def test_cleanup_tolerates_virtual_model_deleted_concurrently(reconciler, mock_models_sdk.inference.virtual_models.list = MagicMock( return_value=_AsyncPaginator([_virtual_model("model-a", default_model_entity="ws/model-a")]) ) + # The VirtualModel delete goes through the umbrella SDK, so it raises the + # Stainless NotFoundError, not the typed client's. mock_models_sdk.inference.virtual_models.delete = AsyncMock( - side_effect=NotFoundError("gone", response=MagicMock(), body=None) + side_effect=StainlessNotFoundError("gone", response=MagicMock(), body=None) ) vm_snapshot, _ = await reconciler._load_virtual_models() @@ -2918,12 +3012,12 @@ async def test_provider_skipped_before_discovery_keeps_served_models_unresolved( @pytest.mark.asyncio -async def test_entity_linked_by_two_providers_is_written_once(reconciler, mock_models_sdk, entity_cache): +async def test_entity_linked_by_two_providers_is_written_once(reconciler, mock_models_client, entity_cache): """Two providers serving one model produce a single Model Entity write.""" # Already carries everything except the provider links, so the only pending # change is the link each provider contributes. await seed_entity_cache( - mock_models_sdk, + mock_models_client, entity_cache, [ _existing_entity( @@ -2935,7 +3029,7 @@ async def test_entity_linked_by_two_providers_is_written_once(reconciler, mock_m ) ], ) - mock_models_sdk.models.update = AsyncMock() + mock_models_client.update_model = AsyncMock(return_value=_response(MagicMock())) ctxs = [] for provider_name in ("provider-a", "provider-b"): @@ -2961,18 +3055,18 @@ async def test_entity_linked_by_two_providers_is_written_once(reconciler, mock_m ): await reconcile_and_flush(reconciler, entity_cache, ctxs) - mock_models_sdk.models.update.assert_awaited_once_with( - workspace="test-ns", - name="shared-model", - model_providers=["test-ns/provider-a", "test-ns/provider-b"], - ) + mock_models_client.update_model.assert_awaited_once() + call = mock_models_client.update_model.call_args + assert call.kwargs["workspace"] == "test-ns" + assert call.kwargs["name"] == "shared-model" + assert call.kwargs["body"].model_providers == ["test-ns/provider-a", "test-ns/provider-b"] @pytest.mark.asyncio -async def test_converged_entities_are_not_rewritten(reconciler, mock_models_sdk, entity_cache): +async def test_converged_entities_are_not_rewritten(reconciler, mock_models_client, entity_cache): """A pass that changes nothing performs no Model Entity writes.""" await seed_entity_cache( - mock_models_sdk, + mock_models_client, entity_cache, [ _existing_entity( @@ -2984,8 +3078,8 @@ async def test_converged_entities_are_not_rewritten(reconciler, mock_models_sdk, ) ], ) - mock_models_sdk.models.update = AsyncMock() - mock_models_sdk.models.create = AsyncMock() + mock_models_client.update_model = AsyncMock(return_value=_response(MagicMock())) + mock_models_client.create_model = AsyncMock(return_value=_response(MagicMock())) provider = MagicMock() provider.workspace = "test-ns" @@ -3007,8 +3101,8 @@ async def test_converged_entities_are_not_rewritten(reconciler, mock_models_sdk, ): await reconcile_and_flush(reconciler, entity_cache, ctx and [ctx]) - mock_models_sdk.models.update.assert_not_awaited() - mock_models_sdk.models.create.assert_not_awaited() + mock_models_client.update_model.assert_not_awaited() + mock_models_client.create_model.assert_not_awaited() @pytest.mark.asyncio diff --git a/services/core/models/tests/unit/sidecars/test_adapters_controller.py b/services/core/models/tests/unit/sidecars/test_adapters_controller.py index a1503eb9f1..6bdc0d434b 100644 --- a/services/core/models/tests/unit/sidecars/test_adapters_controller.py +++ b/services/core/models/tests/unit/sidecars/test_adapters_controller.py @@ -11,9 +11,26 @@ from unittest.mock import MagicMock, patch import pytest +from nemo_platform_plugin.models.types import Adapter, FinetuningType from nmp.core.models.sidecars.adapters.main import ADAPTER_META_FILENAME, AdaptersController +class _Response: + def __init__(self, data): + self._data = data + + def data(self): + return self._data + + +class _ItemsResponse: + def __init__(self, items): + self._items = items + + def items(self): + return iter(self._items) + + def _make_adapter( name: str, fileset: str, @@ -43,6 +60,8 @@ def controller(tmp_path): ctrl = AdaptersController() ctrl.nim_peft_source = str(tmp_path) ctrl._sdk = MagicMock() + ctrl._models_client = MagicMock() + ctrl._models_client.list_models.return_value = _ItemsResponse([]) ctrl.workspace = "default" ctrl.model_name = "base-model" # Default to NIM behavior (no rewrite, no eager vLLM load); vLLM tests @@ -399,7 +418,7 @@ def test_redownloads_when_fileset_changes(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -431,7 +450,7 @@ def test_skips_download_when_metadata_matches(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) dirs_to_keep: set[str] = set() controller._update_lora_adapters(dirs_to_keep) @@ -448,7 +467,7 @@ def test_downloads_new_adapter(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -477,7 +496,7 @@ def test_no_orphaned_temp_dirs_after_download(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -498,7 +517,7 @@ def test_failed_download_leaves_no_adapter_dir(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [] @@ -527,7 +546,7 @@ def test_failed_download_preserves_old_adapter(self, controller, tmp_path): mock_model_entity = MagicMock() mock_model_entity.workspace = "default" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [] @@ -556,7 +575,7 @@ def test_two_adapters_same_name_different_workspaces_coexist(self, controller, t mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter_a, adapter_b] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -578,7 +597,7 @@ def test_dir_name_uses_adapter_workspace_not_base_model_workspace(self, controll mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -613,7 +632,7 @@ def test_bare_fileset_for_cross_workspace_adapter_fetches_from_adapter_workspace mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -641,14 +660,14 @@ def test_step_gc_removes_stale_dir_after_adapter_workspace_change(self, controll mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] controller._sdk.files.list.return_value = mock_files_response # No prompt-tuned models in this scenario. - controller._sdk.models.list.return_value = [] + controller._models_client.list_models.return_value = _ItemsResponse([]) controller.step() @@ -670,7 +689,7 @@ def test_adapter_changed_meta_check_works_against_new_dir_path(self, controller, mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity + controller._models_client.get_model.return_value = _Response(mock_model_entity) dirs_to_keep: set[str] = set() controller._update_lora_adapters(dirs_to_keep) @@ -692,7 +711,7 @@ def _model_entity(self, controller, adapter): me = MagicMock() me.workspace = "default" me.adapters = [adapter] - controller._sdk.models.retrieve.return_value = me + controller._models_client.get_model.return_value = _Response(me) files_resp = MagicMock() files_resp.data = [MagicMock()] controller._sdk.files.list.return_value = files_resp @@ -823,7 +842,7 @@ def test_step_unloads_removed_adapter_before_delete(self, controller, tmp_path): adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._models_client.list_models.return_value = _ItemsResponse([]) # no prompt-tuned models with patch.object(controller, "_vllm_api_call", return_value=(200, "")) as api: controller.step() @@ -843,7 +862,7 @@ def test_step_keeps_removed_adapter_dir_when_vllm_unreachable(self, controller, adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._models_client.list_models.return_value = _ItemsResponse([]) # no prompt-tuned models # vLLM unreachable: both the kept adapter's load and the stale one's unload # hit a transport error. @@ -864,7 +883,7 @@ def test_step_deletes_removed_adapter_dir_when_vllm_answers_non_200(self, contro adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._models_client.list_models.return_value = _ItemsResponse([]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": @@ -887,7 +906,7 @@ def test_step_keeps_removed_adapter_dir_when_vllm_server_error(self, controller, adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._models_client.list_models.return_value = _ItemsResponse([]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": @@ -926,73 +945,15 @@ def test_unload_return_contract(self, controller, endpoint, api_result, expected assert controller._unload_vllm_adapter("default--x") is expected -class TestResolveAdapterWorkspaceFallback: - """Tests for the temporary ``Adapter.workspace`` SDK-schema gap. - - AALGO-117 introduces first-class :class:`Adapter` entities with their own - ``workspace`` in the entity store, but at the time of writing the public - SDK ``Adapter`` schema does not yet expose the field. AALGO-129 needs the - adapter workspace to encode the directory layout, so the sidecar falls - back to the base model's workspace until the SDK schema gains - ``workspace``. These tests pin both halves of that contract: the fallback - must engage on the current SDK shape, and the real value must take over - the moment the field becomes readable. - """ - - def test_fallback_uses_base_model_workspace_when_attribute_missing(self): - """Bare object with no ``workspace`` attribute: fall back to the base model workspace.""" - - class _AdapterWithoutWorkspace: - name = "my-adapter" - - adapter = _AdapterWithoutWorkspace() - assert AdaptersController._resolve_adapter_workspace(adapter, "base-ws") == "base-ws" - - def test_fallback_uses_base_model_workspace_when_attribute_is_none(self): - """``workspace=None`` (e.g. older payload deserialized via Optional[str]): fall back too.""" - adapter = MagicMock() - adapter.workspace = None - assert AdaptersController._resolve_adapter_workspace(adapter, "base-ws") == "base-ws" - - def test_fallback_uses_base_model_workspace_when_attribute_is_empty_string(self): - """``workspace=""`` is treated identically to ``None`` — an empty string can never form a valid - ``{ws}--{name}`` directory anchor, so the safest behavior is the same fallback as for ``None``. - """ - adapter = MagicMock() - adapter.workspace = "" - assert AdaptersController._resolve_adapter_workspace(adapter, "base-ws") == "base-ws" - - def test_explicit_workspace_takes_priority_over_base_model(self): - """Once the SDK schema exposes ``workspace``, the real value must be used verbatim.""" - adapter = MagicMock() - adapter.workspace = "adapter-ws" - assert AdaptersController._resolve_adapter_workspace(adapter, "base-ws") == "adapter-ws" - - def test_update_lora_adapters_falls_back_to_base_model_workspace(self, controller, tmp_path): - """End-to-end fallback path: an SDK Adapter without ``workspace`` is laid down under - ``{base_model_workspace}--{adapter_name}`` so the wire format stays decodable. - """ +class TestResolveAdapterWorkspace: + """Tests for the typed ``Adapter.workspace`` contract.""" - class _LegacyAdapter: - def __init__(self, name: str, fileset: str): - self.name = name - self.fileset = fileset - self.enabled = True - self.updated_at = None - - adapter = _LegacyAdapter("legacy-adapter", "base-ws/fs") - - mock_model_entity = MagicMock() - mock_model_entity.workspace = "base-ws" - mock_model_entity.adapters = [adapter] - controller._sdk.models.retrieve.return_value = mock_model_entity - - mock_files_response = MagicMock() - mock_files_response.data = [MagicMock()] - controller._sdk.files.list.return_value = mock_files_response - - dirs_to_keep: set[str] = set() - controller._update_lora_adapters(dirs_to_keep) + def test_returns_typed_adapter_workspace(self): + adapter = Adapter( + name="my-adapter", + workspace="adapter-ws", + fileset="adapter-ws/fileset", + finetuning_type=FinetuningType.LORA, + ) - assert dirs_to_keep == {"base-ws--legacy-adapter"} - assert (tmp_path / "base-ws--legacy-adapter").is_dir() + assert AdaptersController._resolve_adapter_workspace(adapter) == "adapter-ws" diff --git a/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py b/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py index aa36b48cf1..4d1d980dc1 100644 --- a/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py +++ b/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py @@ -10,8 +10,10 @@ import logging +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.models.client import ModelsClient from nmp.common.sdk_factory import get_platform_sdk -from nmp.guardrails.entities.values._private import RailsConfig +from nmp.guardrails.entities.values._private import ModelParameters, RailsConfig logger = logging.getLogger(__name__) @@ -59,7 +61,8 @@ def build_openai_gateway_url(model_entity_ref: str) -> str: # Use SDK helper to build IGW OpenAI-compatible URL # IGW handles routing the request to the correct Model Provider sdk = get_platform_sdk() - url = sdk.models.get_openai_route_base_url(workspace=workspace) + models = client_from_platform(sdk, ModelsClient) + url = models.get_openai_route_base_url(workspace=workspace) return url @@ -90,14 +93,14 @@ def resolve_model_entity_references(rails_config: RailsConfig) -> RailsConfig: model_ref = model.model parsed = parse_model_entity_reference(model_ref) - if parsed: + if parsed and model_ref is not None: # Resolve to IGW OpenAI-compatible URL gateway_url = build_openai_gateway_url(model_ref) # Set `parameters.base_url` - if model.parameters is None: - model.parameters = {} - model.parameters["base_url"] = gateway_url + parameters = ModelParameters.model_validate(model.parameters or {}) + parameters.base_url = gateway_url + model.parameters = parameters logger.debug(f"Resolved model '{model_ref}' to use Inference Gateway base URL: {gateway_url}") diff --git a/services/guardrails/tests/utils/test_config_utils.py b/services/guardrails/tests/utils/test_config_utils.py index cfe7b9cf53..ca777494fa 100644 --- a/services/guardrails/tests/utils/test_config_utils.py +++ b/services/guardrails/tests/utils/test_config_utils.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch import pytest +from nemo_platform_plugin.models.client import ModelsClient from nmp.guardrails.app.utils.config_utils import ( _load_and_execute_py_config, configure_rails_config, @@ -163,13 +164,18 @@ def test_enrich_config_with_data_returns_original(self): class TestConfigureRailsConfig: @pytest.fixture def mock_platform_sdk(self): - with patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") as mock: - mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( + with ( + patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") as get_sdk, + patch("nmp.guardrails.app.utils.model_routing.client_from_platform") as make_client, + ): + sdk = MagicMock() + models = MagicMock() + models.get_openai_route_base_url.return_value = ( "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" ) - mock.return_value = mock_sdk - yield mock + get_sdk.return_value = sdk + make_client.return_value = models + yield sdk, models, make_client def test_resolves_model_entity_references(self, mock_platform_sdk): """Test that configure_rails_config resolves Model Entity references.""" @@ -183,6 +189,9 @@ def test_resolves_model_entity_references(self, mock_platform_sdk): result = configure_rails_config(rails_config, model) + sdk, models, make_client = mock_platform_sdk + make_client.assert_called_once_with(sdk, ModelsClient) + models.get_openai_route_base_url.assert_called_once_with(workspace="default") main_model = result.models[0] assert ( main_model.parameters["base_url"] diff --git a/services/guardrails/tests/utils/test_model_routing.py b/services/guardrails/tests/utils/test_model_routing.py index dfc91c00d3..f6920f4b17 100644 --- a/services/guardrails/tests/utils/test_model_routing.py +++ b/services/guardrails/tests/utils/test_model_routing.py @@ -52,30 +52,21 @@ class TestBuildOpenAIGatewayUrl: def test_url_construction(self, mock_get_sdk): """Test URL construction for Model Entity reference.""" mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk url = build_openai_gateway_url("default/my-model") assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - mock_sdk.models.get_openai_route_base_url.assert_called_with(workspace="default") # Test non-default workspace - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/custom-workspace/openai/-/v1" - ) url = build_openai_gateway_url("custom-workspace/my-model") assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/custom-workspace/openai/-/v1" - mock_sdk.models.get_openai_route_base_url.assert_called_with(workspace="custom-workspace") @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") def test_v1_suffix_preserved(self, mock_get_sdk): - """Test /v1 suffix is preserved from SDK URL.""" + """Test /v1 suffix is preserved from typed client URL.""" mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk url = build_openai_gateway_url("default/model") @@ -83,17 +74,14 @@ def test_v1_suffix_preserved(self, mock_get_sdk): assert url.endswith("/v1") @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") - def test_url_without_v1_unchanged(self, mock_get_sdk): - """Test URL without /v1 suffix is returned unchanged.""" + def test_typed_client_adds_v1(self, mock_get_sdk): + """Test the typed Models client helper adds the OpenAI /v1 suffix.""" mock_sdk = MagicMock() - # Simulate SDK returning URL without /v1 (future-proofing) - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk url = build_openai_gateway_url("default/model") - assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-" + assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" def test_invalid_reference_raises(self): """Test invalid reference raises ValueError.""" @@ -108,9 +96,7 @@ class TestResolveModelEntityReferences: def test_resolve_single_model(self, mock_get_sdk): """Test resolving a single model with Model Entity reference.""" mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk rails_config = RailsConfig( @@ -129,9 +115,7 @@ def test_resolve_single_model(self, mock_get_sdk): def test_resolve_all_models(self, mock_get_sdk): """Test that ALL models in config get resolved (multiple models use case).""" mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk rails_config = RailsConfig( @@ -187,9 +171,7 @@ def test_explicit_base_url_preserved_for_entity_reference(self): def test_mixed_config(self, mock_get_sdk): """Test config with one Model Entity ref, one explicit URLs.""" mock_sdk = MagicMock() - mock_sdk.models.get_openai_route_base_url.return_value = ( - "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - ) + mock_sdk.base_url = "http://localhost:8000" mock_get_sdk.return_value = mock_sdk rails_config = RailsConfig( diff --git a/services/platform-seed/src/nmp/platform_seed/tasks/seed/run.py b/services/platform-seed/src/nmp/platform_seed/tasks/seed/run.py index 79bec73814..7a4f3a3500 100644 --- a/services/platform-seed/src/nmp/platform_seed/tasks/seed/run.py +++ b/services/platform-seed/src/nmp/platform_seed/tasks/seed/run.py @@ -10,7 +10,11 @@ from nemo_platform import AsyncNeMoPlatform from nemo_platform.resources.entities import AsyncEntitiesResource +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.errors import ConflictError from nemo_platform_plugin.discovery import discover_seed_jobs +from nemo_platform_plugin.models.client import AsyncModelsClient +from nemo_platform_plugin.models.types import CreateModelProviderRequest from nmp.common.config import get_platform_config from nmp.common.entities import EntityClient from nmp.common.sdk_factory import get_async_platform_sdk @@ -57,14 +61,15 @@ async def seed_auth(entity_client: EntityClient, config: PlatformSeedConfig) -> async def seed_model_provider(sdk: AsyncNeMoPlatform) -> None: """Seed the default nvidia-build model provider. Idempotent.""" - from nemo_platform import ConflictError - + models = client_from_platform(sdk, AsyncModelsClient) try: - await sdk.inference.providers.create( - name="nvidia-build", + await models.create_provider( workspace="system", - host_url="https://integrate.api.nvidia.com", - api_key_secret_name="ngc-api-key", + body=CreateModelProviderRequest( + name="nvidia-build", + host_url="https://integrate.api.nvidia.com", + api_key_secret_name="ngc-api-key", + ), ) logger.info("nvidia-build model provider created") except ConflictError: diff --git a/services/platform-seed/tests/test_platform_seed_runner.py b/services/platform-seed/tests/test_platform_seed_runner.py index be094e1a77..3c6a2f3133 100644 --- a/services/platform-seed/tests/test_platform_seed_runner.py +++ b/services/platform-seed/tests/test_platform_seed_runner.py @@ -6,8 +6,10 @@ from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest -from nemo_platform import ConflictError +from nemo_platform_plugin.client.errors import ConflictError +from nemo_platform_plugin.models.types import CreateModelProviderRequest from nmp.platform_seed.config import PlatformSeedConfig from nmp.platform_seed.tasks.seed import run_platform_seed @@ -166,13 +168,19 @@ async def test_run_platform_seed_auth_called(config_enabled, entity_client, sdk) async def test_run_platform_seed_model_provider_called(config_enabled, entity_client, sdk): """Model provider seed is called when model_provider_enabled is True.""" config_enabled.model_provider_enabled = True + models = MagicMock() + models.create_provider = AsyncMock() - result = await run_platform_seed(entity_client, sdk, config_enabled) - sdk.inference.providers.create.assert_awaited_once_with( - name="nvidia-build", + with patch("nmp.platform_seed.tasks.seed.run.client_from_platform", return_value=models): + result = await run_platform_seed(entity_client, sdk, config_enabled) + + models.create_provider.assert_awaited_once_with( workspace="system", - host_url="https://integrate.api.nvidia.com", - api_key_secret_name="ngc-api-key", + body=CreateModelProviderRequest( + name="nvidia-build", + host_url="https://integrate.api.nvidia.com", + api_key_secret_name="ngc-api-key", + ), ) assert result.models_ok is True assert result.errors == [] @@ -182,10 +190,12 @@ async def test_run_platform_seed_model_provider_called(config_enabled, entity_cl async def test_run_platform_seed_model_provider_conflict(config_enabled, entity_client, sdk): """ConflictError when creating provider is handled gracefully (not a failure).""" config_enabled.model_provider_enabled = True + models = MagicMock() + response = httpx.Response(409, request=httpx.Request("POST", "http://models/providers")) + models.create_provider = AsyncMock(side_effect=ConflictError(response)) + + with patch("nmp.platform_seed.tasks.seed.run.client_from_platform", return_value=models): + result = await run_platform_seed(entity_client, sdk, config_enabled) - sdk.inference.providers.create.side_effect = ConflictError( - message="already exists", response=MagicMock(), body=None - ) - result = await run_platform_seed(entity_client, sdk, config_enabled) assert result.models_ok is True assert result.errors == [] diff --git a/services/rl/src/nmp/rl/app/jobs/compiler.py b/services/rl/src/nmp/rl/app/jobs/compiler.py index df8cd12a70..a40f681124 100644 --- a/services/rl/src/nmp/rl/app/jobs/compiler.py +++ b/services/rl/src/nmp/rl/app/jobs/compiler.py @@ -22,7 +22,6 @@ import logging from nemo_platform import AsyncNeMoPlatform -from nemo_platform.types.models.model_entity import ModelEntity from nemo_platform_plugin.integrations import IntegrationsSpec from nemo_platform_plugin.jobs.api_factory import ( ContainerSpec, @@ -37,6 +36,7 @@ ResourcesSpec, ) from nemo_platform_plugin.jobs.exceptions import PlatformJobCompilationError +from nemo_platform_plugin.models.types import ModelEntity from nmp.common.jobs.constants import DEFAULT_JOB_STORAGE_PATH, PERSISTENT_JOB_STORAGE_PATH_ENVVAR from nmp.customization_common.integrations import ( collect_integration_secret_envs, diff --git a/services/rl/tests/test_compiler.py b/services/rl/tests/test_compiler.py index f4c784b19c..61f26330ec 100644 --- a/services/rl/tests/test_compiler.py +++ b/services/rl/tests/test_compiler.py @@ -12,9 +12,9 @@ import pytest from nemo_platform import AsyncNeMoPlatform -from nemo_platform.types.models.model_entity import ModelEntity from nemo_platform_plugin.integrations import IntegrationsSpec, MlflowIntegration, WandbIntegration from nemo_platform_plugin.jobs.exceptions import PlatformJobCompilationError +from nemo_platform_plugin.models.types import ModelEntity from nmp.common.entities.utils import get_random_id from nmp.rl.app.jobs.compiler import ( _build_training_step, diff --git a/services/unsloth/src/nmp/unsloth/app/jobs/compiler.py b/services/unsloth/src/nmp/unsloth/app/jobs/compiler.py index 63b1bffea6..c125f82160 100644 --- a/services/unsloth/src/nmp/unsloth/app/jobs/compiler.py +++ b/services/unsloth/src/nmp/unsloth/app/jobs/compiler.py @@ -17,7 +17,6 @@ import logging from nemo_platform import AsyncNeMoPlatform -from nemo_platform.types.models.model_entity import ModelEntity from nemo_platform_plugin.jobs.api_factory import ( ContainerSpec, CPUExecutionProviderSpec, @@ -29,6 +28,7 @@ ResourcesSpec, ) from nemo_platform_plugin.jobs.exceptions import PlatformJobCompilationError +from nemo_platform_plugin.models.types import ModelEntity from nmp.common.jobs.constants import DEFAULT_JOB_STORAGE_PATH, PERSISTENT_JOB_STORAGE_PATH_ENVVAR from nmp.customization_common.schemas.file_io import ( DownloadItem, diff --git a/services/unsloth/tests/test_compiler_validation_path.py b/services/unsloth/tests/test_compiler_validation_path.py index bf85e6491e..02241f5e1b 100644 --- a/services/unsloth/tests/test_compiler_validation_path.py +++ b/services/unsloth/tests/test_compiler_validation_path.py @@ -5,10 +5,11 @@ from __future__ import annotations -import types +from datetime import datetime from unittest.mock import AsyncMock, MagicMock import pytest +from nemo_platform_plugin.models.types import ModelEntity from nmp.unsloth.app.constants import DEFAULT_DATASET_PATH, DEFAULT_VALIDATION_DATASET_PATH from nmp.unsloth.app.jobs.compiler import platform_job_config_compiler from nmp.unsloth.schemas import ( @@ -22,6 +23,18 @@ ) +def _model_entity() -> ModelEntity: + return ModelEntity( + id="model-qwen3", + workspace="default", + name="qwen3-1.7b", + fileset="default/qwen3-1.7b", + trust_remote_code=False, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + def _spec(*, validation_path: str | None) -> UnslothJobOutput: return UnslothJobOutput( model=ModelLoadSpec(name="default/qwen3-1.7b"), @@ -46,14 +59,7 @@ async def test_training_step_gets_local_validation_path_for_same_fileset() -> No from nmp.unsloth.app.jobs import compiler as compiler_mod original_fetch = compiler_mod.fetch_model_entity - compiler_mod.fetch_model_entity = AsyncMock( - return_value=types.SimpleNamespace( - workspace="default", - name="qwen3-1.7b", - fileset="default/qwen3-1.7b", - trust_remote_code=False, - ), - ) + compiler_mod.fetch_model_entity = AsyncMock(return_value=_model_entity()) try: job = await platform_job_config_compiler( workspace="default", @@ -75,14 +81,7 @@ async def test_training_step_gets_separate_validation_path_for_different_fileset from nmp.unsloth.app.jobs import compiler as compiler_mod original_fetch = compiler_mod.fetch_model_entity - compiler_mod.fetch_model_entity = AsyncMock( - return_value=types.SimpleNamespace( - workspace="default", - name="qwen3-1.7b", - fileset="default/qwen3-1.7b", - trust_remote_code=False, - ), - ) + compiler_mod.fetch_model_entity = AsyncMock(return_value=_model_entity()) try: job = await platform_job_config_compiler( workspace="default", @@ -104,14 +103,7 @@ async def test_upload_step_stamps_output_metadata() -> None: from nmp.unsloth.app.jobs import compiler as compiler_mod original_fetch = compiler_mod.fetch_model_entity - compiler_mod.fetch_model_entity = AsyncMock( - return_value=types.SimpleNamespace( - workspace="default", - name="qwen3-1.7b", - fileset="default/qwen3-1.7b", - trust_remote_code=False, - ), - ) + compiler_mod.fetch_model_entity = AsyncMock(return_value=_model_entity()) try: job = await platform_job_config_compiler( workspace="default", diff --git a/services/unsloth/tests/test_model_entity.py b/services/unsloth/tests/test_model_entity.py index d5231e19d3..854f479bf7 100644 --- a/services/unsloth/tests/test_model_entity.py +++ b/services/unsloth/tests/test_model_entity.py @@ -14,11 +14,27 @@ from __future__ import annotations +import json import types +from datetime import datetime from pathlib import Path from unittest.mock import MagicMock, patch import pytest +from nemo_platform_plugin.files.client import FilesClient +from nemo_platform_plugin.models.client import ModelsClient +from nemo_platform_plugin.models.types import ( + CreateModelAdapterRequest, + CreateModelDeploymentConfigRequest, + CreateModelDeploymentRequest, + CreateModelEntityRequest, + Engine, + ModelDeploymentStatus, + ModelEntity, + UpdateAdapterRequest, + UpdateModelDeploymentConfigRequest, + UpdateModelEntityRequest, +) def _make_job_ctx(workspace: str = "default"): @@ -49,6 +65,33 @@ def _make_sdk() -> MagicMock: return sdk +def _response(data: object) -> MagicMock: + response = MagicMock() + response.data.return_value = data + return response + + +def _page(items: list[object]) -> MagicMock: + response = MagicMock() + response.items.return_value = items + return response + + +def _configure_clients(mock_client_from_platform: MagicMock) -> tuple[MagicMock, MagicMock]: + models = MagicMock() + files = MagicMock() + + def make_client(_sdk: object, client_type: type) -> MagicMock: + if client_type is ModelsClient: + return models + if client_type is FilesClient: + return files + raise AssertionError(f"Unexpected client type: {client_type}") + + mock_client_from_platform.side_effect = make_client + return models, files + + def _raise_runner_conflict() -> None: """Raise the ``ConflictError`` class the runner is bound against. @@ -71,6 +114,18 @@ def _model_entity(*, workspace: str = "default", name: str = "base", spec: objec return me +def _compiler_model_entity(*, workspace: str = "default", name: str = "base") -> ModelEntity: + return ModelEntity( + id=f"model-{name}", + workspace=workspace, + name=name, + fileset=f"{workspace}/{name}-fileset", + trust_remote_code=False, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + # --------------------------------------------------------------------------- # sanitize_name # --------------------------------------------------------------------------- @@ -111,11 +166,10 @@ def test_creates_model_entity_for_full_sft(self, mock_cfp) -> None: from nmp.customization_common.schemas.model_entity import ModelEntityTaskConfig sdk = _make_sdk() - mock_fc = MagicMock() - mock_cfp.return_value = mock_fc - sdk.models.retrieve.return_value = _model_entity(name="base-model") + models, files = _configure_clients(mock_cfp) + models.get_model.return_value = _response(_model_entity(name="base-model")) new_me = _model_entity(name="trained-model") - sdk.models.create.return_value = new_me + models.create_model.return_value = _response(new_me) runner = _make_runner(sdk) config = ModelEntityTaskConfig( @@ -128,8 +182,17 @@ def test_creates_model_entity_for_full_sft(self, mock_cfp) -> None: result, deploy_target = runner.create_model_entity(config) - mock_fc.get_fileset.assert_called_once_with(workspace="default", name="trained-model") - sdk.models.create.assert_called_once() + files.get_fileset.assert_called_once_with(workspace="default", name="trained-model") + models.get_model.assert_called_once_with(name="base-model", workspace="default") + models.create_model.assert_called_once() + create_call = models.create_model.call_args + assert create_call.kwargs["workspace"] == "default" + body = create_call.kwargs["body"] + assert isinstance(body, CreateModelEntityRequest) + assert body.name == "trained-model" + assert body.fileset == "default/trained-model" + assert body.base_model is None + assert body.trust_remote_code is False assert deploy_target is new_me assert result is not None @@ -139,10 +202,11 @@ def test_conflict_falls_back_to_update(self, mock_cfp) -> None: from nmp.customization_common.schemas.model_entity import ModelEntityTaskConfig sdk = _make_sdk() - mock_cfp.return_value = MagicMock() - sdk.models.retrieve.return_value = _model_entity(name="base-model") - sdk.models.create.side_effect = lambda **_: _raise_runner_conflict() - sdk.models.update.return_value = _model_entity(name="trained-model") + models, _files = _configure_clients(mock_cfp) + models.get_model.return_value = _response(_model_entity(name="base-model")) + models.create_model.side_effect = lambda **_: _raise_runner_conflict() + updated_me = _model_entity(name="trained-model") + models.update_model.return_value = _response(updated_me) runner = _make_runner(sdk) config = ModelEntityTaskConfig( @@ -155,10 +219,15 @@ def test_conflict_falls_back_to_update(self, mock_cfp) -> None: _, _ = runner.create_model_entity(config) - sdk.models.update.assert_called_once() - update_call = sdk.models.update.call_args + models.update_model.assert_called_once() + update_call = models.update_model.call_args assert update_call.kwargs["name"] == "trained-model" assert update_call.kwargs["workspace"] == "default" + body = update_call.kwargs["body"] + assert isinstance(body, UpdateModelEntityRequest) + assert body.fileset == "default/trained-model" + assert body.base_model is None + assert body.trust_remote_code is False @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") def test_missing_fileset_raises_creation_error(self, mock_cfp) -> None: @@ -166,9 +235,8 @@ def test_missing_fileset_raises_creation_error(self, mock_cfp) -> None: from nmp.customization_common.schemas.model_entity import ModelEntityCreationError, ModelEntityTaskConfig sdk = _make_sdk() - mock_fc = MagicMock() - mock_fc.get_fileset.side_effect = RuntimeError("fileset missing") - mock_cfp.return_value = mock_fc + models, files = _configure_clients(mock_cfp) + files.get_fileset.side_effect = RuntimeError("fileset missing") runner = _make_runner(sdk) config = ModelEntityTaskConfig( name="x", @@ -180,6 +248,9 @@ def test_missing_fileset_raises_creation_error(self, mock_cfp) -> None: with pytest.raises(ModelEntityCreationError, match="does not exist or is not accessible"): runner.create_model_entity(config) + models.get_model.assert_not_called() + models.create_model.assert_not_called() + # --------------------------------------------------------------------------- # ModelEntityRunner.create_model_entity — LoRA adapter path @@ -194,10 +265,10 @@ def test_creates_adapter_for_lora(self, mock_cfp) -> None: from nmp.unsloth.entities.values import FinetuningType sdk = _make_sdk() - mock_cfp.return_value = MagicMock() + models, _files = _configure_clients(mock_cfp) base_me = _model_entity(name="base-model") - sdk.models.retrieve.return_value = base_me - sdk.models.adapters.create.return_value = _model_entity(name="adapter-x") + models.get_model.return_value = _response(base_me) + models.create_model_adapter.return_value = _response(_model_entity(name="adapter-x")) runner = _make_runner(sdk) config = ModelEntityTaskConfig( @@ -210,7 +281,18 @@ def test_creates_adapter_for_lora(self, mock_cfp) -> None: _result, deploy_target = runner.create_model_entity(config) - sdk.models.adapters.create.assert_called_once() + models.create_model_adapter.assert_called_once() + create_call = models.create_model_adapter.call_args + assert create_call.kwargs["model_name"] == "base-model" + assert create_call.kwargs["workspace"] == "default" + body = create_call.kwargs["body"] + assert isinstance(body, CreateModelAdapterRequest) + assert body.name == "adapter-x" + assert body.fileset == "default/adapter-x" + assert body.lora_config is not None + assert body.lora_config.rank == 8 + assert body.lora_config.alpha == 16 + assert body.enabled is True assert deploy_target is base_me @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") @@ -220,10 +302,10 @@ def test_adapter_conflict_falls_back_to_update(self, mock_cfp) -> None: from nmp.unsloth.entities.values import FinetuningType sdk = _make_sdk() - mock_cfp.return_value = MagicMock() - sdk.models.retrieve.return_value = _model_entity(name="base-model") - sdk.models.adapters.create.side_effect = lambda **_: _raise_runner_conflict() - sdk.models.adapters.update.return_value = _model_entity(name="adapter-x") + models, _files = _configure_clients(mock_cfp) + models.get_model.return_value = _response(_model_entity(name="base-model")) + models.create_model_adapter.side_effect = lambda **_: _raise_runner_conflict() + models.update_model_adapter.return_value = _response(_model_entity(name="adapter-x")) runner = _make_runner(sdk) config = ModelEntityTaskConfig( @@ -236,7 +318,15 @@ def test_adapter_conflict_falls_back_to_update(self, mock_cfp) -> None: runner.create_model_entity(config) - sdk.models.adapters.update.assert_called_once() + models.update_model_adapter.assert_called_once() + update_call = models.update_model_adapter.call_args + assert update_call.kwargs["adapter"] == "adapter-x" + assert update_call.kwargs["model_name"] == "base-model" + assert update_call.kwargs["workspace"] == "default" + body = update_call.kwargs["body"] + assert isinstance(body, UpdateAdapterRequest) + assert body.fileset == "default/adapter-x" + assert body.enabled is True # --------------------------------------------------------------------------- @@ -245,11 +335,13 @@ def test_adapter_conflict_falls_back_to_update(self, mock_cfp) -> None: class TestLaunchModel: - def test_no_deployment_config_returns_early(self) -> None: + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") + def test_no_deployment_config_returns_early(self, mock_cfp) -> None: from nmp.customization_common.schemas.file_io import FileSetRef from nmp.customization_common.schemas.model_entity import ModelEntityTaskConfig sdk = _make_sdk() + models, _files = _configure_clients(mock_cfp) runner = _make_runner(sdk) me = _model_entity(name="x") config = ModelEntityTaskConfig( @@ -262,60 +354,124 @@ def test_no_deployment_config_returns_early(self) -> None: runner.launch_model(config, me) - sdk.inference.deployments.create.assert_not_called() - sdk.inference.deployment_configs.create.assert_not_called() + models.create_deployment.assert_not_called() + models.create_deployment_config.assert_not_called() - def test_inline_params_creates_config_then_deployment(self) -> None: + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") + def test_inline_params_creates_config_then_deployment(self, mock_cfp) -> None: from nmp.customization_common.schemas.file_io import FileSetRef from nmp.customization_common.schemas.model_entity import DeploymentParameters, ModelEntityTaskConfig sdk = _make_sdk() - sdk.inference.deployment_configs.create.return_value = types.SimpleNamespace( - workspace="default", - name="sft-cfg-x", - ) - sdk.inference.deployments.create.return_value = types.SimpleNamespace( - workspace="default", + models, _files = _configure_clients(mock_cfp) + deployment_config = types.SimpleNamespace(workspace="other", name="sft-cfg-x") + deployment = types.SimpleNamespace(workspace="other", name="sft-deploy-x") + deployment_status = types.SimpleNamespace( + workspace="other", name="sft-deploy-x", + status=ModelDeploymentStatus.PENDING, ) - sdk.inference.deployments.retrieve.return_value = types.SimpleNamespace( - workspace="default", - name="sft-deploy-x", - status="PENDING", + models.create_deployment_config.return_value = _response(deployment_config) + models.create_deployment.return_value = _response(deployment) + models.get_deployment.return_value = _response(deployment_status) + + runner = _make_runner(sdk) + me = _model_entity( + workspace="other", + name="x", + spec=types.SimpleNamespace(family="llama", base_num_parameters=1_000_000_000), + ) + config = ModelEntityTaskConfig( + name="x", + workspace="other", + fileset=FileSetRef(workspace="other", name="x"), + model_entity="other/base", + deployment_config=DeploymentParameters(gpu=1, image_name="img", image_tag="1.0"), + ) + + runner.launch_model(config, me) + + config_call = models.create_deployment_config.call_args + assert config_call.kwargs["workspace"] == "other" + config_body = config_call.kwargs["body"] + assert isinstance(config_body, CreateModelDeploymentConfigRequest) + assert config_body.name == "sft-cfg-x" + assert config_body.engine is Engine.NIM + assert config_body.model_spec.model_name == "x" + assert config_body.model_spec.model_namespace == "other" + assert config_body.executor_config.gpu == 1 + assert config_body.executor_config.image_name == "img" + assert config_body.executor_config.image_tag == "1.0" + + deployment_call = models.create_deployment.call_args + assert deployment_call.kwargs["workspace"] == "other" + deployment_body = deployment_call.kwargs["body"] + assert isinstance(deployment_body, CreateModelDeploymentRequest) + assert deployment_body.name == "sft-deploy-x" + assert deployment_body.config == "sft-cfg-x" + models.get_deployment.assert_called_once_with(workspace="other", name="sft-deploy-x") + + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") + def test_inline_config_conflict_updates_before_deployment(self, mock_cfp) -> None: + from nmp.customization_common.schemas.file_io import FileSetRef + from nmp.customization_common.schemas.model_entity import DeploymentParameters, ModelEntityTaskConfig + + sdk = _make_sdk() + models, _files = _configure_clients(mock_cfp) + models.create_deployment_config.side_effect = lambda **_: _raise_runner_conflict() + updated_config = types.SimpleNamespace(workspace="default", name="sft-cfg-x") + deployment = types.SimpleNamespace(workspace="default", name="sft-deploy-x") + models.update_deployment_config.return_value = _response(updated_config) + models.create_deployment.return_value = _response(deployment) + models.get_deployment.return_value = _response( + types.SimpleNamespace( + workspace="default", + name="sft-deploy-x", + status=ModelDeploymentStatus.PENDING, + ) ) runner = _make_runner(sdk) - me = _model_entity(name="x", spec=types.SimpleNamespace(family="llama", base_num_parameters=1_000_000_000)) + me = _model_entity( + name="x", + spec=types.SimpleNamespace(family="llama", base_num_parameters=1_000_000_000), + ) config = ModelEntityTaskConfig( name="x", workspace="default", fileset=FileSetRef(workspace="default", name="x"), model_entity="default/base", - deployment_config=DeploymentParameters(gpu=1, image_name="img", image_tag="1.0"), + deployment_config=DeploymentParameters(gpu=2), ) runner.launch_model(config, me) - sdk.inference.deployment_configs.create.assert_called_once() - sdk.inference.deployments.create.assert_called_once() + update_call = models.update_deployment_config.call_args + assert update_call.kwargs["workspace"] == "default" + assert update_call.kwargs["name"] == "sft-cfg-x" + body = update_call.kwargs["body"] + assert isinstance(body, UpdateModelDeploymentConfigRequest) + assert body.engine is Engine.NIM + assert body.executor_config.gpu == 2 + models.create_deployment.assert_called_once() - def test_string_ref_resolves_existing_config(self) -> None: + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") + def test_string_ref_resolves_existing_config(self, mock_cfp) -> None: from nmp.customization_common.schemas.file_io import FileSetRef from nmp.customization_common.schemas.model_entity import ModelEntityTaskConfig sdk = _make_sdk() - sdk.inference.deployment_configs.retrieve.return_value = types.SimpleNamespace( - workspace="default", - name="existing-cfg", - ) - sdk.inference.deployments.create.return_value = types.SimpleNamespace( - workspace="default", - name="sft-deploy-x", - ) - sdk.inference.deployments.retrieve.return_value = types.SimpleNamespace( - workspace="default", - name="sft-deploy-x", - status="PENDING", + models, _files = _configure_clients(mock_cfp) + deployment_config = types.SimpleNamespace(workspace="shared", name="existing-cfg") + deployment = types.SimpleNamespace(workspace="shared", name="sft-deploy-x") + models.get_deployment_config.return_value = _response(deployment_config) + models.create_deployment.return_value = _response(deployment) + models.get_deployment.return_value = _response( + types.SimpleNamespace( + workspace="shared", + name="sft-deploy-x", + status=ModelDeploymentStatus.PENDING, + ) ) runner = _make_runner(sdk) @@ -325,19 +481,19 @@ def test_string_ref_resolves_existing_config(self) -> None: workspace="default", fileset=FileSetRef(workspace="default", name="x"), model_entity="default/base", - deployment_config="existing-cfg", + deployment_config="shared/existing-cfg", ) runner.launch_model(config, me) - sdk.inference.deployment_configs.retrieve.assert_called_once_with( - workspace="default", - name="existing-cfg", - ) - sdk.inference.deployment_configs.create.assert_not_called() - sdk.inference.deployments.create.assert_called_once() + models.get_deployment_config.assert_called_once_with(workspace="shared", name="existing-cfg") + models.create_deployment_config.assert_not_called() + deployment_call = models.create_deployment.call_args + assert deployment_call.kwargs["workspace"] == "shared" + assert deployment_call.kwargs["body"].config == "existing-cfg" - def test_lora_with_active_deployment_skips(self) -> None: + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") + def test_lora_with_active_deployment_skips(self, mock_cfp) -> None: from nmp.customization_common.schemas.file_io import FileSetRef from nmp.customization_common.schemas.model_entity import ( DeploymentParameters, @@ -347,30 +503,41 @@ def test_lora_with_active_deployment_skips(self) -> None: from nmp.unsloth.entities.values import FinetuningType sdk = _make_sdk() - # Active deployment exists → launch_model should return without creating anything. - existing_config = types.SimpleNamespace(workspace="default", name="cfg-1") - active_deployment = types.SimpleNamespace(status="READY") - sdk.inference.deployment_configs.list.return_value = types.SimpleNamespace(data=[existing_config]) - sdk.inference.deployments.list.return_value = types.SimpleNamespace(data=[active_deployment]) + models, _files = _configure_clients(mock_cfp) + existing_config = types.SimpleNamespace(workspace="other", name="cfg-1") + active_deployment = types.SimpleNamespace(status=ModelDeploymentStatus.READY) + models.list_deployment_configs.return_value = _page([existing_config]) + models.list_deployments.return_value = _page([active_deployment]) runner = _make_runner(sdk) - me = _model_entity(name="base") + me = _model_entity(workspace="other", name="base") config = ModelEntityTaskConfig( name="adapter", - workspace="default", - fileset=FileSetRef(workspace="default", name="adapter"), - model_entity="default/base", + workspace="other", + fileset=FileSetRef(workspace="other", name="adapter"), + model_entity="other/base", peft=PEFTConfig(type=FinetuningType.LORA, rank=8, alpha=16), deployment_config=DeploymentParameters(), ) runner.launch_model(config, me) - sdk.inference.deployment_configs.create.assert_not_called() - sdk.inference.deployments.create.assert_not_called() + config_call = models.list_deployment_configs.call_args + assert config_call.kwargs["workspace"] == "other" + assert json.loads(config_call.kwargs["query_params"]["filter"]) == {"model_entity_id": "other/base"} + deployment_call = models.list_deployments.call_args + assert deployment_call.kwargs["workspace"] == "other" + assert json.loads(deployment_call.kwargs["query_params"]["filter"]) == { + "config": "cfg-1", + "workspace": "other", + } + models.create_deployment_config.assert_not_called() + models.create_deployment.assert_not_called() + @patch("nmp.customization_common.tasks.model_entity.run.client_from_platform") def test_lora_with_lora_enabled_false_warns_and_skips( self, + mock_cfp, caplog: pytest.LogCaptureFixture, ) -> None: from nmp.customization_common.schemas.file_io import FileSetRef @@ -382,7 +549,8 @@ def test_lora_with_lora_enabled_false_warns_and_skips( from nmp.unsloth.entities.values import FinetuningType sdk = _make_sdk() - sdk.inference.deployment_configs.list.return_value = types.SimpleNamespace(data=[]) + models, _files = _configure_clients(mock_cfp) + models.list_deployment_configs.return_value = _page([]) runner = _make_runner(sdk) me = _model_entity(name="base") @@ -399,7 +567,7 @@ def test_lora_with_lora_enabled_false_warns_and_skips( runner.launch_model(config, me) assert any("lora_enabled is false" in r.getMessage() for r in caplog.records) - sdk.inference.deployments.create.assert_not_called() + models.create_deployment.assert_not_called() # --------------------------------------------------------------------------- @@ -437,14 +605,7 @@ async def test_inline_params_pass_through_to_model_entity_step(self) -> None: from nmp.unsloth.app.jobs import compiler as compiler_mod original_fetch = compiler_mod.fetch_model_entity - compiler_mod.fetch_model_entity = AsyncMock( - return_value=types.SimpleNamespace( - workspace="default", - name="base", - fileset="default/base-fileset", - trust_remote_code=False, - ) - ) + compiler_mod.fetch_model_entity = AsyncMock(return_value=_compiler_model_entity()) try: job_spec = await platform_job_config_compiler( workspace="default", @@ -489,14 +650,7 @@ async def test_string_ref_passes_through_unchanged(self) -> None: from nmp.unsloth.app.jobs import compiler as compiler_mod original_fetch = compiler_mod.fetch_model_entity - compiler_mod.fetch_model_entity = AsyncMock( - return_value=types.SimpleNamespace( - workspace="default", - name="base", - fileset="default/base-fileset", - trust_remote_code=False, - ) - ) + compiler_mod.fetch_model_entity = AsyncMock(return_value=_compiler_model_entity()) try: job_spec = await platform_job_config_compiler( workspace="default",