diff --git a/e2e/backends/docker_compose.py b/e2e/backends/docker_compose.py index a6f3345d66..6cdfcd2844 100644 --- a/e2e/backends/docker_compose.py +++ b/e2e/backends/docker_compose.py @@ -11,10 +11,11 @@ from typing import Literal, TextIO import httpx -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, httpx_tls_config_from_env ComposeLifecycle = Literal["fresh", "reuse"] _DIAGNOSTIC_COMMAND_TIMEOUT_SECONDS = 60 +_CA_BUNDLE_ENVVARS = (NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "REQUESTS_CA_BUNDLE", "SSL_CERT_FILE") def _compose_env(env: dict[str, str] | None) -> dict[str, str]: @@ -164,19 +165,14 @@ def start(self) -> None: self._wait_ready() def _wait_ready(self) -> None: - verify = ( - self.env.get(NMP_CLIENT_SSL_CERT_FILE_ENVVAR) - or self.env.get("REQUESTS_CA_BUNDLE") - or self.env.get("SSL_CERT_FILE") - or True - ) + tls_config = httpx_tls_config_from_env(self.env, cert_file_envvars=_CA_BUNDLE_ENVVARS) deadline = time.monotonic() + self.wait_timeout_seconds pending = list(dict.fromkeys(self.wait_urls)) last_results: dict[str, str] = {} while time.monotonic() < deadline and pending: for wait_url in list(pending): try: - response = httpx.get(wait_url, timeout=5, verify=verify) + response = httpx.get(wait_url, timeout=5, **tls_config) if response.status_code == 200: pending.remove(wait_url) else: diff --git a/e2e/services_pool.py b/e2e/services_pool.py index 4b0803f337..feff078762 100644 --- a/e2e/services_pool.py +++ b/e2e/services_pool.py @@ -24,6 +24,7 @@ import pytest import yaml from _pytest.nodes import Node +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, httpx_tls_config_from_env from nmp.testing.e2e import Docker as DockerE2EBackend from nmp.testing.e2e.config import deep_merge @@ -42,6 +43,7 @@ # deployments orphan cleanup cannot delete peer platforms' docker containers. _DEFAULT_E2E_DISABLE_DEPLOYMENTS_ORPHAN_CLEANUP = _E2E_REPO_ROOT / "e2e/configs/disable-deployments-orphan-cleanup.yaml" _E2E_COMPOSE_LIFECYCLE_ENV = "NMP_E2E_COMPOSE_LIFECYCLE" +_CA_BUNDLE_ENVVARS = (NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "REQUESTS_CA_BUNDLE", "SSL_CERT_FILE") def admin_headers() -> dict[str, str]: @@ -780,18 +782,6 @@ def _wait_for_auth_ready(url: str, proc: subprocess.Popen[Any] | None, timeout: return False -def _request_verify_from_env(env: Mapping[str, str] | None = None) -> str | bool: - source = dict(os.environ) - if env: - source.update(env) - return ( - source.get("NMP_CLIENT_SSL_CERT_FILE") - or source.get("REQUESTS_CA_BUNDLE") - or source.get("SSL_CERT_FILE") - or True - ) - - def _wait_for_auth_ready_url( url: str, proc: subprocess.Popen[Any] | None, @@ -800,12 +790,12 @@ def _wait_for_auth_ready_url( timeout: float = _AUTH_READY_TIMEOUT, ) -> bool: deadline = time.monotonic() + timeout - verify = _request_verify_from_env(env) + tls_config = httpx_tls_config_from_env(env, cert_file_envvars=_CA_BUNDLE_ENVVARS) while time.monotonic() < deadline: if proc is not None and _process_exited(proc): return False try: - response = httpx.get(url, timeout=5.0, verify=verify) + response = httpx.get(url, timeout=5.0, **tls_config) if response.status_code == 200: return True except httpx.RequestError as exc: diff --git a/packages/nemo_platform/pyproject.toml b/packages/nemo_platform/pyproject.toml index 2ebb781296..199e59bcde 100644 --- a/packages/nemo_platform/pyproject.toml +++ b/packages/nemo_platform/pyproject.toml @@ -665,10 +665,6 @@ customization = "nemo_customizer.cli:CustomizationCLI" experimentalist = "nemo_experimentalist_plugin.cli:ExperimentalistCLI" analyst = "nemo_insights_plugin.analyst.cli:AnalystCLI" -# Generated from [tool.bundle-package]; do not edit this table by hand. -[project.entry-points."nemo.client_provider"] -platform = "nmp.common.client_factory:PlatformNemoClientProvider" - # Generated from [tool.bundle-package]; do not edit this table by hand. [project.entry-points."nemo.controllers"] agents-deployment = "nemo_agents_plugin.runner.controller:AgentDeploymentController" diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/device_flow.py b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/device_flow.py index d45ed21ea5..e05dda02d3 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/device_flow.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/device_flow.py @@ -9,11 +9,11 @@ from dataclasses import dataclass import httpx +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from rich.console import Console from rich.panel import Panel from nemo_platform_ext.auth.token_provider import refresh_token_grant -from nemo_platform_ext.client.tls import client_verify_from_env console = Console() @@ -74,7 +74,7 @@ def __init__( async def start_device_authorization(self) -> DeviceCodeResponse: """Start the device authorization flow.""" - async with httpx.AsyncClient(verify=client_verify_from_env()) as client: + async with httpx.AsyncClient(**httpx_tls_config_from_env()) as client: response = await client.post( self.device_authorization_endpoint, data={ @@ -104,7 +104,7 @@ async def poll_for_token( """Poll the token endpoint until authorization is complete.""" start_time = time.time() - async with httpx.AsyncClient(verify=client_verify_from_env()) as client: + async with httpx.AsyncClient(**httpx_tls_config_from_env()) as client: while time.time() - start_time < expires_in: await _async_pause(interval) @@ -284,7 +284,7 @@ def authenticate_with_password_grant( "password": password, "scope": scope, } - with httpx.Client(verify=client_verify_from_env()) as client: + with httpx.Client(**httpx_tls_config_from_env()) as client: response = client.post(token_endpoint, data=data, timeout=30.0) if response.status_code != 200: diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/helpers.py b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/helpers.py index 9c3281afe0..730810e3d7 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/helpers.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/helpers.py @@ -24,8 +24,7 @@ from typing import Any import httpx - -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env DEFAULT_OAUTH_SCOPES = "openid profile email offline_access" @@ -157,7 +156,7 @@ def discover_nmp_config(base_url: str, timeout: float = 10.0) -> NMPOIDCConfig: response = httpx.get( f"{base_url.rstrip('/')}/apis/auth/discovery", timeout=timeout, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) response.raise_for_status() data = response.json() diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/token_provider.py b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/token_provider.py index bcc7b7c5c8..c049c0ccfc 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/token_provider.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/token_provider.py @@ -13,10 +13,10 @@ from dataclasses import dataclass, field import httpx +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from typing_extensions import Self from nemo_platform_ext.auth.helpers import decode_jwt_claims -from nemo_platform_ext.client.tls import client_verify_from_env logger = logging.getLogger(__name__) @@ -56,7 +56,7 @@ def refresh_token_grant( if scope: data["scope"] = scope - response = httpx.post(token_endpoint, data=data, timeout=timeout, verify=client_verify_from_env()) + response = httpx.post(token_endpoint, data=data, timeout=timeout, **httpx_tls_config_from_env()) if response.status_code != 200: error_data: dict[str, str] = {} diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py index bdffa3cf54..400e3b344e 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py @@ -24,9 +24,9 @@ WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, subject_token_type_for_exchange, ) +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from nemo_platform_ext.auth.token_provider import DEFAULT_REFRESH_MARGIN_SECONDS, TokenSet -from nemo_platform_ext.client.tls import client_verify_from_env logger = logging.getLogger(__name__) @@ -103,7 +103,7 @@ def token_exchange_grant( if scope: data["scope"] = scope - response = httpx.post(token_endpoint, data=data, timeout=timeout, verify=client_verify_from_env()) + response = httpx.post(token_endpoint, data=data, timeout=timeout, **httpx_tls_config_from_env()) if response.status_code != 200: error_data: dict[str, object] = {} diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/setup.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/setup.py index 2901c2d2af..bf86dc3958 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/setup.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/setup.py @@ -29,6 +29,7 @@ from nemo_platform import NeMoPlatform from nemo_platform_plugin.capabilities import probe_docker from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.tls import HttpxTLSConfig, httpx_tls_config_from_env from nemo_platform_plugin.secrets.client import SecretsClient from nemo_platform_plugin.secrets.types import PlatformSecretCreateRequest, PlatformSecretUpdateRequest from nemo_platform_plugin.workspaces.client import WorkspacesClient @@ -49,7 +50,6 @@ from nemo_platform_ext.cli.docker_preflight import DOCKER_PREFLIGHT_MESSAGE, require_docker_for_default_local from nemo_platform_ext.cli.telemetry import emit from nemo_platform_ext.cli.telemetry.events import OnboardingStepEvent, TaskStatusEnum -from nemo_platform_ext.client.tls import client_verify_from_env from nemo_platform_ext.config.config import Config from nemo_platform_ext.config.models import DEFAULT_BASE_URL, ConfigFile, ConfigParams, LocalServicesConfig, NoAuthUser from nemo_platform_ext.local.install import services_extra_install_command @@ -323,11 +323,11 @@ def _check_platform_reachable(base_url: str, timeout: float = 5.0) -> bool: Local ``nemo services run`` publishes ``/status``. Hosted deployments may only expose ``/cluster-info`` on ingress, so try both. """ - verify = client_verify_from_env() + tls_config = httpx_tls_config_from_env() root = base_url.rstrip("/") for path in _PLATFORM_REACHABILITY_PATHS: try: - resp = httpx.get(f"{root}{path}", timeout=timeout, verify=verify) + resp = httpx.get(f"{root}{path}", timeout=timeout, **tls_config) if resp.status_code == 200: return True except Exception: @@ -466,10 +466,10 @@ def _platform_request_headers(cli_context: CLIContext) -> dict[str, str] | None: return {key: value for key, value in headers.items() if isinstance(key, str) and isinstance(value, str)} -def _hosted_platform_without_status(base_url: str, *, timeout: float, verify: str | bool) -> bool: +def _hosted_platform_without_status(base_url: str, *, timeout: float, tls_config: HttpxTLSConfig) -> bool: """Return True when ``/cluster-info`` confirms a hosted platform that omits ``/status``.""" try: - resp = httpx.get(f"{base_url.rstrip('/')}/cluster-info", timeout=timeout, verify=verify) + resp = httpx.get(f"{base_url.rstrip('/')}/cluster-info", timeout=timeout, **tls_config) except Exception: return False return resp.status_code == 200 @@ -485,13 +485,13 @@ def _check_controller_health(base_url: str, timeout: float = 5.0) -> tuple[bool, If ``controllers.status`` is empty on the first call (startup timing race), waits ``_CONTROLLER_HEALTH_RETRY_DELAY`` seconds and retries once. """ - verify = client_verify_from_env() + tls_config = httpx_tls_config_from_env() root = base_url.rstrip("/") for attempt in range(2): try: - resp = httpx.get(f"{root}/status", timeout=timeout, verify=verify) + resp = httpx.get(f"{root}/status", timeout=timeout, **tls_config) if resp.status_code == 404: - if _hosted_platform_without_status(root, timeout=timeout, verify=verify): + if _hosted_platform_without_status(root, timeout=timeout, tls_config=tls_config): return True, "Hosted deployment does not publish /status." return False, "Unexpected status 404 from /status endpoint." if resp.status_code != 200: diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/client/enhanced.py b/packages/nemo_platform_ext/src/nemo_platform_ext/client/enhanced.py index 3620b99dde..efc3697c7d 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/client/enhanced.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/client/enhanced.py @@ -21,8 +21,11 @@ not_given, ) from nemo_platform._base_client import AsyncAPIClient, SyncAPIClient +from nemo_platform_plugin.client.auth import AsyncTokenProvider, TokenProvider, TokenProviderAuth from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR -from nemo_platform_plugin.client.tls import client_verify_from_env +from nemo_platform_plugin.client.platform_options import AsyncPlatformClientOptions, SyncPlatformClientOptions +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env +from nemo_platform_plugin.client.types import RetryPolicy def _should_bootstrap_config( @@ -54,14 +57,42 @@ def _copy_requires_bootstrap( context_name: str | None, access_token: str | None, ) -> bool: - return ( - config_path is not None - or context_name is not None - or access_token is not None - or bool(os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)) + return config_path is not None or context_name is not None or access_token is not None + + +def _typed_client_default_headers( + custom_headers: Mapping[str, object], + transport_headers: Mapping[str, str], +) -> Mapping[str, str] | None: + headers = {str(key): value for key, value in custom_headers.items() if isinstance(value, str)} + if headers: + return headers + + default_transport_headers = {"accept", "accept-encoding", "connection", "user-agent", "host"} + headers = { + str(key): value + for key, value in transport_headers.items() + if str(key).lower() not in default_transport_headers and isinstance(value, str) + } + return headers or None + + +def _typed_client_retry(max_retries: int) -> RetryPolicy: + return RetryPolicy( + max_retries=max_retries, + retryable_status_codes=(408, 409, 429), + retry_all_server_errors=True, + respect_retry_decision_headers=True, + respect_retry_after_headers=True, ) +def _typed_client_timeout(timeout: float | Timeout | None) -> float | Timeout | None: + if timeout is None: + return httpx.Timeout(None) + return timeout + + class NeMoPlatform(SyncAPIClient): def __init__( self, @@ -76,6 +107,7 @@ def __init__( max_retries: int = DEFAULT_MAX_RETRIES, default_headers: Mapping[str, str] | None = None, default_query: Mapping[str, object] | None = None, + token_provider: TokenProvider | None = None, # Configure a custom httpx client. # We provide a `DefaultHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`. # See the [httpx documentation](https://www.python-httpx.org/api/#client) for more details. @@ -171,6 +203,7 @@ def __init__( if workspace is None: workspace = client_init_kwargs.workspace default_headers = client_init_kwargs.default_headers + token_provider = client_init_kwargs.token_provider if client_init_kwargs.http_client is not None and not isinstance( client_init_kwargs.http_client, httpx.Client ): @@ -185,11 +218,13 @@ def __init__( if base_url is None: raise RuntimeError("NeMoPlatform client initialization failed: base_url is required") - client_verify = client_verify_from_env() - if http_client is None and client_verify is not True: - http_client = DefaultHttpxClient(verify=client_verify) + tls_config = httpx_tls_config_from_env() + if http_client is None and tls_config: + http_client = DefaultHttpxClient(**tls_config) self.workspace = workspace + self._token_provider = token_provider + self._token_provider_auth = TokenProviderAuth(token_provider) if token_provider is not None else None super().__init__( version=__version__, @@ -205,6 +240,30 @@ def __init__( # TODO: needs to be removed self.inference_base_url = self._enforce_trailing_slash(httpx.URL(inference_base_url or base_url)) + @property + def custom_auth(self) -> TokenProviderAuth | None: + return self._token_provider_auth + + @property + def http_client(self) -> httpx.Client: + return self._client + + @property + def token_provider(self) -> TokenProvider | None: + return self._token_provider + + def typed_client_options(self) -> SyncPlatformClientOptions: + return SyncPlatformClientOptions( + base_url=str(self.base_url).rstrip("/"), + workspace=self.workspace, + default_headers=_typed_client_default_headers(self._custom_headers, self._client.headers), + timeout=_typed_client_timeout(self.timeout), + retry=_typed_client_retry(self.max_retries), + http_client=self._client, + url_resolver=self._prepare_url, + auth=self._token_provider, + ) + def __getattr__(self, name: str) -> Any: from nemo_platform_plugin.discovery import discover_sdk @@ -259,12 +318,16 @@ def copy( elif set_default_query is not None: params = set_default_query - if http_client is None and not _copy_requires_bootstrap( + requires_bootstrap = _copy_requires_bootstrap( config_path=config_path, context_name=context_name, access_token=access_token, - ): + ) + if http_client is None and not requires_bootstrap: http_client = self._client + extra_kwargs = dict(_extra_kwargs) + if not requires_bootstrap: + extra_kwargs.setdefault("token_provider", self._token_provider) return self.__class__( workspace=workspace or self.workspace, base_url=base_url or self.base_url, @@ -277,7 +340,7 @@ def copy( max_retries=self.max_retries if isinstance(max_retries, NotGiven) else max_retries, default_headers=headers, default_query=params, - **_extra_kwargs, + **extra_kwargs, ) @@ -298,6 +361,7 @@ def __init__( max_retries: int = DEFAULT_MAX_RETRIES, default_headers: Mapping[str, str] | None = None, default_query: Mapping[str, object] | None = None, + token_provider: TokenProvider | AsyncTokenProvider | None = None, # Configure a custom httpx client. # We provide a `DefaultAsyncHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`. # See the [httpx documentation](https://www.python-httpx.org/api/#asyncclient) for more details. @@ -412,6 +476,7 @@ async def main() -> None: if workspace is None: workspace = client_init_kwargs.workspace default_headers = client_init_kwargs.default_headers + token_provider = client_init_kwargs.token_provider if client_init_kwargs.http_client is not None and not isinstance( client_init_kwargs.http_client, httpx.AsyncClient ): @@ -426,11 +491,13 @@ async def main() -> None: if base_url is None: raise RuntimeError("NeMoPlatform client initialization failed: base_url is required") - client_verify = client_verify_from_env() - if http_client is None and client_verify is not True: - http_client = DefaultAsyncHttpxClient(verify=client_verify) + tls_config = httpx_tls_config_from_env() + if http_client is None and tls_config: + http_client = DefaultAsyncHttpxClient(**tls_config) self.workspace = workspace + self._token_provider = token_provider + self._token_provider_auth = TokenProviderAuth(token_provider) if token_provider is not None else None super().__init__( version=__version__, @@ -449,6 +516,30 @@ async def main() -> None: # TODO: needs to be removed self.inference_base_url = self._enforce_trailing_slash(httpx.URL(inference_base_url or base_url)) + @property + def custom_auth(self) -> TokenProviderAuth | None: + return self._token_provider_auth + + @property + def http_client(self) -> httpx.AsyncClient: + return self._client + + @property + def token_provider(self) -> TokenProvider | AsyncTokenProvider | None: + return self._token_provider + + def typed_client_options(self) -> AsyncPlatformClientOptions: + return AsyncPlatformClientOptions( + base_url=str(self.base_url).rstrip("/"), + workspace=self.workspace, + default_headers=_typed_client_default_headers(self._custom_headers, self._client.headers), + timeout=_typed_client_timeout(self.timeout), + retry=_typed_client_retry(self.max_retries), + http_client=self._client, + url_resolver=self._prepare_url, + auth=self._token_provider, + ) + def __getattr__(self, name: str) -> Any: from nemo_platform_plugin.discovery import discover_sdk @@ -503,12 +594,16 @@ def copy( elif set_default_query is not None: params = set_default_query - if http_client is None and not _copy_requires_bootstrap( + requires_bootstrap = _copy_requires_bootstrap( config_path=config_path, context_name=context_name, access_token=access_token, - ): + ) + if http_client is None and not requires_bootstrap: http_client = self._client + extra_kwargs = dict(_extra_kwargs) + if not requires_bootstrap: + extra_kwargs.setdefault("token_provider", self._token_provider) return self.__class__( workspace=workspace or self.workspace, base_url=base_url or self.base_url, @@ -521,5 +616,5 @@ def copy( max_retries=self.max_retries if isinstance(max_retries, NotGiven) else max_retries, default_headers=headers, default_query=params, - **_extra_kwargs, + **extra_kwargs, ) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py b/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py index 64c001c9ac..9e881c2bc1 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/client/factory.py @@ -67,6 +67,7 @@ not_given, ) from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from nemo_platform_ext.auth.helpers import NMPOIDCConfig, build_effective_scope, discover_nmp_config from nemo_platform_ext.auth.token_provider import ( @@ -74,7 +75,6 @@ TokenSet, ) from nemo_platform_ext.auth.workload_exchange import WorkloadTokenExchangeProvider -from nemo_platform_ext.client.tls import client_verify_from_env logger = logging.getLogger(__name__) @@ -103,13 +103,15 @@ class ClientInitConfig: For non-OAuth users this just carries base_url/workspace/headers. For OAuth users it also includes a custom httpx client with an event - hook that injects/refreshes the Bearer token on every request. + hook that injects/refreshes the Bearer token on every request, plus the + token provider itself for typed client adapters. """ base_url: str workspace: str | None default_headers: Mapping[str, str] | None = None http_client: httpx.Client | httpx.AsyncClient | None = None + token_provider: _AccessTokenProvider | None = None @dataclass(frozen=True) @@ -634,13 +636,14 @@ def build_client_init_kwargs( http_client = DefaultHttpxClient( event_hooks={"request": [hook], "response": []}, follow_redirects=True, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, default_headers=headers or None, http_client=http_client, + token_provider=bootstrap.token_provider, ) @@ -677,13 +680,14 @@ def build_async_client_init_kwargs( http_client = DefaultAsyncHttpxClient( event_hooks={"request": [hook], "response": []}, follow_redirects=True, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) return ClientInitConfig( base_url=bootstrap.base_url, workspace=bootstrap.workspace, default_headers=headers or None, http_client=http_client, + token_provider=bootstrap.token_provider, ) @@ -721,6 +725,7 @@ def create_client( workspace=client_init_kwargs.workspace, default_headers=client_init_kwargs.default_headers, http_client=http_client, + token_provider=client_init_kwargs.token_provider, max_retries=max_retries, timeout=timeout, ) diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/client/tls.py b/packages/nemo_platform_ext/src/nemo_platform_ext/client/tls.py deleted file mode 100644 index b2bc998d2e..0000000000 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/client/tls.py +++ /dev/null @@ -1,16 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""TLS configuration shared by NeMo Platform SDK and CLI clients.""" - -from __future__ import annotations - -import os - -NMP_CLIENT_SSL_CERT_FILE_ENVVAR = "NMP_CLIENT_SSL_CERT_FILE" - - -def client_verify_from_env() -> str | bool: - """Return the httpx verify setting for NeMo Platform client requests.""" - cert_file = os.environ.get(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "").strip() - return cert_file or True diff --git a/packages/nemo_platform_ext/tests/auth/test_device_flow.py b/packages/nemo_platform_ext/tests/auth/test_device_flow.py index 4672af4d7d..04f1003362 100644 --- a/packages/nemo_platform_ext/tests/auth/test_device_flow.py +++ b/packages/nemo_platform_ext/tests/auth/test_device_flow.py @@ -26,7 +26,7 @@ authenticate_with_password_grant, refresh_access_token, ) -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR class TestDeviceCodeResponse: diff --git a/packages/nemo_platform_ext/tests/auth/test_token_provider.py b/packages/nemo_platform_ext/tests/auth/test_token_provider.py index 785d14cdf3..555a435a55 100644 --- a/packages/nemo_platform_ext/tests/auth/test_token_provider.py +++ b/packages/nemo_platform_ext/tests/auth/test_token_provider.py @@ -15,7 +15,7 @@ TokenSet, refresh_token_grant, ) -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR def _make_jwt(claims: dict, header: dict | None = None) -> str: diff --git a/packages/nemo_platform_ext/tests/auth/test_utils.py b/packages/nemo_platform_ext/tests/auth/test_utils.py index 6c66577652..180c91bbd3 100644 --- a/packages/nemo_platform_ext/tests/auth/test_utils.py +++ b/packages/nemo_platform_ext/tests/auth/test_utils.py @@ -19,7 +19,7 @@ normalize_scope_prefix, validate_requested_scopes_granted, ) -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR from pytest_httpserver import HTTPServer diff --git a/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py b/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py index 475ee61c64..23d8e8c2de 100644 --- a/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py +++ b/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py @@ -23,8 +23,8 @@ read_subject_token_file, token_exchange_grant, ) -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR def test_workload_exchange_module_has_no_nmp_common_dependency(): diff --git a/packages/nemo_platform_ext/tests/cli/commands/test_setup.py b/packages/nemo_platform_ext/tests/cli/commands/test_setup.py index 9fdfa34886..18ed4963d9 100644 --- a/packages/nemo_platform_ext/tests/cli/commands/test_setup.py +++ b/packages/nemo_platform_ext/tests/cli/commands/test_setup.py @@ -198,7 +198,7 @@ def test_reachable_via_status(self): mock_resp = MagicMock() mock_resp.status_code = 200 with ( - patch(f"{SETUP_MOD}.client_verify_from_env", return_value="/tmp/custom-ca.pem"), + patch(f"{SETUP_MOD}.httpx_tls_config_from_env", return_value={"verify": "/tmp/custom-ca.pem"}), patch(f"{SETUP_MOD}.httpx.get", return_value=mock_resp) as mock_get, ): assert _check_platform_reachable("http://localhost:8080") is True @@ -224,7 +224,7 @@ def _get(url, **kwargs): raise AssertionError(f"unexpected url: {url}") with ( - patch(f"{SETUP_MOD}.client_verify_from_env", return_value="/tmp/custom-ca.pem"), + patch(f"{SETUP_MOD}.httpx_tls_config_from_env", return_value={"verify": "/tmp/custom-ca.pem"}), patch(f"{SETUP_MOD}.httpx.get", side_effect=_get), ): assert _check_platform_reachable("https://nemo-platform-freeplay.dev.aire.nvidia.com") is True diff --git a/packages/nemo_platform_ext/tests/client/test_client.py b/packages/nemo_platform_ext/tests/client/test_client.py index e4b79332f9..cac1ba47f1 100644 --- a/packages/nemo_platform_ext/tests/client/test_client.py +++ b/packages/nemo_platform_ext/tests/client/test_client.py @@ -17,8 +17,8 @@ from nemo_platform import AsyncNeMoPlatform, DefaultHttpxClient, NeMoPlatform, not_given from nemo_platform_ext.auth.helpers import NMPOIDCConfig, decode_jwt_claims from nemo_platform_ext.client.factory import create_client -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR def _make_jwt(claims: dict) -> str: @@ -810,7 +810,7 @@ async def test_async_constructor_passes_context_name_to_bootstrap(self, mock_bui class TestAsyncNeMoPlatformInit: @pytest.mark.asyncio - @patch("nemo_platform.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) + @patch("nemo_platform_ext.client.factory.discover_nmp_config", return_value=_MOCK_NMP_CONFIG) async def test_async_client_uses_config_for_api_key(self, _mock_discover, tmp_path): config_path = _write_config(tmp_path, user_type="api-key", api_key="nvapi-test-key-123") diff --git a/packages/nemo_platform_ext/tests/local/test_health_child.py b/packages/nemo_platform_ext/tests/local/test_health_child.py index 04a975d2f0..0830896592 100644 --- a/packages/nemo_platform_ext/tests/local/test_health_child.py +++ b/packages/nemo_platform_ext/tests/local/test_health_child.py @@ -53,7 +53,11 @@ def controller_run(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -87,7 +91,11 @@ def stubborn_controller(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -107,25 +115,21 @@ def stubborn_controller(stop_signal: threading.Event) -> None: @pytest.mark.integration -def test_lifespan_cleanup_runs_on_app_shutdown() -> None: - """``close_shared_http_clients`` should be called during lifespan teardown.""" - cleanup_called = threading.Event() - +def test_lifespan_shutdown_marks_controller_stop_signal() -> None: + """Lifespan teardown should mark the controller stop signal.""" with ( patch("nmp.platform_runner.server.get_platform_config") as mock_pc, patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, - patch("nmp.platform_runner.server.close_shared_http_clients") as mock_close, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") - async def fake_close(): - cleanup_called.set() - - mock_close.side_effect = fake_close - from nmp.platform_runner.server import create_app app = create_app(services=[]) @@ -135,7 +139,7 @@ async def fake_close(): with TestClient(app): pass - assert cleanup_called.is_set(), "close_shared_http_clients was not called during shutdown" + assert app.state.controller_stop_signal.is_set() # --------------------------------------------------------------------------- diff --git a/packages/nemo_platform_ext/tests/local/test_services_contract.py b/packages/nemo_platform_ext/tests/local/test_services_contract.py index acfdc227ca..28d8c34b07 100644 --- a/packages/nemo_platform_ext/tests/local/test_services_contract.py +++ b/packages/nemo_platform_ext/tests/local/test_services_contract.py @@ -274,7 +274,7 @@ def get_routers(self): return [] dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": sidecar_run_func} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": sidecar_run_func} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -298,6 +298,7 @@ def get_routers(self): monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) diff --git a/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py b/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py index 33cadabaa7..85d4d7a642 100644 --- a/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py +++ b/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py @@ -42,7 +42,7 @@ def get_routers(self): def _sidecar_with_events(started: threading.Event, stopped: threading.Event) -> Callable[[threading.Event], None]: - """Return a sidecar ``run(stop_signal)`` that signals start/stop via events.""" + """Return a sidecar run function that signals start/stop via events.""" def run(stop_signal: threading.Event) -> None: started.set() @@ -71,7 +71,7 @@ def patched_registry( test sidecar, plus minimal auth/platform config stubs.""" started, stopped = sidecar_events dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": _sidecar_with_events(started, stopped)} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": _sidecar_with_events(started, stopped)} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -95,6 +95,7 @@ def patched_registry( monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) 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 410a1ae6fc..fe17d0eeed 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 @@ -17,24 +17,40 @@ def make_sync_resource(platform: NeMoPlatform) -> NemoClient: from __future__ import annotations -from collections.abc import Callable from typing import TypeVar, overload -import httpx from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient -from nemo_platform_plugin.client.types import RetryPolicy +from nemo_platform_plugin.client.platform_options import AsyncPlatformClientOptions, SyncPlatformClientOptions SyncT = TypeVar("SyncT", bound=NemoClient) AsyncT = TypeVar("AsyncT", bound=AsyncNemoClient) -def _url_resolver_from_platform(platform: NeMoPlatform | AsyncNeMoPlatform) -> Callable[[str], str | httpx.URL]: - router = getattr(platform, "_nmp_request_router", None) - resolver = getattr(router, "resolve", None) - if resolver is not None: - return resolver - return platform._prepare_url +def _sync_client_from_options(client_cls: type[SyncT], options: SyncPlatformClientOptions) -> SyncT: + return client_cls( + base_url=options.base_url, + workspace=options.workspace, + default_headers=options.default_headers, + timeout=options.timeout, + retry=options.retry, + http_client=options.http_client, + url_resolver=options.url_resolver, + auth=options.auth, + ) + + +def _async_client_from_options(client_cls: type[AsyncT], options: AsyncPlatformClientOptions) -> AsyncT: + return client_cls( + base_url=options.base_url, + workspace=options.workspace, + default_headers=options.default_headers, + timeout=options.timeout, + retry=options.retry, + http_client=options.http_client, + url_resolver=options.url_resolver, + auth=options.auth, + ) @overload @@ -51,58 +67,11 @@ def client_from_platform( The overloads ensure callers get the correct concrete return type. """ - # Prefer _custom_headers (set via with_options/set_default_headers), - # fall back to the httpx client's actual headers (set at construction, - # e.g. TestClient(headers={...})), filtering out httpx defaults. - # _custom_headers and _client are private Stainless SDK attrs present on both - # NeMoPlatform and AsyncNeMoPlatform but not visible to the type checker. - headers = platform._custom_headers # type: ignore[union-attr] - if not headers: - _skip = {"accept", "accept-encoding", "connection", "user-agent", "host"} - headers = {k: v for k, v in platform._client.headers.items() if k.lower() not in _skip} # type: ignore[union-attr] - - 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) - - # Carry the platform's timeout across as a per-request override. The shared - # httpx client keeps whatever timeout it was built with, so a caller's - # ``platform.with_options(timeout=...)`` would otherwise be silently dropped - # on the way to the typed client — the httpx client it hands over is the - # *same* object, with the *original* timeout still on it. - timeout = platform.timeout - if timeout is None: - # ``None`` on the platform means "no timeout at all", but the typed - # client reads None as "defer to the transport". Say the same thing in - # the form httpx itself uses, so the override survives. - timeout = httpx.Timeout(None) - if isinstance(platform, AsyncNeMoPlatform): if not issubclass(client_cls, AsyncNemoClient): raise TypeError("AsyncNeMoPlatform requires an AsyncNemoClient class") - return client_cls( - base_url=str(platform.base_url).rstrip("/"), - workspace=platform.workspace, - default_headers=headers or None, - timeout=timeout, - retry=retry, - http_client=platform._client, - url_resolver=url_resolver, - ) + return _async_client_from_options(client_cls, platform.typed_client_options()) if not issubclass(client_cls, NemoClient): raise TypeError("NeMoPlatform requires a NemoClient class") - return client_cls( - base_url=str(platform.base_url).rstrip("/"), - workspace=platform.workspace, - default_headers=headers or None, - timeout=timeout, - retry=retry, - http_client=platform._client, - url_resolver=url_resolver, - ) + return _sync_client_from_options(client_cls, platform.typed_client_options()) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/auth.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/auth.py index f0640ef4ec..e3a4928793 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/auth.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/auth.py @@ -13,9 +13,8 @@ from __future__ import annotations import asyncio -import inspect from collections.abc import AsyncGenerator, Generator -from typing import Protocol, cast, runtime_checkable +from typing import Protocol, runtime_checkable import httpx @@ -32,10 +31,10 @@ def get_access_token(self) -> str: ... @runtime_checkable -class AsyncTokenProvider(Protocol): +class AsyncTokenProvider(TokenProvider, Protocol): """Async protocol for objects that can supply an access token.""" - async def get_access_token(self) -> str: ... + async def get_access_token_async(self) -> str: ... # --------------------------------------------------------------------------- @@ -67,16 +66,12 @@ async def resolve_token_async(provider: TokenProvider | AsyncTokenProvider) -> s Three cases, in priority order: 1. Provider has ``get_access_token_async()`` (e.g. OIDCTokenProvider) — use it. - 2. ``get_access_token()`` is a coroutine function — await it. - 3. ``get_access_token()`` is sync — run in a thread, since it may perform IO + 2. ``get_access_token()`` is sync — run in a thread, since it may perform IO such as a token refresh. """ - get_async = getattr(provider, "get_access_token_async", None) - if get_async is not None and callable(get_async): - return await get_async() - if inspect.iscoroutinefunction(provider.get_access_token): - return await provider.get_access_token() - return await asyncio.to_thread(cast(TokenProvider, provider).get_access_token) + if isinstance(provider, AsyncTokenProvider): + return await provider.get_access_token_async() + return await asyncio.to_thread(provider.get_access_token) # --------------------------------------------------------------------------- @@ -99,14 +94,13 @@ class TokenProviderAuth(httpx.Auth): def __init__(self, provider: TokenProvider | AsyncTokenProvider) -> None: self._provider = provider + @property + def provider(self) -> TokenProvider | AsyncTokenProvider: + return self._provider + def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]: if "Authorization" not in request.headers: token = self._provider.get_access_token() - if inspect.isawaitable(token): - raise TypeError( - "Async token provider used on a synchronous transport; " - "use AsyncNemoClient with an AsyncTokenProvider." - ) request.headers["Authorization"] = f"Bearer {token}" yield request diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py index d39f09b665..d82f13c678 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py @@ -4,11 +4,13 @@ """Shared client constants and env checks.""" import os +from pathlib import Path WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE" JWT_WORKLOAD_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = "urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof" OPAQUE_DOCKER_PROOF_PREFIX = "nmp_obo_v1" +NMP_PRINCIPAL_ENVVAR = "NMP_PRINCIPAL" def is_workload_identity_token_file_set() -> bool: @@ -20,3 +22,13 @@ def subject_token_type_for_exchange(subject_token: str) -> str: if subject_token.startswith(f"{OPAQUE_DOCKER_PROOF_PREFIX}."): return DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE return JWT_WORKLOAD_SUBJECT_TOKEN_TYPE + + +def workload_identity_token_file_from_env() -> Path | None: + token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR) + return Path(token_file) if token_file else None + + +def require_workload_identity_without_principal_env() -> None: + if is_workload_identity_token_file_set() and os.environ.get(NMP_PRINCIPAL_ENVVAR): + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} and {NMP_PRINCIPAL_ENVVAR} are mutually exclusive") diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py index 1319655690..4ad6e2ef44 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py @@ -38,7 +38,7 @@ WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, subject_token_type_for_exchange, ) -from nemo_platform_plugin.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env logger = logging.getLogger(__name__) @@ -159,7 +159,7 @@ def discover_nmp_config(base_url: str, timeout: float = 10.0) -> NMPOIDCConfig: response = httpx.get( f"{base_url.rstrip('/')}/apis/auth/discovery", timeout=timeout, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) response.raise_for_status() data = response.json() @@ -266,7 +266,7 @@ def refresh_token_grant( if scope: data["scope"] = scope - response = httpx.post(token_endpoint, data=data, timeout=timeout, verify=client_verify_from_env()) + response = httpx.post(token_endpoint, data=data, timeout=timeout, **httpx_tls_config_from_env()) if response.status_code != 200: error_data: dict[str, str] = {} @@ -325,7 +325,7 @@ def token_exchange_grant( if scope: data["scope"] = scope - response = httpx.post(token_endpoint, data=data, timeout=timeout, verify=client_verify_from_env()) + response = httpx.post(token_endpoint, data=data, timeout=timeout, **httpx_tls_config_from_env()) if response.status_code != 200: error_data: dict[str, object] = {} if response.headers.get("content-type", "").startswith("application/json"): diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/platform_options.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/platform_options.py new file mode 100644 index 0000000000..bee110d7f7 --- /dev/null +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/platform_options.py @@ -0,0 +1,39 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Typed options exported by the platform SDK for typed client construction.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from dataclasses import dataclass + +import httpx +from nemo_platform_plugin.client.auth import AsyncTokenProvider, TokenProvider +from nemo_platform_plugin.client.types import RetryPolicy + +URLResolver = Callable[[str], str | httpx.URL] + + +@dataclass(frozen=True) +class SyncPlatformClientOptions: + base_url: str + workspace: str | None + default_headers: Mapping[str, str] | None + timeout: float | httpx.Timeout | None + retry: RetryPolicy + http_client: httpx.Client + url_resolver: URLResolver + auth: TokenProvider | None = None + + +@dataclass(frozen=True) +class AsyncPlatformClientOptions: + base_url: str + workspace: str | None + default_headers: Mapping[str, str] | None + timeout: float | httpx.Timeout | None + retry: RetryPolicy + http_client: httpx.AsyncClient + url_resolver: URLResolver + auth: TokenProvider | AsyncTokenProvider | None = None diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/tls.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/tls.py index 86f8e5871f..b6f8ef4191 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/tls.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/tls.py @@ -6,11 +6,26 @@ from __future__ import annotations import os +from collections.abc import Mapping, Sequence +from typing import TypedDict NMP_CLIENT_SSL_CERT_FILE_ENVVAR = "NMP_CLIENT_SSL_CERT_FILE" -def client_verify_from_env() -> str | bool: - """Return the httpx verify setting for NeMo Platform client requests.""" - cert_file = os.environ.get(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "").strip() - return cert_file or True +class HttpxTLSConfig(TypedDict, total=False): + """TLS kwargs passed to HTTPX client and request calls.""" + + verify: str + + +def httpx_tls_config_from_env( + env: Mapping[str, str] | None = None, + *, + cert_file_envvars: Sequence[str] = (NMP_CLIENT_SSL_CERT_FILE_ENVVAR,), +) -> HttpxTLSConfig: + """Return HTTPX TLS kwargs for NeMo Platform client requests.""" + for envvar in cert_file_envvars: + cert_file = (env[envvar] if env is not None and envvar in env else os.environ.get(envvar, "")).strip() + if cert_file: + return {"verify": cert_file} + return {} diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py index e59ba4dc46..cdd84534e7 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py @@ -1,33 +1,27 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""NemoClient factory for task containers and services — the plugin-side -interface for building authenticated -:class:`~nemo_platform_plugin.client.client.NemoClient` / -:class:`~nemo_platform_plugin.client.client.AsyncNemoClient` handles. +"""Standalone NemoClient factory for plugin-only task containers. -This is the :class:`~nemo_platform_plugin.client.client.NemoClient` sibling of -:mod:`nemo_platform_plugin.sdk_provider`. Plugin authors call -:func:`get_nemo_client` / :func:`get_async_nemo_client` here instead of -importing from ``nmp.common``. This keeps ``nemo-platform-plugin`` free of any -``nmp-common`` dependency while still allowing the platform to register a richer -provider (URL routing, shared HTTP clients, OTEL headers, workload identity, -...) when ``nmp-common`` is installed. +Most platform/service code should acquire a ``NeMoPlatform`` SDK through +``sdk_provider`` and adapt it to a typed service client with +``client_from_platform(sdk, ServiceClient)``. This module exists for +plugin-only environments that cannot import ``nmp.common`` and need direct +:class:`~nemo_platform_plugin.client.client.NemoClient` or +:class:`~nemo_platform_plugin.client.client.AsyncNemoClient` handles. Lookup order for the provider ----------------------------- 1. **Explicit override** — set via :func:`set_nemo_client_provider` (for tests). -2. **Entry-point discovery** — scans the ``nemo.client_provider`` group. - When ``nmp-common`` is installed in the image (platform deployment), its - provider is picked up automatically. +2. **Entry-point discovery** — scans the ``nemo.client_provider`` group for + optional third-party overrides. 3. **Built-in default** — :class:`DefaultNemoClientProvider`, an env-var-based implementation that reads ``NMP_BASE_URL`` and ``NMP_PRINCIPAL``. Works for local development and gateway-routed task containers. -For user-facing / CLI usage, prefer ``NemoClient.from_config()`` which reads -``~/.config/nmp/config.yaml`` and wires up OIDC token refresh / workload -identity token exchange. +For user-facing / CLI usage, prefer ``nemo_platform.NeMoPlatform`` or +``nemo_platform.AsyncNeMoPlatform``. """ from __future__ import annotations @@ -42,14 +36,17 @@ from nemo_platform_plugin.client.auth import TokenProvider from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_plugin.client.constants import ( + NMP_PRINCIPAL_ENVVAR, WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, is_workload_identity_token_file_set, + require_workload_identity_without_principal_env, + workload_identity_token_file_from_env, ) logger = logging.getLogger(__name__) _INTERNAL_REQUEST_HEADER = "X-NMP-Internal" -_NMP_PRINCIPAL_ENVVAR = "NMP_PRINCIPAL" +_NMP_PRINCIPAL_ENVVAR = NMP_PRINCIPAL_ENVVAR # --------------------------------------------------------------------------- @@ -223,10 +220,25 @@ def _workload_identity_auth(base_url: str) -> TokenProvider: """ from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider - token_file = os.environ[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] + token_file = workload_identity_token_file_from_env() + if token_file is None: + raise RuntimeError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} is not set") return resolve_workload_exchange_provider(base_url=base_url, subject_token_file=Path(token_file)) +def _workload_identity_headers(*, internal: bool) -> dict[str, str]: + return {_INTERNAL_REQUEST_HEADER: "true"} if internal else {} + + +def _ensure_no_trusted_headers_for_workload_identity( + *, + as_service: str | None, + on_behalf_of: str | None, +) -> None: + if as_service is not None or on_behalf_of is not None: + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} cannot be combined with trusted principal headers") + + def _base_url() -> str: return os.environ.get("NMP_BASE_URL", "http://localhost:8080") @@ -247,8 +259,19 @@ def get_nemo_client( on_behalf_of: str | None = None, workspace: str | None = None, ) -> NemoClient: + base_url = _base_url() + if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() + _ensure_no_trusted_headers_for_workload_identity(as_service=as_service, on_behalf_of=on_behalf_of) + return NemoClient( + base_url=base_url, + workspace=workspace, + auth=_workload_identity_auth(base_url), + default_headers=_workload_identity_headers(internal=internal) or None, + ) + headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return NemoClient(base_url=_base_url(), workspace=workspace, default_headers=headers or None) + return NemoClient(base_url=base_url, workspace=workspace, default_headers=headers or None) def get_async_nemo_client( self, @@ -258,8 +281,19 @@ def get_async_nemo_client( on_behalf_of: str | None = None, workspace: str | None = None, ) -> AsyncNemoClient: + base_url = _base_url() + if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() + _ensure_no_trusted_headers_for_workload_identity(as_service=as_service, on_behalf_of=on_behalf_of) + return AsyncNemoClient( + base_url=base_url, + workspace=workspace, + auth=_workload_identity_auth(base_url), + default_headers=_workload_identity_headers(internal=internal) or None, + ) + headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return AsyncNemoClient(base_url=_base_url(), workspace=workspace, default_headers=headers or None) + return AsyncNemoClient(base_url=base_url, workspace=workspace, default_headers=headers or None) def get_task_nemo_client( self, @@ -269,11 +303,12 @@ def get_task_nemo_client( ) -> NemoClient: base_url = _base_url() if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() return NemoClient( base_url=base_url, workspace=workspace, auth=_workload_identity_auth(base_url), - default_headers={_INTERNAL_REQUEST_HEADER: "true"}, + default_headers=_workload_identity_headers(internal=True), ) return NemoClient( base_url=base_url, @@ -289,11 +324,12 @@ def get_async_task_nemo_client( ) -> AsyncNemoClient: base_url = _base_url() if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() return AsyncNemoClient( base_url=base_url, workspace=workspace, auth=_workload_identity_auth(base_url), - default_headers={_INTERNAL_REQUEST_HEADER: "true"}, + default_headers=_workload_identity_headers(internal=True), ) return AsyncNemoClient( base_url=base_url, @@ -391,7 +427,7 @@ def get_nemo_client( Delegates to the resolved :class:`NemoClientProvider`. Under the built-in default this reads ``NMP_BASE_URL`` (default ``http://localhost:8080``) and ``NMP_PRINCIPAL`` from the environment; under the platform provider it - additionally routes service URLs, reuses the shared HTTP client, and injects + additionally routes service URLs, uses endpoint-aware HTTP clients, and injects OTEL headers. Args: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/config.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/config.py index 7a0c994084..f538dcdab9 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/config.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/config.py @@ -340,9 +340,6 @@ def get_platform_config_class() -> Type[NemoPlatformConfig]: # Platform configuration types # --------------------------------------------------------------------------- -# Regex for env vars that set per-service URLs (e.g. NMP_FILES_URL). -_PLATFORM_SERVICE_URL_ENV_PATTERN = re.compile(r"^NMP_([A-Z0-9]+)_URL$") - # Addresses that refer to localhost/loopback interfaces LOOPBACK_ADDRESSES = ("localhost", "0.0.0.0", "::1", "127.0.0.1") @@ -477,9 +474,8 @@ class NemoPlatformConfig(ServiceConfig): """Platform-wide configuration settings. It inherits from ServiceConfig and provides Platform-centric settings, which may be used by other microservices to interact with other Platform services. - Environment variables NMP__URL (e.g. NMP_FILES_URL) are read and merged into - service_discovery with the service name lowercased; NMP_BASE_URL sets base_url and is not added to - service_discovery. + ``base_url`` is the default platform URL; ``service_discovery`` can pin + specific services to different URLs. """ model_config = SettingsConfigDict( @@ -532,8 +528,7 @@ def get(cls) -> NemoPlatformConfig: default_factory=dict, description=( "Map of service names to their URLs. Used to discover services by name (e.g. 'files': 'http://files-service:8080'). " - "Environment variables NMP__URL (e.g. NMP_FILES_URL) are read and merged " - "into this map with the service name lowercased; NMP_BASE_URL is not added here (it sets base_url)." + "Per-service environment URL overrides are resolved by endpoint/client factories, not merged into this config." ), ) @@ -641,25 +636,6 @@ def get_service_url(self, api_name: str) -> str: def create_service_pattern(self) -> re.Pattern[str] | None: return re.compile(r"/apis/([a-z]+(?:-[a-z]+)*)/") - @model_validator(mode="before") - @classmethod - def merge_service_url_env_vars(cls, values: Any) -> Any: - if not isinstance(values, dict): - return values - sd = dict(values.get("service_discovery") or {}) - for key, value in environ.items(): - if not value: - continue - match = _PLATFORM_SERVICE_URL_ENV_PATTERN.match(key) - if not match: - continue - service_name = match.group(1) - if service_name == "BASE": - continue - sd[service_name.lower()] = value - values["service_discovery"] = sd - return values - @model_validator(mode="before") @classmethod def handle_docker_env_vars(cls, values: Any) -> Any: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py index 156143b79d..e0753c789e 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py @@ -54,14 +54,14 @@ def get_sdk_client() -> "AsyncNeMoPlatform": def get_nemo_client() -> AsyncNemoClient: - """FastAPI dependency for getting the async NemoClient. + """Legacy FastAPI dependency placeholder for getting an async NemoClient. - This is a placeholder. The actual client is injected via - app.dependency_overrides in Service.create_app(). + Platform services should depend on ``get_sdk_client`` and adapt to typed + service clients with ``client_from_platform(sdk, ServiceClient)`` instead. """ raise RuntimeError( - "get_nemo_client() was called without being overridden. " - "Ensure your Service subclass calls super().create_app()." + "get_nemo_client() is not wired by platform services. " + "Depend on get_sdk_client() and adapt with client_from_platform(sdk, ServiceClient)." ) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/jobs/result_manager.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/jobs/result_manager.py index cac46051d9..75120bd8c5 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/jobs/result_manager.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/jobs/result_manager.py @@ -7,7 +7,7 @@ from abc import ABC from dataclasses import dataclass, field from pathlib import Path -from typing import Generic, Literal, Type, TypeVar, overload +from typing import Generic, Literal, Type, TypeVar, cast, overload from filesets import parse_fileset_ref from nemo_platform import AsyncNeMoPlatform, NeMoPlatform @@ -37,6 +37,16 @@ class FileDoesNotExist(Exception): ... PlatformSDKT = TypeVar("PlatformSDKT", "NeMoPlatform", "AsyncNeMoPlatform") +def _owned_sdks_once(*sdk_ownerships: tuple[PlatformSDKT, bool]) -> list[PlatformSDKT]: + owned_sdks: list[PlatformSDKT] = [] + seen_sdk_ids: set[int] = set() + for sdk, should_close in sdk_ownerships: + if should_close and id(sdk) not in seen_sdk_ids: + owned_sdks.append(sdk) + seen_sdk_ids.add(id(sdk)) + return owned_sdks + + @dataclass class BaseResultManager(Generic[FileManagerClsT, PlatformSDKT], ABC): """ @@ -50,6 +60,8 @@ class BaseResultManager(Generic[FileManagerClsT, PlatformSDKT], ABC): files_sdk: PlatformSDKT jobs_sdk: PlatformSDKT attempt_id: str | None = field(default=None) + owns_files_sdk: bool = field(default=False) + owns_jobs_sdk: bool = field(default=False) def _validate_local_path(self, artifact_local_path: str | Path) -> Path: if isinstance(artifact_local_path, str): @@ -68,6 +80,13 @@ def _result_remote_path(self, attempt_id: str, result_name: str, base: str | Non @dataclass class ResultManager(BaseResultManager[Type[FilesetFileManager], NeMoPlatform]): + def close(self) -> None: + for sdk in _owned_sdks_once( + (self.files_sdk, self.owns_files_sdk), + (self.jobs_sdk, self.owns_jobs_sdk), + ): + sdk.close() + def _fetch_job_metadata(self) -> tuple[str, str, str | None]: """Fetch job and return (attempt_id, fileset_name, output_location).""" jobs = client_from_platform(self.jobs_sdk, JobsClient) @@ -135,6 +154,13 @@ def download_artifact(self, artifact_url: str, local_dir: str | Path | None = No @dataclass class AsyncResultManager(BaseResultManager[Type[AsyncFilesetFileManager], AsyncNeMoPlatform]): + async def aclose(self) -> None: + for sdk in _owned_sdks_once( + (self.files_sdk, self.owns_files_sdk), + (self.jobs_sdk, self.owns_jobs_sdk), + ): + await sdk.close() + async def _fetch_job_metadata(self) -> tuple[str, str, str | None]: """Fetch job and return (attempt_id, fileset_name, output_location).""" jobs = client_from_platform(self.jobs_sdk, AsyncJobsClient) @@ -254,18 +280,31 @@ def result_manager_factory( if workspace is None: workspace = _get_job_workspace() + if is_async: + if jobs_sdk is None: + jobs_sdk = files_sdk + async_files_sdk = cast(AsyncNeMoPlatform, files_sdk) + async_jobs_sdk = cast(AsyncNeMoPlatform, jobs_sdk) + return AsyncResultManager( + job_name=job_name, + workspace=workspace, + attempt_id=attempt_id, + file_manager_cls=AsyncFilesetFileManager, + files_sdk=async_files_sdk, + jobs_sdk=async_jobs_sdk, + ) + if jobs_sdk is None: jobs_sdk = files_sdk - - file_manager_cls = AsyncFilesetFileManager if is_async else FilesetFileManager - result_manager_cls = AsyncResultManager if is_async else ResultManager - return result_manager_cls( + sync_files_sdk = cast(NeMoPlatform, files_sdk) + sync_jobs_sdk = cast(NeMoPlatform, jobs_sdk) + return ResultManager( job_name=job_name, workspace=workspace, attempt_id=attempt_id, - file_manager_cls=file_manager_cls, - files_sdk=files_sdk, - jobs_sdk=jobs_sdk, + file_manager_cls=FilesetFileManager, + files_sdk=sync_files_sdk, + jobs_sdk=sync_jobs_sdk, ) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py index 49379fa2a2..fa86ba3b4f 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py @@ -7,8 +7,8 @@ Plugin authors use :func:`get_task_sdk` in their ``__main__.py`` entrypoints instead of importing from ``nmp.common.sdk_factory``. This keeps the ``nemo-platform-plugin`` package free of ``nmp-common`` dependencies while -still allowing the platform to register a richer provider (with URL routing, -shared HTTP clients, OTEL headers, etc.) when ``nmp-common`` is installed. +still allowing the platform to register a richer provider (with platform auth +context, OTEL headers, etc.) when ``nmp-common`` is installed. Lookup order for the provider ----------------------------- @@ -36,21 +36,25 @@ import logging import os from importlib.metadata import entry_points -from typing import Any, Protocol, TypeVar, runtime_checkable - -from nemo_platform import AsyncNeMoPlatform, NeMoPlatform -from nemo_platform_plugin.client.constants import is_workload_identity_token_file_set - -_SDKT = TypeVar("_SDKT", NeMoPlatform, AsyncNeMoPlatform) +from typing import Any, Protocol, overload, runtime_checkable + +from nemo_platform import AsyncNeMoPlatform, DefaultAsyncHttpxClient, DefaultHttpxClient, NeMoPlatform +from nemo_platform_plugin.client.auth import TokenProvider, TokenProviderAuth +from nemo_platform_plugin.client.constants import ( + NMP_PRINCIPAL_ENVVAR, + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + is_workload_identity_token_file_set, + require_workload_identity_without_principal_env, + workload_identity_token_file_from_env, +) logger = logging.getLogger(__name__) # Header name the platform uses to mark internal (service-to-service) requests. _INTERNAL_REQUEST_HEADER = "X-NMP-Internal" -# Environment variable the jobs backend writes with the job creator's -# principal (JSON-serialised). -_NMP_PRINCIPAL_ENVVAR = "NMP_PRINCIPAL" +# Environment variable the jobs backend writes with the job creator's principal. +_NMP_PRINCIPAL_ENVVAR = NMP_PRINCIPAL_ENVVAR # --------------------------------------------------------------------------- @@ -162,6 +166,40 @@ def _workload_identity_headers(*, internal: bool) -> dict[str, str]: return {_INTERNAL_REQUEST_HEADER: "true"} if internal else {} +def _workload_identity_token_provider(base_url: str) -> TokenProvider: + from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider + + subject_token_file = workload_identity_token_file_from_env() + if subject_token_file is None: + raise RuntimeError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} is not set") + return resolve_workload_exchange_provider(base_url=base_url, subject_token_file=subject_token_file) + + +def _ensure_no_trusted_headers_for_workload_identity( + *, + as_service: str | None, + on_behalf_of: str | None, + on_behalf_of_headers: dict[str, str] | None, +) -> None: + if as_service is not None or on_behalf_of is not None or on_behalf_of_headers is not None: + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} cannot be combined with trusted principal headers") + + +def _task_on_behalf_of_headers() -> dict[str, str] | None: + principal = _read_principal_from_env() + return _on_behalf_of_headers(principal) if principal is not None else None + + +def _warn_missing_task_principal(*, service_name: str, async_sdk: bool) -> None: + qualifier = "async task SDK" if async_sdk else "task SDK" + logger.warning( + "%s not set; %s will authenticate as service:%s without on-behalf-of delegation", + _NMP_PRINCIPAL_ENVVAR, + qualifier, + service_name, + ) + + class DefaultSDKProvider: """Env-var-based provider that ships with the plugin package. @@ -172,75 +210,127 @@ class DefaultSDKProvider: def get_task_sdk(self, service_name: str) -> NeMoPlatform: if is_workload_identity_token_file_set(): - return NeMoPlatform( - base_url=self._base_url(), - default_headers=_workload_identity_headers(internal=True), - ) - - headers: dict[str, str] = { - "X-NMP-Principal-Id": f"service:{service_name}", - _INTERNAL_REQUEST_HEADER: "true", - } - - principal = _read_principal_from_env() - if principal is not None: - headers.update(_on_behalf_of_headers(principal)) - else: - logger.warning( - "%s not set; task SDK will authenticate as service:%s without on-behalf-of delegation", - _NMP_PRINCIPAL_ENVVAR, - service_name, - ) - - return NeMoPlatform( - base_url=self._base_url(), - default_headers=headers, + require_workload_identity_without_principal_env() + return self._make_workload_identity_sdk(NeMoPlatform, internal=True) + + on_behalf_of_headers = _task_on_behalf_of_headers() + if on_behalf_of_headers is None: + _warn_missing_task_principal(service_name=service_name, async_sdk=False) + return self._make_sdk( + NeMoPlatform, + as_service=service_name, + internal=True, + on_behalf_of_headers=on_behalf_of_headers, ) def get_async_task_sdk(self, service_name: str) -> AsyncNeMoPlatform: # Async mirror of get_task_sdk: identical headers (service principal, # internal marker, and full on-behalf-of id/email/groups), async client. if is_workload_identity_token_file_set(): - return AsyncNeMoPlatform( - base_url=self._base_url(), - default_headers=_workload_identity_headers(internal=True), - ) - - headers: dict[str, str] = { - "X-NMP-Principal-Id": f"service:{service_name}", - _INTERNAL_REQUEST_HEADER: "true", - } + require_workload_identity_without_principal_env() + return self._make_workload_identity_sdk(AsyncNeMoPlatform, internal=True) + + on_behalf_of_headers = _task_on_behalf_of_headers() + if on_behalf_of_headers is None: + _warn_missing_task_principal(service_name=service_name, async_sdk=True) + return self._make_sdk( + AsyncNeMoPlatform, + as_service=service_name, + internal=True, + on_behalf_of_headers=on_behalf_of_headers, + ) - principal = _read_principal_from_env() - if principal is not None: - headers.update(_on_behalf_of_headers(principal)) - else: - logger.warning( - "%s not set; async task SDK will authenticate as service:%s without on-behalf-of delegation", - _NMP_PRINCIPAL_ENVVAR, - service_name, - ) + @overload + def _make_sdk( + self, + cls: type[NeMoPlatform], + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + on_behalf_of_headers: dict[str, str] | None = None, + ) -> NeMoPlatform: ... - return AsyncNeMoPlatform( - base_url=self._base_url(), - default_headers=headers, - ) + @overload + def _make_sdk( + self, + cls: type[AsyncNeMoPlatform], + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + on_behalf_of_headers: dict[str, str] | None = None, + ) -> AsyncNeMoPlatform: ... def _make_sdk( self, - cls: type[_SDKT], + cls: type[NeMoPlatform] | type[AsyncNeMoPlatform], *, as_service: str | None = None, internal: bool = False, on_behalf_of: str | None = None, - ) -> _SDKT: - if as_service is None and on_behalf_of is None and is_workload_identity_token_file_set(): - headers = _workload_identity_headers(internal=internal) - return cls(base_url=self._base_url(), default_headers=headers or None) + on_behalf_of_headers: dict[str, str] | None = None, + ) -> NeMoPlatform | AsyncNeMoPlatform: + if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() + _ensure_no_trusted_headers_for_workload_identity( + as_service=as_service, + on_behalf_of=on_behalf_of, + on_behalf_of_headers=on_behalf_of_headers, + ) + if cls is NeMoPlatform: + return self._make_workload_identity_sdk(NeMoPlatform, internal=internal) + if cls is AsyncNeMoPlatform: + return self._make_workload_identity_sdk(AsyncNeMoPlatform, internal=internal) + raise TypeError(f"Unsupported SDK class for workload identity: {cls!r}") headers = self._build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) + if on_behalf_of_headers: + headers.update(on_behalf_of_headers) return cls(base_url=self._base_url(), default_headers=headers or None) + @overload + def _make_workload_identity_sdk( + self, + cls: type[NeMoPlatform], + *, + internal: bool, + ) -> NeMoPlatform: ... + + @overload + def _make_workload_identity_sdk( + self, + cls: type[AsyncNeMoPlatform], + *, + internal: bool, + ) -> AsyncNeMoPlatform: ... + + def _make_workload_identity_sdk( + self, + cls: type[NeMoPlatform] | type[AsyncNeMoPlatform], + *, + internal: bool, + ) -> NeMoPlatform | AsyncNeMoPlatform: + base_url = self._base_url() + headers = _workload_identity_headers(internal=internal) + token_provider = _workload_identity_token_provider(base_url) + auth = TokenProviderAuth(token_provider) + if cls is AsyncNeMoPlatform: + return AsyncNeMoPlatform( + base_url=base_url, + default_headers=headers or None, + http_client=DefaultAsyncHttpxClient(auth=auth), + token_provider=token_provider, + ) + if cls is NeMoPlatform: + return NeMoPlatform( + base_url=base_url, + default_headers=headers or None, + http_client=DefaultHttpxClient(auth=auth), + token_provider=token_provider, + ) + raise TypeError(f"Unsupported SDK class for workload identity: {cls!r}") + def get_platform_sdk( self, *, @@ -391,9 +481,8 @@ def get_async_task_sdk(service_name: str) -> AsyncNeMoPlatform: ``NMP_PRINCIPAL`` is set, on behalf of the job creator with the full delegated identity (on-behalf-of id, email, and groups) — wire-identical to :func:`get_task_sdk`. - A dedicated provider method (not a wrapper over :func:`get_async_platform_sdk`) so each provider - mirrors its own sync :meth:`SDKProvider.get_task_sdk` exactly; the platform provider routes URLs - and reuses its shared async client, the default provider uses env-var headers. + Delegates to the active provider so each provider can preserve its own SDK + construction lifecycle while keeping task and platform SDK auth policy aligned. """ return _resolve_provider().get_async_task_sdk(service_name) @@ -450,7 +539,8 @@ def get_forwarding_headers(sdk: NeMoPlatform | AsyncNeMoPlatform) -> dict[str, s per-request headers, pass a request-scoped SDK built with ``sdk.with_options(set_default_headers=...)``. """ - # _custom_headers holds the headers passed at construction time - # (service principal, internal marker, on-behalf-of, OTEL, etc.) - # — everything the platform SDK factory injects. - return dict(sdk._custom_headers) + return { + name: value + for name, value in sdk.default_headers.items() + if name.startswith("X-NMP-") and isinstance(value, str) + } diff --git a/packages/nemo_platform_plugin/tests/client/test_adapter.py b/packages/nemo_platform_plugin/tests/client/test_adapter.py index 7f82c0e985..b4c536adf3 100644 --- a/packages/nemo_platform_plugin/tests/client/test_adapter.py +++ b/packages/nemo_platform_plugin/tests/client/test_adapter.py @@ -4,11 +4,24 @@ from __future__ import annotations import httpx -from nemo_platform import NeMoPlatform +import pytest +from nemo_platform import AsyncNeMoPlatform, 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 +from nemo_platform_plugin.workspaces.client import AsyncWorkspacesClient, WorkspacesClient + + +class _Provider: + def __init__(self, token: str) -> None: + self._token = token + + def get_access_token(self) -> str: + return self._token + + async def get_access_token_async(self) -> str: + return self._token def test_client_from_platform_preserves_stainless_retry_policy() -> None: @@ -32,26 +45,6 @@ def test_client_from_platform_preserves_stainless_retry_policy() -> None: ) -def test_client_from_platform_prefers_platform_request_router() -> None: - http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) - platform = NeMoPlatform( - base_url="http://gateway", - workspace="default", - http_client=http_client, - ) - - class RequestRouter: - def resolve(self, url: str) -> str: - return url.replace("http://gateway/apis/jobs", "http://127.0.0.1:8080/apis/jobs") - - platform._nmp_request_router = RequestRouter() # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] - - client = client_from_platform(platform, JobsClient) - - request = endpoints.list_steps(workspace="default", name="job-1") - assert client._resolve_path(request) == ("http://127.0.0.1:8080/apis/jobs/v2/workspaces/default/jobs/job-1/steps") - - def test_client_from_platform_falls_back_to_sdk_prepare_url() -> None: http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) platform = NeMoPlatform( @@ -132,3 +125,75 @@ def test_client_from_platform_carries_disabled_timeout() -> None: # Not the transport's 60s: httpx reads an all-None Timeout as "wait forever". assert client._timeout == httpx.Timeout(None) + + +def test_client_from_platform_preserves_token_provider_auth() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + request=request, + json={ + "id": "workspace-id", + "name": "default", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", + }, + ) + + http_client = httpx.Client(transport=httpx.MockTransport(handler)) + platform = NeMoPlatform( + base_url="http://test", + workspace="default", + default_headers={"X-NMP-Internal": "true"}, + http_client=http_client, + token_provider=_Provider("adapter-token"), + ) + + client = client_from_platform(platform, WorkspacesClient) + workspace = client.get_workspace(name="default").data() + + assert workspace.name == "default" + assert requests[0].headers["Authorization"] == "Bearer adapter-token" + assert requests[0].headers["X-NMP-Internal"] == "true" + assert "X-NMP-Principal-Id" not in requests[0].headers + + +@pytest.mark.asyncio +async def test_async_client_from_platform_preserves_token_provider_auth() -> None: + requests: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + request=request, + json={ + "id": "workspace-id", + "name": "default", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", + }, + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + platform = AsyncNeMoPlatform( + base_url="http://test", + workspace="default", + default_headers={"X-NMP-Internal": "true"}, + http_client=http_client, + token_provider=_Provider("async-adapter-token"), + ) + + try: + client = client_from_platform(platform, AsyncWorkspacesClient) + workspace = (await client.get_workspace(name="default")).data() + + assert workspace.name == "default" + assert requests[0].headers["Authorization"] == "Bearer async-adapter-token" + assert requests[0].headers["X-NMP-Internal"] == "true" + assert "X-NMP-Principal-Id" not in requests[0].headers + finally: + await http_client.aclose() diff --git a/packages/nemo_platform_plugin/tests/client/test_tls.py b/packages/nemo_platform_plugin/tests/client/test_tls.py new file mode 100644 index 0000000000..e025a50e7b --- /dev/null +++ b/packages/nemo_platform_plugin/tests/client/test_tls.py @@ -0,0 +1,43 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, httpx_tls_config_from_env + + +def test_httpx_tls_config_from_env_defaults_to_certificate_validation(monkeypatch): + monkeypatch.delenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, raising=False) + + assert httpx_tls_config_from_env() == {} + + +def test_httpx_tls_config_from_env_ignores_blank_ca_bundle(monkeypatch): + monkeypatch.setenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, " ") + + assert httpx_tls_config_from_env() == {} + + +def test_httpx_tls_config_from_env_uses_custom_ca_bundle(monkeypatch): + monkeypatch.setenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "/tmp/nemo-ca.pem") + + assert httpx_tls_config_from_env() == {"verify": "/tmp/nemo-ca.pem"} + + +def test_httpx_tls_config_from_env_uses_env_overlay(monkeypatch): + monkeypatch.setenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "/tmp/default-ca.pem") + + assert httpx_tls_config_from_env({NMP_CLIENT_SSL_CERT_FILE_ENVVAR: "/tmp/override-ca.pem"}) == { + "verify": "/tmp/override-ca.pem" + } + + +def test_httpx_tls_config_from_env_uses_first_nonblank_configured_envvar(monkeypatch): + monkeypatch.delenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, raising=False) + + assert httpx_tls_config_from_env( + { + NMP_CLIENT_SSL_CERT_FILE_ENVVAR: " ", + "REQUESTS_CA_BUNDLE": "/tmp/requests-ca.pem", + "SSL_CERT_FILE": "/tmp/ssl-ca.pem", + }, + cert_file_envvars=(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "REQUESTS_CA_BUNDLE", "SSL_CERT_FILE"), + ) == {"verify": "/tmp/requests-ca.pem"} diff --git a/packages/nemo_platform_plugin/tests/test_client_auth.py b/packages/nemo_platform_plugin/tests/test_client_auth.py index b6785acb4e..822c060447 100644 --- a/packages/nemo_platform_plugin/tests/test_client_auth.py +++ b/packages/nemo_platform_plugin/tests/test_client_auth.py @@ -208,11 +208,14 @@ def test_constructor_workload_exchange_does_not_override_principal_header(self, class TestAsyncNemoClientAuth: @respx.mock def test_async_provider_called(self): - """AsyncNemoClient(auth=AsyncProvider()) calls async get_access_token().""" + """AsyncNemoClient(auth=AsyncProvider()) calls async get_access_token_async().""" route = respx.get("http://localhost:8080/test").mock(return_value=httpx.Response(200, json={"ok": True})) class AsyncProvider: - async def get_access_token(self) -> str: + def get_access_token(self) -> str: + raise AssertionError("sync token path should not be used") + + async def get_access_token_async(self) -> str: return "async-token" client = AsyncNemoClient(base_url="http://localhost:8080", auth=AsyncProvider()) @@ -420,7 +423,6 @@ def test_token_exchange_grant_sends_rfc8693_request(self, monkeypatch): "scope": "openid email groups", }, timeout=5.0, - verify=True, ) def test_token_exchange_grant_uses_docker_opaque_subject_token_type(self, monkeypatch): @@ -926,16 +928,16 @@ def test_async_client_uses_workload_exchange_provider_from_env(self, monkeypatch subject_token_file=subject_token_file, ) - def test_workload_exchange_provider_does_not_override_explicit_service(self, monkeypatch, tmp_path): + def test_workload_exchange_provider_rejects_explicit_service(self, monkeypatch, tmp_path): subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token\n", encoding="utf-8") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) monkeypatch.delenv("NMP_PRINCIPAL", raising=False) with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider: - client = get_nemo_client(as_service="jobs") + with pytest.raises(ValueError, match="trusted principal headers"): + get_nemo_client(as_service="jobs") - assert client._auth is None resolve_provider.assert_not_called() diff --git a/packages/nemo_platform_plugin/tests/test_client_provider.py b/packages/nemo_platform_plugin/tests/test_client_provider.py index 0c043f136c..d8127762f4 100644 --- a/packages/nemo_platform_plugin/tests/test_client_provider.py +++ b/packages/nemo_platform_plugin/tests/test_client_provider.py @@ -3,9 +3,8 @@ """Tests for :mod:`nemo_platform_plugin.client_provider`. -Covers the env-var default provider and the provider/entry-point resolution -seam. The rich platform provider (``nmp.common.client_factory``) is tested in -``packages/nmp_common/tests/client_factory``. +Covers the env-var default provider and optional provider/entry-point resolution. +Platform code should prefer ``sdk_provider`` plus ``client_from_platform``. """ from __future__ import annotations @@ -15,6 +14,7 @@ import pytest from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR from nemo_platform_plugin.client_provider import ( DefaultNemoClientProvider, NemoClientProvider, @@ -174,6 +174,24 @@ def test_async_workspace_passthrough(self, monkeypatch): client = DefaultNemoClientProvider().get_async_nemo_client(workspace="team-a") assert client.workspace == "team-a" + def test_workload_identity_rejects_principal_env(self, monkeypatch, tmp_path): + token_file = tmp_path / "token" + token_file.write_text("subject-token", encoding="utf-8") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@ex.com"})) + + with pytest.raises(ValueError, match="mutually exclusive"): + DefaultNemoClientProvider().get_nemo_client() + + def test_workload_identity_rejects_trusted_principal_headers(self, monkeypatch, tmp_path): + token_file = tmp_path / "token" + token_file.write_text("subject-token", encoding="utf-8") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with pytest.raises(ValueError, match="trusted principal headers"): + DefaultNemoClientProvider().get_nemo_client(as_service="evaluator", internal=True) + # --------------------------------------------------------------------------- # Provider resolution @@ -335,6 +353,14 @@ def get_async_nemo_client(self, **kwargs): captured.update(kwargs) return AsyncNemoClient(base_url="http://x") + def get_task_nemo_client(self, service_name, **kwargs): + captured.update(kwargs) + return NemoClient(base_url="http://x") + + def get_async_task_nemo_client(self, service_name, **kwargs): + captured.update(kwargs) + return AsyncNemoClient(base_url="http://x") + set_nemo_client_provider(_CapturingProvider()) get_nemo_client(as_service="svc", internal=True, on_behalf_of="u@x", workspace="ws1") assert captured == {"as_service": "svc", "internal": True, "on_behalf_of": "u@x", "workspace": "ws1"} diff --git a/packages/nemo_platform_plugin/tests/test_dependencies.py b/packages/nemo_platform_plugin/tests/test_dependencies.py index 5b5bef300c..a9734882bc 100644 --- a/packages/nemo_platform_plugin/tests/test_dependencies.py +++ b/packages/nemo_platform_plugin/tests/test_dependencies.py @@ -8,7 +8,7 @@ def test_get_nemo_client_requires_platform_override() -> None: - with pytest.raises(RuntimeError, match=r"get_nemo_client\(\) was called without being overridden"): + with pytest.raises(RuntimeError, match=r"get_nemo_client\(\) is not wired by platform services"): get_nemo_client() diff --git a/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py b/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py index f31021c72f..8bc4ec4657 100644 --- a/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py +++ b/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py @@ -129,7 +129,7 @@ def test_task_client_uses_workload_identity(monkeypatch, tmp_path, _stub_workloa token_file = tmp_path / "token" token_file.write_text("subject-token") monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) - monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) # must be ignored in WI mode + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) client = get_task_nemo_client("evaluator") @@ -143,10 +143,21 @@ def test_task_client_uses_workload_identity(monkeypatch, tmp_path, _stub_workloa assert client._default_headers.get("X-NMP-Internal") == "true" +def test_task_client_rejects_workload_identity_with_principal(monkeypatch, tmp_path): + token_file = tmp_path / "token" + token_file.write_text("subject-token") + monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) + + with pytest.raises(ValueError, match="mutually exclusive"): + get_task_nemo_client("evaluator") + + async def test_async_task_client_uses_workload_identity(monkeypatch, tmp_path, _stub_workload_exchange): token_file = tmp_path / "token" token_file.write_text("subject-token") monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) client = get_async_task_nemo_client("evaluator") assert isinstance(client._auth, _FakeExchangeProvider) diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index c4ff61c53b..fe7d9328b1 100644 --- a/packages/nemo_platform_plugin/tests/test_sdk_provider.py +++ b/packages/nemo_platform_plugin/tests/test_sdk_provider.py @@ -6,6 +6,7 @@ from __future__ import annotations import json +import logging from unittest.mock import patch import pytest @@ -21,6 +22,30 @@ ) +class _FakeExchangeProvider: + def get_access_token(self) -> str: + return "exchanged-token" + + async def get_access_token_async(self) -> str: + return "exchanged-token" + + +@pytest.fixture +def stub_workload_exchange(monkeypatch: pytest.MonkeyPatch) -> dict[str, str]: + captured: dict[str, str] = {} + + def _fake(*, base_url, subject_token_file): + captured["base_url"] = base_url + captured["subject_token_file"] = str(subject_token_file) + return _FakeExchangeProvider() + + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + _fake, + ) + return captured + + def _xnmp(sdk) -> dict[str, str]: return {k: v for k, v in sdk.default_headers.items() if k.startswith("X-NMP-")} @@ -128,15 +153,15 @@ def test_get_task_sdk_default_base_url(self, monkeypatch): sdk = provider.get_task_sdk("test") assert sdk.base_url == "http://localhost:8080" - def test_get_task_sdk_uses_workload_identity_when_token_file_configured(self, monkeypatch, tmp_path): + def test_get_task_sdk_uses_workload_identity_when_token_file_configured( + self, monkeypatch, tmp_path, caplog, stub_workload_exchange + ): subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv( - "NMP_PRINCIPAL", - json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), - ) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + caplog.set_level(logging.WARNING, logger="nemo_platform_plugin.sdk_provider") provider = DefaultSDKProvider() sdk = provider.get_task_sdk("evaluator") @@ -144,8 +169,17 @@ def test_get_task_sdk_uses_workload_identity_when_token_file_configured(self, mo assert sdk.default_headers["X-NMP-Internal"] == "true" assert "X-NMP-Principal-Id" not in sdk.default_headers assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers + request = sdk._client.build_request("GET", "http://test:9090/apis/entities/v2/workspaces/default") + auth = sdk._client.auth + assert auth is not None + response = next(auth.sync_auth_flow(request)) + assert response is request + assert request.headers["Authorization"] == "Bearer exchanged-token" finally: sdk.close() + assert stub_workload_exchange["base_url"] == "http://test:9090" + assert stub_workload_exchange["subject_token_file"] == str(subject_token_file) + assert "will authenticate as service:evaluator without on-behalf-of delegation" not in caplog.text def test_get_platform_sdk_as_service(self, monkeypatch): monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") @@ -157,6 +191,29 @@ def test_get_platform_sdk_as_service(self, monkeypatch): assert sdk.default_headers["X-NMP-Principal-Id"] == "service:my-svc" assert sdk.default_headers["X-NMP-Internal"] == "true" + def test_get_platform_sdk_rejects_workload_identity_with_principal_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), + ) + + with pytest.raises(ValueError, match="mutually exclusive"): + DefaultSDKProvider().get_platform_sdk() + + def test_get_platform_sdk_rejects_workload_identity_with_service_headers(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with pytest.raises(ValueError, match="trusted principal headers"): + DefaultSDKProvider().get_platform_sdk(as_service="my-svc", internal=True) + def test_get_platform_sdk_on_behalf_of(self, monkeypatch): monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") monkeypatch.delenv("NMP_PRINCIPAL", raising=False) @@ -218,23 +275,44 @@ def test_parity_with_sync_task_sdk(self, monkeypatch): assert _xnmp(provider.get_async_task_sdk("evaluator")) == _xnmp(provider.get_task_sdk("evaluator")) @pytest.mark.asyncio - async def test_async_task_sdk_uses_workload_identity_when_token_file_configured(self, monkeypatch, tmp_path): + async def test_async_task_sdk_uses_workload_identity_when_token_file_configured( + self, monkeypatch, tmp_path, stub_workload_exchange + ): subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv( - "NMP_PRINCIPAL", - json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), - ) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) sdk = DefaultSDKProvider().get_async_task_sdk("evaluator") try: assert sdk.default_headers["X-NMP-Internal"] == "true" assert "X-NMP-Principal-Id" not in sdk.default_headers assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers + request = sdk._client.build_request("GET", "http://test:9090/apis/entities/v2/workspaces/default") + auth = sdk._client.auth + assert auth is not None + response = await anext(auth.async_auth_flow(request)) + assert response is request + assert request.headers["Authorization"] == "Bearer exchanged-token" finally: await sdk.close() + assert stub_workload_exchange["base_url"] == "http://test:9090" + assert stub_workload_exchange["subject_token_file"] == str(subject_token_file) + + @pytest.mark.asyncio + async def test_async_task_sdk_rejects_workload_identity_with_principal_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), + ) + + with pytest.raises(ValueError, match="mutually exclusive"): + DefaultSDKProvider().get_async_task_sdk("evaluator") # --------------------------------------------------------------------------- @@ -246,9 +324,15 @@ class _CustomProvider: def get_task_sdk(self, service_name: str) -> NeMoPlatform: return NeMoPlatform(base_url="http://custom:1234") + def get_async_task_sdk(self, service_name: str) -> AsyncNeMoPlatform: + return AsyncNeMoPlatform(base_url="http://custom:1234") + def get_platform_sdk(self, **kwargs) -> NeMoPlatform: return NeMoPlatform(base_url="http://custom:1234") + def get_async_platform_sdk(self, **kwargs) -> AsyncNeMoPlatform: + return AsyncNeMoPlatform(base_url="http://custom:1234") + class _FakeEntryPoint: def __init__( diff --git a/packages/nmp_common/pyproject.toml b/packages/nmp_common/pyproject.toml index ab7f6f2655..4a2c337f7f 100644 --- a/packages/nmp_common/pyproject.toml +++ b/packages/nmp_common/pyproject.toml @@ -72,8 +72,5 @@ dev-dependencies = [ [project.entry-points."nemo.sdk_provider"] platform = "nmp.common.sdk_factory:PlatformSDKProvider" -[project.entry-points."nemo.client_provider"] -platform = "nmp.common.client_factory:PlatformNemoClientProvider" - [tool.hatch.build.targets.wheel] packages = ["src/nmp_common", "src/nmp"] diff --git a/packages/nmp_common/src/nmp/common/auth/access_key_lifecycle.py b/packages/nmp_common/src/nmp/common/auth/access_key_lifecycle.py index 864d6bacdd..25deb11e24 100644 --- a/packages/nmp_common/src/nmp/common/auth/access_key_lifecycle.py +++ b/packages/nmp_common/src/nmp/common/auth/access_key_lifecycle.py @@ -6,6 +6,7 @@ import logging import math import time +from collections.abc import Sequence import httpx import jwt @@ -20,7 +21,7 @@ from nmp.common.config import AuthConfig from .token_claims import TokenClaims -from .token_resolver import ResolvedBearerToken +from .token_resolver import ResolvedBearerToken, ResolvedTokenKind logger = logging.getLogger(__name__) ACCESS_KEY_LIFECYCLE_CIRCUIT_FAILURE_THRESHOLD = 3 @@ -47,18 +48,21 @@ def __init__(self, config: AuthConfig, http_client: httpx.AsyncClient | None = N self._failure_count = 0 self._circuit_open_until = 0.0 + async def aclose(self) -> None: + sdk = self._sdk + self._sdk = None + if sdk is not None: + await sdk.close() + def _get_sdk(self) -> AsyncNeMoPlatform: if self._sdk is None: - # Import lazily to avoid an auth -> SDK factory import cycle. The - # factory attaches PlatformRequestRouter, which owns service - # discovery and transport selection for /apis/auth requests. - from nmp.common.sdk_factory import get_async_platform_sdk, with_options_preserving_request_router - - sdk = get_async_platform_sdk(http_client=self._http_client) - self._sdk = with_options_preserving_request_router( - sdk, + # Import lazily to avoid an auth -> SDK factory import cycle. + from nmp.common.sdk_factory import get_async_platform_sdk + + self._sdk = get_async_platform_sdk( + http_client=self._http_client, max_retries=0, - _extra_kwargs={"_strict_response_validation": True}, + _strict_response_validation=True, ) return self._sdk @@ -85,6 +89,15 @@ def _unavailable(self, status_code: int, detail: str) -> AccessKeyLifecycleUnava async def authenticate(self, token: str) -> ResolvedBearerToken | None: """Return trusted access-key claims, or None when auth rejects the token.""" + return await self.authenticate_token(token, token_kinds=("access_key",)) + + async def authenticate_token( + self, + token: str, + *, + token_kinds: Sequence[ResolvedTokenKind], + ) -> ResolvedBearerToken | None: + """Return trusted claims for allowed auth-service token kinds.""" now = time.monotonic() if now < self._circuit_open_until: retry_after = self._retry_after() @@ -117,10 +130,10 @@ async def authenticate(self, token: str) -> ResolvedBearerToken | None: logger.error("Access-key lifecycle validation failed at %s: %s", exc.request.url, exc) raise self._unavailable(503, "Access-key lifecycle validation unavailable") from exc - if result.token_kind != "access_key": + if result.token_kind not in token_kinds: self._record_success() return None - if not result.jti: + if result.token_kind == "access_key" and not result.jti: logger.error("Access-key lifecycle validator returned an invalid success response") raise self._unavailable(503, "Access-key lifecycle validation unavailable") @@ -138,19 +151,20 @@ async def authenticate(self, token: str) -> ResolvedBearerToken | None: ) except jwt.PyJWTError: unverified_payload = {} + raw_claims = { + **unverified_payload, + "nmp_token_type": result.token_kind, + "sub": result.principal, + **({"jti": result.jti} if result.jti else {}), + **({"email": result.email} if result.email is not None else {}), + } return ResolvedBearerToken( claims=TokenClaims( subject=result.principal, email=result.email, groups=result.groups or [], scopes=result.scopes or [], - raw_claims={ - **unverified_payload, - "nmp_token_type": "access_key", - "jti": result.jti, - "sub": result.principal, - **({"email": result.email} if result.email is not None else {}), - }, + raw_claims=raw_claims, ), - token_kind="access_key", + token_kind=result.token_kind, ) diff --git a/packages/nmp_common/src/nmp/common/auth/access_keys.py b/packages/nmp_common/src/nmp/common/auth/access_keys.py index d07656ae96..792546b282 100644 --- a/packages/nmp_common/src/nmp/common/auth/access_keys.py +++ b/packages/nmp_common/src/nmp/common/auth/access_keys.py @@ -93,6 +93,13 @@ def clear_access_key_signing_key_cache() -> None: _ACCESS_KEY_JWKS_CLIENTS.clear() +async def close_access_key_jwks_clients() -> None: + clients = list(_ACCESS_KEY_JWKS_CLIENTS.values()) + _ACCESS_KEY_JWKS_CLIENTS.clear() + for client in clients: + await client.aclose() + + def _groups_claim_for_gateway_header(groups: list[str]) -> str | None: groups_claim = ",".join(group.strip() for group in groups if group.strip()) return groups_claim or None diff --git a/packages/nmp_common/src/nmp/common/auth/jwks.py b/packages/nmp_common/src/nmp/common/auth/jwks.py index 44e8eda346..cbabbd6d7e 100644 --- a/packages/nmp_common/src/nmp/common/auth/jwks.py +++ b/packages/nmp_common/src/nmp/common/auth/jwks.py @@ -8,8 +8,9 @@ import time from typing import Any +import httpx import jwt -from nmp.common import http_clients +from nemo_platform import DefaultAsyncHttpxClient from .loading_cache import AsyncCoalescingLoader @@ -46,9 +47,17 @@ def signing_jwk_from_jwks(token: str, jwks: dict[str, Any]) -> Any: class AsyncJWKSClient: """Async JWKS client with TTL caching and one refresh on unknown key IDs.""" - def __init__(self, jwks_uri: str, *, lifespan: int = DEFAULT_JWKS_CACHE_LIFESPAN) -> None: + def __init__( + self, + jwks_uri: str, + *, + lifespan: int = DEFAULT_JWKS_CACHE_LIFESPAN, + http_client: httpx.AsyncClient | None = None, + ) -> None: self._jwks_uri = jwks_uri self._lifespan = lifespan + self._owns_http_client = http_client is None + self._http_client: httpx.AsyncClient = http_client if http_client is not None else DefaultAsyncHttpxClient() self._jwks: dict[str, Any] | None = None self._jwks_cache_time = 0.0 self._unknown_kid_refresh_loader: AsyncCoalescingLoader[dict[str, Any]] = AsyncCoalescingLoader( @@ -76,6 +85,10 @@ def clear_cache(self) -> None: min_interval_seconds=UNKNOWN_KID_REFRESH_MIN_INTERVAL_SECONDS ) + async def aclose(self) -> None: + if self._owns_http_client: + await self._http_client.aclose() + async def _refresh_jwks_for_unknown_kid(self) -> dict[str, Any]: return await self._unknown_kid_refresh_loader.load( self._force_refresh_jwks, @@ -91,17 +104,21 @@ async def _force_refresh_jwks(self) -> dict[str, Any]: jwks, _ = await self._fetch_jwks(refresh=True) return jwks + async def _request_jwks(self, http_client: httpx.AsyncClient) -> dict[str, Any]: + response = await http_client.get(self._jwks_uri, timeout=10.0) + response.raise_for_status() + jwks = response.json() + if not isinstance(jwks, dict): + raise jwt.InvalidTokenError("JWKS response was not an object") + return jwks + async def _fetch_jwks(self, *, refresh: bool = False) -> tuple[dict[str, Any], bool]: now = time.monotonic() if self._jwks is not None and not refresh and self._lifespan > 0: if now - self._jwks_cache_time < self._lifespan: return self._jwks, True - response = await http_clients.shared_async_http_client().get(self._jwks_uri, timeout=10.0) - response.raise_for_status() - jwks = response.json() - if not isinstance(jwks, dict): - raise jwt.InvalidTokenError("JWKS response was not an object") + jwks = await self._request_jwks(self._http_client) validate_jwks(jwks) if self._lifespan > 0: self._jwks = jwks diff --git a/packages/nmp_common/src/nmp/common/auth/jwt.py b/packages/nmp_common/src/nmp/common/auth/jwt.py index de0801389f..c07dee71df 100644 --- a/packages/nmp_common/src/nmp/common/auth/jwt.py +++ b/packages/nmp_common/src/nmp/common/auth/jwt.py @@ -47,6 +47,12 @@ def __init__(self, config: AuthConfig): self._discovery_cache: JsonObject | None = None self._discovery_cache_time: float = 0.0 + async def aclose(self) -> None: + jwks_client = self._jwks_client + self._jwks_client = None + if jwks_client is not None: + await jwks_client.aclose() + async def _discover_oidc_config(self) -> JsonObject: """Fetch OIDC discovery document from issuer. diff --git a/packages/nmp_common/src/nmp/common/auth/middleware.py b/packages/nmp_common/src/nmp/common/auth/middleware.py index bfb3245784..ad70b111fa 100644 --- a/packages/nmp_common/src/nmp/common/auth/middleware.py +++ b/packages/nmp_common/src/nmp/common/auth/middleware.py @@ -24,7 +24,7 @@ from .dependencies import auth_client_context from .exceptions import InvalidPrincipalHeader, InvalidScopeFormatError from .models import Principal -from .token_resolver import ResolvedBearerToken, resolve_bearer_token +from .token_resolver import ResolvedBearerToken, ResolvedTokenKind, resolve_bearer_token logger = logging.getLogger(__name__) @@ -505,6 +505,11 @@ async def _handle_bearer_token_request(self, request: Request, call_next: Callab ) if resolved is None: + if self.config.oidc.workload_token_exchange_enabled: + resolved_or_error = await self._authenticate_workload_token_lifecycle(token) + if isinstance(resolved_or_error, Response): + return resolved_or_error + return await self._handle_resolved_bearer_token(request, call_next, resolved_or_error) if jwt_validator is None: logger.warning( "Bearer token rejected: OIDC is not configured and the token did not pass " @@ -531,9 +536,26 @@ def _access_key_lifecycle_error_response( async def _authenticate_access_key_lifecycle( self, token: str, + ) -> ResolvedBearerToken | Response: + return await self._authenticate_token_lifecycle(token, token_kinds=("access_key",)) + + async def _authenticate_workload_token_lifecycle( + self, + token: str, + ) -> ResolvedBearerToken | Response: + return await self._authenticate_token_lifecycle( + token, + token_kinds=("workload_access_token", "workload_subject_token"), + ) + + async def _authenticate_token_lifecycle( + self, + token: str, + *, + token_kinds: tuple[ResolvedTokenKind, ...], ) -> ResolvedBearerToken | Response: try: - resolved = await self._access_key_lifecycle.authenticate(token) + resolved = await self._access_key_lifecycle.authenticate_token(token, token_kinds=token_kinds) except AccessKeyLifecycleUnavailableError as exc: return self._access_key_lifecycle_error_response( exc.status_code, diff --git a/packages/nmp_common/src/nmp/common/client_factory.py b/packages/nmp_common/src/nmp/common/client_factory.py deleted file mode 100644 index 26d883a306..0000000000 --- a/packages/nmp_common/src/nmp/common/client_factory.py +++ /dev/null @@ -1,351 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Rich NemoClient factory backed by platform internals. - -This is the :class:`~nemo_platform_plugin.client.client.NemoClient` sibling of -:mod:`nmp.common.sdk_factory`. It builds typed clients that reuse the same -platform machinery the SDK factory uses: - -- base URL from :class:`~nmp.common.config.Configuration`; -- per-service URL routing via :class:`~nmp.common.sdk_factory.PlatformRequestRouter`; -- the shared sync/async HTTP clients (connection-pool + SSL-context reuse); -- principal / auth + internal-request headers via ``_get_default_headers``; -- OTEL trace-propagation headers captured on the current request. - -:class:`PlatformNemoClientProvider` is registered under the ``nemo.client_provider`` -entry-point group so :func:`nemo_platform_plugin.client_provider.get_nemo_client` -discovers it automatically whenever ``nmp-common`` is installed. -""" - -from __future__ import annotations - -import logging -import os -from collections.abc import Callable -from pathlib import Path - -import httpx -from nemo_platform_plugin.client.auth import TokenProvider -from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient -from nemo_platform_plugin.client.constants import ( - WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, - is_workload_identity_token_file_set, -) -from nmp.common.auth import Principal, principal_from_env -from nmp.common.config import get_platform_config -from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client -from nmp.common.observability import MARK_INTERNAL_REQUEST_HEADERS -from nmp.common.observability.otel import get_otel_headers -from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint -from nmp.common.sdk_factory import PlatformRequestRouter, _get_default_headers, _should_bootstrap_workload_identity - -logger = logging.getLogger(__name__) - -# Test-only: async HTTP client to use for NemoClient requests in test context. -# Set by test fixtures to route requests through the in-process test transport, -# mirroring ``nmp.common.sdk_factory._test_http_client``. -_test_http_client: httpx.AsyncClient | None = None - - -def _sync_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.Client | None, -) -> httpx.Client: - """Endpoint-aware sync client: honour an explicit client, else a UDS - transport for ``unix://`` endpoints, else the shared TCP client.""" - if http_client is not None: - return http_client - if endpoint.transport == "uds": - return endpoint.sync_http_client() - return shared_sync_http_client() - - -def _async_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.AsyncClient | None, -) -> httpx.AsyncClient: - """Async counterpart of :func:`_sync_http_client_for_endpoint`. - - Preserves the module-level ``_test_http_client`` fixture hook ahead of the - UDS / shared-client selection (mirrors ``sdk_factory``). - """ - if http_client is not None: - return http_client - if _test_http_client is not None: - return _test_http_client - if endpoint.transport == "uds": - return endpoint.async_http_client() - return shared_async_http_client() - - -def _workload_identity_auth(base_url: str) -> TokenProvider: - """Build a workload-identity token-exchange auth provider. - - Only call when :func:`is_workload_identity_token_file_set` is true. - """ - from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider - - token_file = os.environ[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] - return resolve_workload_exchange_provider(base_url=base_url, subject_token_file=Path(token_file)) - - -def _workload_identity_headers(internal: bool) -> dict[str, str]: - return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} - - -def _absolute_url(url: str) -> httpx.URL: - """Default resolver for the request router. - - :class:`NemoClient` hands its ``url_resolver`` the fully-qualified request - URL (``base_url`` + path), so — unlike the generated SDK's ``_prepare_url``, - which resolves a relative path — the router just needs to parse it. - """ - return httpx.URL(url) - - -def _platform_url_resolver() -> Callable[[str], httpx.URL]: - """Build a per-service URL router bound to the current platform config.""" - router = PlatformRequestRouter( - platform_config=get_platform_config(), - default_resolver=_absolute_url, - ) - return router.resolve - - -def _platform_headers( - as_service: str | None, - internal: bool, - on_behalf_of: str | Principal | None, -) -> dict[str, str]: - """Auth / internal headers plus OTEL trace-propagation headers. - - ``_get_default_headers`` supplies the principal + internal-request markers - (wire-identical to the SDK factory); ``get_otel_headers`` layers on the - trace-propagation context captured on the current request (empty outside a - request scope). - """ - headers = _get_default_headers(as_service, internal, on_behalf_of) - for name, value in get_otel_headers().items(): - normalized_name = name.lower() - if normalized_name == "x-nmp-internal" or normalized_name.startswith("x-nmp-principal-"): - continue - headers[name] = value - return headers - - -def get_nemo_client( - *, - as_service: str | None = None, - internal: bool = False, - on_behalf_of: str | Principal | None = None, - workspace: str | None = None, - http_client: httpx.Client | None = None, -) -> NemoClient: - """Build a sync :class:`NemoClient` configured with platform internals. - - Args: - as_service: If provided, authenticate as ``service:{as_service}``. - If ``None``, propagate the current request's / env principal. - internal: Mark requests as internal (service-to-service). - on_behalf_of: Principal (or id) to act on behalf of. Passing a - :class:`~nmp.common.auth.Principal` (rather than a bare id string) - is only reachable through this direct entry point; the plugin-facing - :class:`~nemo_platform_plugin.client_provider.NemoClientProvider` - protocol narrows ``on_behalf_of`` to ``str | None``. - workspace: Default workspace used to fill ``{workspace}`` path params. - http_client: Optional sync HTTP client; defaults to the shared client. - - Note: - OTEL trace-propagation headers are captured once, at construction, from - the current request context. Build a fresh client per request scope - rather than caching one across requests, or its ``traceparent`` will be - stale (mirrors ``get_platform_sdk``). - """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( - as_service=as_service, - on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - return NemoClient( - base_url=endpoint.connect_base_url, - workspace=workspace, - auth=_workload_identity_auth(endpoint.connect_base_url), - default_headers=_workload_identity_headers(internal) or None, - http_client=_sync_http_client_for_endpoint(endpoint, http_client), - url_resolver=_platform_url_resolver(), - ) - return NemoClient( - base_url=endpoint.connect_base_url, - workspace=workspace, - default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, - http_client=_sync_http_client_for_endpoint(endpoint, http_client), - url_resolver=_platform_url_resolver(), - ) - - -def get_async_nemo_client( - *, - as_service: str | None = None, - internal: bool = False, - on_behalf_of: str | Principal | None = None, - workspace: str | None = None, - http_client: httpx.AsyncClient | None = None, -) -> AsyncNemoClient: - """Async counterpart of :func:`get_nemo_client`. - - Uses the explicitly provided ``http_client`` (e.g. from a test fixture), then - the module-level ``_test_http_client`` fallback, then the shared async - client — mirroring ``nmp.common.sdk_factory.get_async_platform_sdk``. - """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( - as_service=as_service, - on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - return AsyncNemoClient( - base_url=endpoint.connect_base_url, - workspace=workspace, - auth=_workload_identity_auth(endpoint.connect_base_url), - default_headers=_workload_identity_headers(internal) or None, - http_client=_async_http_client_for_endpoint(endpoint, http_client), - url_resolver=_platform_url_resolver(), - ) - return AsyncNemoClient( - base_url=endpoint.connect_base_url, - workspace=workspace, - default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, - http_client=_async_http_client_for_endpoint(endpoint, http_client), - url_resolver=_platform_url_resolver(), - ) - - -def get_task_nemo_client( - service_name: str, - *, - workspace: str | None = None, - http_client: httpx.Client | None = None, -) -> NemoClient: - """Build a sync :class:`NemoClient` for use inside a task container. - - NemoClient counterpart of :func:`nmp.common.sdk_factory.get_task_sdk`: - reads the job creator's principal from ``NMP_PRINCIPAL`` and authenticates - as ``service:{service_name}`` while acting on behalf of that creator, or -- - when ``NMP_WORKLOAD_IDENTITY_TOKEN_FILE`` is set -- bootstraps - workload-identity bearer-token exchange (via :func:`get_nemo_client` with - ``internal=True``) instead of trusted ``X-NMP-*`` principal headers. - """ - if http_client is None and is_workload_identity_token_file_set(): - return get_nemo_client(internal=True, workspace=workspace) - if http_client is None: - http_client = resolve_platform_endpoint().sync_sdk_http_client() - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; task NemoClient will authenticate as service:%s without on-behalf-of delegation", - service_name, - ) - return get_nemo_client( - as_service=service_name, - internal=True, - on_behalf_of=principal.effective_principal if principal else None, - workspace=workspace, - http_client=http_client, - ) - - -def get_async_task_nemo_client( - service_name: str, - *, - workspace: str | None = None, - http_client: httpx.AsyncClient | None = None, -) -> AsyncNemoClient: - """Async counterpart of :func:`get_task_nemo_client`. Wire-identical.""" - if http_client is None and is_workload_identity_token_file_set(): - return get_async_nemo_client(internal=True, workspace=workspace) - if http_client is None: - http_client = resolve_platform_endpoint().async_sdk_http_client() - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; async task NemoClient will authenticate as service:%s without on-behalf-of delegation", - service_name, - ) - return get_async_nemo_client( - as_service=service_name, - internal=True, - on_behalf_of=principal.effective_principal if principal else None, - workspace=workspace, - http_client=http_client, - ) - - -# --------------------------------------------------------------------------- -# Entry-point provider for nemo_platform_plugin.client_provider -# --------------------------------------------------------------------------- - - -class PlatformNemoClientProvider: - """Rich :class:`~nemo_platform_plugin.client_provider.NemoClientProvider` - that uses platform internals (shared HTTP clients, URL routing, OTEL - headers, auth context). - - Registered as a ``nemo.client_provider`` entry-point so it is discovered - automatically when ``nmp-common`` is installed. - """ - - def get_nemo_client( - self, - *, - as_service: str | None = None, - internal: bool = False, - on_behalf_of: str | Principal | None = None, - workspace: str | None = None, - http_client: httpx.Client | None = None, - ) -> NemoClient: - return get_nemo_client( - as_service=as_service, - internal=internal, - on_behalf_of=on_behalf_of, - workspace=workspace, - http_client=http_client, - ) - - def get_async_nemo_client( - self, - *, - as_service: str | None = None, - internal: bool = False, - on_behalf_of: str | Principal | None = None, - workspace: str | None = None, - http_client: httpx.AsyncClient | None = None, - ) -> AsyncNemoClient: - return get_async_nemo_client( - as_service=as_service, - internal=internal, - on_behalf_of=on_behalf_of, - workspace=workspace, - http_client=http_client, - ) - - def get_task_nemo_client( - self, - service_name: str, - *, - workspace: str | None = None, - http_client: httpx.Client | None = None, - ) -> NemoClient: - return get_task_nemo_client(service_name, workspace=workspace, http_client=http_client) - - def get_async_task_nemo_client( - self, - service_name: str, - *, - workspace: str | None = None, - http_client: httpx.AsyncClient | None = None, - ) -> AsyncNemoClient: - return get_async_task_nemo_client(service_name, workspace=workspace, http_client=http_client) diff --git a/packages/nmp_common/src/nmp/common/http_clients.py b/packages/nmp_common/src/nmp/common/http_clients.py deleted file mode 100644 index 0a8934bb01..0000000000 --- a/packages/nmp_common/src/nmp/common/http_clients.py +++ /dev/null @@ -1,107 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Shared HTTP clients for SDK and internal requests. - -Ideally, most consumers get their HTTP client from DependencyProvider, which -manages per-service client lifecycle. However, code often constructs SDKs in -isolation (tasks, controllers, background jobs) without access to a -DependencyProvider. These shared clients prevent each SDK instantiation from -creating a new HTTP client, avoiding connection pool and SSL context overhead. - -Used by: -- get_platform_sdk() / get_async_platform_sdk() for tasks, controllers, and - other code outside DependencyProvider context -- Platform shutdown cleanup via close_cached_http_clients() - -The wrapper classes make close()/aclose() a no-op so that SDK code can safely -call sdk.close() without affecting other users of the shared client. Actual -cleanup happens at shutdown via close_cached_http_clients(). -""" - -import asyncio -from functools import cache -from typing import cast - -import httpx -from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient - - -class _SharedAsyncHttpClient(DefaultAsyncHttpxClient): - """Shared async HTTP client that ignores aclose() calls.""" - - async def aclose(self) -> None: - pass - - async def _real_close(self) -> None: - await super().aclose() - - -class _SharedSyncHttpClient(DefaultHttpxClient): - """Shared sync HTTP client that ignores close() calls.""" - - def close(self) -> None: - pass - - def _real_close(self) -> None: - super().close() - - -_shared_async_http_clients: dict[asyncio.AbstractEventLoop, _SharedAsyncHttpClient] = {} - - -def shared_async_http_client() -> httpx.AsyncClient: - """Get the shared async HTTP client for SDK requests. - - Returns a cached async HTTP client scoped to the current event loop. - Use this when creating SDK instances that need to share a connection - pool and SSL context within a single loop. - - If called when no event loop is currently running, return a fresh async - client instead of caching one globally. This keeps sync setup code safe - without binding a shared client to an arbitrary thread-local loop. - - The returned client ignores aclose() calls - cleanup happens at shutdown - via close_cached_http_clients(). This allows SDKs to safely call close() - without breaking other users of the shared client. - """ - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return DefaultAsyncHttpxClient() - client = _shared_async_http_clients.get(loop) - if client is None or client.is_closed: - client = _SharedAsyncHttpClient() - _shared_async_http_clients[loop] = client - return client - - -@cache -def shared_sync_http_client() -> httpx.Client: - """Get the shared sync HTTP client for SDK requests. - - Returns a cached sync HTTP client with SDK-compatible defaults. - Use this when creating SDK instances that need to share a global - connection pool and SSL context. - - The returned client ignores close() calls - cleanup happens at shutdown - via close_cached_http_clients(). This allows SDKs to safely call close() - without breaking other users of the shared client. - """ - return _SharedSyncHttpClient() - - -async def close_shared_http_clients() -> None: - """Close shared HTTP clients during graceful shutdown. - - Called from the platform lifespan shutdown to clean up module-level - shared clients used by non-service code (tasks, jobs, legacy callers). - """ - loop = asyncio.get_running_loop() - async_client = _shared_async_http_clients.pop(loop, None) - if async_client is not None and not async_client.is_closed: - await async_client._real_close() - - sync_client = cast(_SharedSyncHttpClient, shared_sync_http_client()) - sync_client._real_close() - shared_sync_http_client.cache_clear() diff --git a/packages/nmp_common/src/nmp/common/immutable_http_client.py b/packages/nmp_common/src/nmp/common/immutable_http_client.py new file mode 100644 index 0000000000..ed29d7623f --- /dev/null +++ b/packages/nmp_common/src/nmp/common/immutable_http_client.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Immutability helpers for SDK-owned HTTP clients. + +NeMo Platform SDK instances reuse their underlying httpx client when callers +derive scoped SDKs via ``with_options()`` or pass the SDK into typed plugin +clients. Those clients must be created with their required transport-level +configuration and then left alone. + +Caller-specific request configuration, including auth headers, belongs on the +SDK instance or on a separate explicit client. Mutating an SDK-owned client +after it has been handed out is a bug because that state can leak into derived +SDKs or requests. These wrappers make those bugs fail immediately. +""" + +from types import MappingProxyType +from typing import NoReturn + +import httpx +from httpx._types import CookieTypes, HeaderTypes +from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient + +_IMMUTABLE_CLIENT_ATTRS = { + "_base_url", + "_cookies", + "_event_hooks", + "_headers", + "_params", + "_timeout", + "_trust_env", + "base_url", + "cookies", + "event_hooks", + "follow_redirects", + "headers", + "max_redirects", + "params", + "timeout", + "trust_env", +} + + +def _raise_immutable_client_mutation_error() -> NoReturn: + raise TypeError( + "SDK HTTP clients are immutable. Pass per-SDK options as SDK constructor " + "arguments, or pass a separate httpx client configured for that use case." + ) + + +class _ImmutableHeaders(httpx.Headers): + def __setitem__(self, key: str, value: str) -> None: + _raise_immutable_client_mutation_error() + + def __delitem__(self, key: str) -> None: + _raise_immutable_client_mutation_error() + + def clear(self) -> None: + _raise_immutable_client_mutation_error() + + def pop(self, key: str, default: object = None) -> str: + _raise_immutable_client_mutation_error() + + def popitem(self) -> tuple[str, str]: + _raise_immutable_client_mutation_error() + + def setdefault(self, key: str, default: str = "") -> str: + _raise_immutable_client_mutation_error() + + def update(self, headers: HeaderTypes | None = None) -> None: + _raise_immutable_client_mutation_error() + + +class _ImmutableCookies(httpx.Cookies): + def extract_cookies(self, response: httpx.Response) -> None: + # httpx normally persists response cookies on the client. SDK clients + # should not carry request/session state between derived SDK handles. + pass + + def set(self, name: str, value: str, domain: str = "", path: str = "/") -> None: + _raise_immutable_client_mutation_error() + + def delete( + self, + name: str, + domain: str | None = None, + path: str | None = None, + ) -> None: + _raise_immutable_client_mutation_error() + + def clear(self, domain: str | None = None, path: str | None = None) -> None: + _raise_immutable_client_mutation_error() + + def update(self, cookies: CookieTypes | None = None) -> None: + _raise_immutable_client_mutation_error() + + def __setitem__(self, name: str, value: str) -> None: + _raise_immutable_client_mutation_error() + + def __delitem__(self, name: str) -> None: + _raise_immutable_client_mutation_error() + + +class ImmutableHttpClientMixin: + """Mixin for httpx client subclasses that are immutable after construction.""" + + _immutable_http_client_frozen = False + + def __setattr__(self, name: str, value: object) -> None: + if self._immutable_http_client_frozen and name in _IMMUTABLE_CLIENT_ATTRS: + raise AttributeError( + "SDK HTTP clients are immutable. Pass a separate httpx client " + "when client-level configuration needs to differ." + ) + super().__setattr__(name, value) + + def _freeze_http_client(self) -> None: + # Assignment blocking is not enough because these attributes are mutable + # containers. Replace them with immutable versions before sharing the + # client through SDK clones or plugin adapters. + self._headers = _ImmutableHeaders(self._headers) + self._cookies = _ImmutableCookies(self._cookies) + self._event_hooks = MappingProxyType( + {hook_name: tuple(hooks) for hook_name, hooks in self._event_hooks.items()} + ) + self._immutable_http_client_frozen = True + + +class ImmutableHttpxClient(ImmutableHttpClientMixin, httpx.Client): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableAsyncHttpxClient(ImmutableHttpClientMixin, httpx.AsyncClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableDefaultHttpxClient(ImmutableHttpClientMixin, DefaultHttpxClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableDefaultAsyncHttpxClient(ImmutableHttpClientMixin, DefaultAsyncHttpxClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() diff --git a/packages/nmp_common/src/nmp/common/jobs/log_client.py b/packages/nmp_common/src/nmp/common/jobs/log_client.py index 6024cb533d..4d265189bc 100644 --- a/packages/nmp_common/src/nmp/common/jobs/log_client.py +++ b/packages/nmp_common/src/nmp/common/jobs/log_client.py @@ -9,6 +9,7 @@ """ import logging +from collections.abc import AsyncIterator from nemo_platform import AsyncNeMoPlatform from nemo_platform_plugin.client.adapter import client_from_platform @@ -35,9 +36,14 @@ def __init__(self, sdk: AsyncNeMoPlatform | None = None): sdk: AsyncNeMoPlatform SDK instance. If not provided, creates one using platform config. """ + self._owns_sdk = sdk is None self._sdk = sdk or get_async_platform_sdk() self._files_client = client_from_platform(self._sdk, AsyncFilesClient) + async def aclose(self) -> None: + if self._owns_sdk: + await self._sdk.close() + async def query_logs( self, fileset: str, @@ -81,6 +87,10 @@ async def query_logs( raise -def dep_job_logs_client() -> JobLogsClient: +async def dep_job_logs_client() -> AsyncIterator[JobLogsClient]: """FastAPI dependency for JobLogsClient.""" - return JobLogsClient() + client = JobLogsClient() + try: + yield client + finally: + await client.aclose() diff --git a/packages/nmp_common/src/nmp/common/jobs/result_manager.py b/packages/nmp_common/src/nmp/common/jobs/result_manager.py index abff02722f..6854897d5e 100644 --- a/packages/nmp_common/src/nmp/common/jobs/result_manager.py +++ b/packages/nmp_common/src/nmp/common/jobs/result_manager.py @@ -3,7 +3,7 @@ import os import tarfile -from typing import Literal, overload +from typing import Literal, cast, overload from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.jobs.constants import NEMO_JOB_WORKSPACE_ENVVAR @@ -60,21 +60,41 @@ def result_manager_factory( raise ValueError(f"{NEMO_JOB_WORKSPACE_ENVVAR} environment variable is not set") workspace = workspace_env + if is_async: + owns_files_sdk = files_sdk is None + if files_sdk is None: + files_sdk = get_async_platform_sdk() + if jobs_sdk is None: + jobs_sdk = files_sdk + async_files_sdk = cast(AsyncNeMoPlatform, files_sdk) + async_jobs_sdk = cast(AsyncNeMoPlatform, jobs_sdk) + return AsyncResultManager( + job_name=job_name, + workspace=workspace, + attempt_id=attempt_id, + file_manager_cls=AsyncFilesetFileManager, + files_sdk=async_files_sdk, + jobs_sdk=async_jobs_sdk, + owns_files_sdk=owns_files_sdk, + owns_jobs_sdk=False, + ) + + owns_files_sdk = files_sdk is None if files_sdk is None: - files_sdk = get_async_platform_sdk() if is_async else get_platform_sdk() - + files_sdk = get_platform_sdk() if jobs_sdk is None: - jobs_sdk = get_async_platform_sdk() if is_async else get_platform_sdk() - - file_manager_cls = AsyncFilesetFileManager if is_async else FilesetFileManager - result_manager_cls = AsyncResultManager if is_async else ResultManager - return result_manager_cls( + jobs_sdk = files_sdk + sync_files_sdk = cast(NeMoPlatform, files_sdk) + sync_jobs_sdk = cast(NeMoPlatform, jobs_sdk) + return ResultManager( job_name=job_name, workspace=workspace, attempt_id=attempt_id, - file_manager_cls=file_manager_cls, - files_sdk=files_sdk, - jobs_sdk=jobs_sdk, + file_manager_cls=FilesetFileManager, + files_sdk=sync_files_sdk, + jobs_sdk=sync_jobs_sdk, + owns_files_sdk=owns_files_sdk, + owns_jobs_sdk=False, ) @@ -92,24 +112,24 @@ async def download_from_result_info( in tests also affects download_from_result_info, preserving the old monkeypatch behavior. """ - if files_sdk is None: - files_sdk = get_async_platform_sdk() - mgr = result_manager_factory( job_name=job_name, workspace=workspace, files_sdk=files_sdk, ) - tmp_dir_path = await mgr.download_artifact(artifact_url=artifact_url) - filename = result_name + try: + tmp_dir_path = await mgr.download_artifact(artifact_url=artifact_url) + filename = result_name - if tmp_dir_path.path.is_dir(): - filename = f"{filename}.tar.gz" - tar_path = tmp_dir_path.tmp_dir / filename - with tarfile.open(tar_path, "w:gz") as tar: - tar.add(tmp_dir_path.path, arcname=os.path.basename(tmp_dir_path.path)) + if tmp_dir_path.path.is_dir(): + filename = f"{filename}.tar.gz" + tar_path = tmp_dir_path.tmp_dir / filename + with tarfile.open(tar_path, "w:gz") as tar: + tar.add(tmp_dir_path.path, arcname=os.path.basename(tmp_dir_path.path)) - tmp_dir_path.path = tar_path + tmp_dir_path.path = tar_path - return filename, tmp_dir_path + return filename, tmp_dir_path + finally: + await mgr.aclose() diff --git a/packages/nmp_common/src/nmp/common/platform_endpoint.py b/packages/nmp_common/src/nmp/common/platform_endpoint.py index b8392881c8..f0f049d999 100644 --- a/packages/nmp_common/src/nmp/common/platform_endpoint.py +++ b/packages/nmp_common/src/nmp/common/platform_endpoint.py @@ -5,17 +5,38 @@ from __future__ import annotations +import logging import os -from dataclasses import dataclass +import re +from collections.abc import Mapping +from dataclasses import dataclass, field from pathlib import Path -from typing import Literal +from types import MappingProxyType +from typing import Literal, cast import httpx from httpx._types import TimeoutTypes -from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from nmp.common.config import PlatformConfig +from nmp.common.immutable_http_client import ( + ImmutableAsyncHttpxClient, + ImmutableDefaultAsyncHttpxClient, + ImmutableDefaultHttpxClient, + ImmutableHttpxClient, +) UDS_BASE_URL = "http://nemo-platform.local" +logger = logging.getLogger(__name__) + + +def _empty_service_endpoints() -> Mapping[str, "PlatformEndpoint"]: + return MappingProxyType({}) + + +def _get_platform_config() -> PlatformConfig: + from nmp.common.config import Configuration + + return cast(PlatformConfig, Configuration.get_platform_config()) @dataclass(frozen=True) @@ -23,6 +44,12 @@ class PlatformEndpoint: connect_base_url: str socket_path: Path | None transport: Literal["tcp", "uds"] + service_pattern: re.Pattern[str] | None = field(default=None, repr=False, compare=False) + service_endpoints: Mapping[str, "PlatformEndpoint"] = field( + default_factory=_empty_service_endpoints, + repr=False, + compare=False, + ) def sync_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Client: if self.transport == "uds": @@ -48,40 +75,113 @@ def async_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Asy return httpx.AsyncClient(follow_redirects=True) return httpx.AsyncClient(follow_redirects=True, timeout=timeout) - def sync_sdk_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Client: + def sync_sdk_http_client( + self, + *, + timeout: TimeoutTypes | None = None, + ) -> httpx.Client: + if self.service_endpoints: + transport = _SyncPlatformEndpointRoutingTransport(endpoint=self) + if timeout is None: + return ImmutableDefaultHttpxClient(transport=transport) + return ImmutableDefaultHttpxClient(transport=transport, timeout=timeout) if self.transport == "uds": - return self.sync_http_client(timeout=timeout) + if self.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + transport = httpx.HTTPTransport(uds=str(self.socket_path)) + if timeout is None: + return ImmutableHttpxClient(transport=transport, follow_redirects=True) + return ImmutableHttpxClient(transport=transport, follow_redirects=True, timeout=timeout) if timeout is None: - return DefaultHttpxClient() - return DefaultHttpxClient(timeout=timeout) + return ImmutableDefaultHttpxClient() + return ImmutableDefaultHttpxClient(timeout=timeout) - def async_sdk_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.AsyncClient: + def async_sdk_http_client( + self, + *, + timeout: TimeoutTypes | None = None, + ) -> httpx.AsyncClient: + if self.service_endpoints: + transport = _AsyncPlatformEndpointRoutingTransport(endpoint=self) + if timeout is None: + return ImmutableDefaultAsyncHttpxClient(transport=transport) + return ImmutableDefaultAsyncHttpxClient(transport=transport, timeout=timeout) if self.transport == "uds": - return self.async_http_client(timeout=timeout) + if self.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + transport = httpx.AsyncHTTPTransport(uds=str(self.socket_path)) + if timeout is None: + return ImmutableAsyncHttpxClient(transport=transport, follow_redirects=True) + return ImmutableAsyncHttpxClient(transport=transport, follow_redirects=True, timeout=timeout) if timeout is None: - return DefaultAsyncHttpxClient() - return DefaultAsyncHttpxClient(timeout=timeout) + return ImmutableDefaultAsyncHttpxClient() + return ImmutableDefaultAsyncHttpxClient(timeout=timeout) + + def route_request_url(self, url: str | httpx.URL) -> "RoutedPlatformEndpointRequest": + """Resolve one outgoing SDK URL using this endpoint's fixed routing table.""" + + request_url = httpx.URL(url) + endpoint = self + service_name = "unknown" + + match = self.service_pattern.search(request_url.path) if self.service_pattern is not None else None + if match is not None: + service_name = match.group(1) + endpoint = self.service_endpoints.get(service_name) or self + routed_url = _url_for_endpoint(request_url, endpoint) + else: + routed_url = request_url + + logger.debug( + "Routing SDK URL to service endpoint" + if service_name != "unknown" + else "Routing SDK URL to default endpoint", + extra={ + "service": service_name, + "path": request_url.path, + "host": routed_url.host, + "port": routed_url.port, + "transport": endpoint.transport, + }, + ) + return RoutedPlatformEndpointRequest(url=routed_url, endpoint=endpoint) + + +@dataclass(frozen=True) +class RoutedPlatformEndpointRequest: + url: httpx.URL + endpoint: PlatformEndpoint def resolve_platform_endpoint(platform_config: PlatformConfig | None = None) -> PlatformEndpoint: """Resolve the default platform endpoint from ``NMP_BASE_URL`` / config.""" if platform_config is None: - from nmp.common.config import Configuration - - platform_config = Configuration.get_platform_config() - return parse_platform_endpoint(platform_config.base_url) + platform_config = _get_platform_config() + default_endpoint = parse_platform_endpoint(platform_config.base_url) + service_endpoints = { + service_name: resolve_service_endpoint(service_name, platform_config) + for service_name in sorted(_service_route_names(platform_config)) + } + return PlatformEndpoint( + connect_base_url=default_endpoint.connect_base_url, + socket_path=default_endpoint.socket_path, + transport=default_endpoint.transport, + service_pattern=platform_config.create_service_pattern(), + service_endpoints=MappingProxyType(service_endpoints), + ) def resolve_service_endpoint(service_name: str, platform_config: PlatformConfig | None = None) -> PlatformEndpoint: - """Resolve a service endpoint using ``NMP__URL`` before ``NMP_BASE_URL``.""" + """Resolve a service endpoint, keeping local services on the local process URL.""" if platform_config is None: - from nmp.common.config import Configuration - - platform_config = Configuration.get_platform_config() - env_name = f"NMP_{service_name.upper().replace('-', '_')}_URL" - endpoint = os.environ.get(env_name) or platform_config.get_service_url(service_name) + platform_config = _get_platform_config() + normalized_name = _normalize_service_name(service_name) + if normalized_name in {_normalize_service_name(local) for local in platform_config.get_services()}: + return parse_platform_endpoint(platform_config.get_service_url(normalized_name)) + env_name = _service_url_env_var_name(service_name) + endpoint = os.environ.get(env_name) or platform_config.get_service_url(normalized_name) return parse_platform_endpoint(endpoint) @@ -109,3 +209,99 @@ def _parse_unix_socket_path(endpoint: str) -> Path: if not raw_path.startswith("/"): raise ValueError(f"UDS endpoint must use an absolute socket path, got {endpoint!r}") return Path(raw_path) + + +def _service_route_names(platform_config: PlatformConfig) -> set[str]: + names = {_normalize_service_name(service_name) for service_name in platform_config.service_discovery} + names.update(_normalize_service_name(service_name) for service_name in platform_config.get_services()) + names.discard("") + return names + + +def _service_url_env_var_name(service_name: str) -> str: + return f"NMP_{_normalize_service_name(service_name).upper().replace('-', '_')}_URL" + + +def _normalize_service_name(service_name: str) -> str: + return service_name.strip().lower().replace("_", "-") + + +def _url_for_endpoint(url: httpx.URL, endpoint: PlatformEndpoint) -> httpx.URL: + if endpoint.transport == "uds": + return url.copy_with(scheme="http", host=httpx.URL(UDS_BASE_URL).host, port=None) + endpoint_url = httpx.URL(endpoint.connect_base_url) + path = url.path + if endpoint_url.path not in ("", "/"): + path = f"{endpoint_url.path.rstrip('/')}/{url.path.lstrip('/')}" + return url.copy_with(scheme=endpoint_url.scheme, host=endpoint_url.host, port=endpoint_url.port, path=path) + + +def _set_request_url(request: httpx.Request, url: httpx.URL) -> None: + if request.url == url: + return + request.url = url + if url.host: + request.headers["Host"] = url.netloc.decode("ascii") + + +def _uds_socket_paths(endpoint: PlatformEndpoint) -> frozenset[Path]: + paths: set[Path] = set() + for candidate in (endpoint, *endpoint.service_endpoints.values()): + if candidate.transport != "uds": + continue + if candidate.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + paths.add(candidate.socket_path) + return frozenset(paths) + + +class _SyncPlatformEndpointRoutingTransport(httpx.BaseTransport): + def __init__(self, *, endpoint: PlatformEndpoint) -> None: + self._endpoint = endpoint + self._tcp_transport = httpx.HTTPTransport(**httpx_tls_config_from_env()) + self._uds_transports = { + socket_path: httpx.HTTPTransport(uds=str(socket_path)) for socket_path in _uds_socket_paths(endpoint) + } + + def handle_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return self._transport_for_endpoint(routed.endpoint).handle_request(request) + + def close(self) -> None: + self._tcp_transport.close() + for transport in self._uds_transports.values(): + transport.close() + + def _transport_for_endpoint(self, endpoint: PlatformEndpoint) -> httpx.HTTPTransport: + if endpoint.transport == "tcp": + return self._tcp_transport + if endpoint.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + return self._uds_transports[endpoint.socket_path] + + +class _AsyncPlatformEndpointRoutingTransport(httpx.AsyncBaseTransport): + def __init__(self, *, endpoint: PlatformEndpoint) -> None: + self._endpoint = endpoint + self._tcp_transport = httpx.AsyncHTTPTransport(**httpx_tls_config_from_env()) + self._uds_transports = { + socket_path: httpx.AsyncHTTPTransport(uds=str(socket_path)) for socket_path in _uds_socket_paths(endpoint) + } + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return await self._transport_for_endpoint(routed.endpoint).handle_async_request(request) + + async def aclose(self) -> None: + await self._tcp_transport.aclose() + for transport in self._uds_transports.values(): + await transport.aclose() + + def _transport_for_endpoint(self, endpoint: PlatformEndpoint) -> httpx.AsyncHTTPTransport: + if endpoint.transport == "tcp": + return self._tcp_transport + if endpoint.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + return self._uds_transports[endpoint.socket_path] diff --git a/packages/nmp_common/src/nmp/common/sdk_factory.py b/packages/nmp_common/src/nmp/common/sdk_factory.py index 57a95891a7..c38ecfc2a6 100644 --- a/packages/nmp_common/src/nmp/common/sdk_factory.py +++ b/packages/nmp_common/src/nmp/common/sdk_factory.py @@ -4,177 +4,287 @@ """SDK factory functions for creating NeMo Platform SDK instances.""" import logging +from collections.abc import Mapping from dataclasses import dataclass -from typing import Any, Callable, Optional, TypeVar, cast +from pathlib import Path +from typing import Any, Generic, Optional, TypeVar import httpx from nemo_platform import AsyncNeMoPlatform, NeMoPlatform -from nemo_platform_plugin.client.constants import is_workload_identity_token_file_set +from nemo_platform_plugin.client.auth import TokenProvider +from nemo_platform_plugin.client.constants import ( + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + is_workload_identity_token_file_set, + require_workload_identity_without_principal_env, + workload_identity_token_file_from_env, +) from nmp.common.auth import Principal, get_principal_auth_headers, principal_from_env -from nmp.common.config import Configuration, PlatformConfig -from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client +from nmp.common.immutable_http_client import ImmutableDefaultAsyncHttpxClient, ImmutableDefaultHttpxClient from nmp.common.observability import MARK_INTERNAL_REQUEST_HEADERS from nmp.common.observability.otel import get_otel_headers -from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint, resolve_service_endpoint +from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint logger = logging.getLogger(__name__) -PlatformSDKT = TypeVar("PlatformSDKT", NeMoPlatform, AsyncNeMoPlatform) +PlatformSDKT = TypeVar("PlatformSDKT", bound=NeMoPlatform | AsyncNeMoPlatform) +_HTTPClientT = TypeVar("_HTTPClientT", httpx.Client, httpx.AsyncClient) -# Test-only: HTTP clients to use for SDK requests in test context. -# Set by test fixtures to route requests through the in-process test transport. -# -# TODO: Remove these module-level variables once all direct get_platform_sdk() / -# get_async_platform_sdk() callers are migrated to use DependencyProvider. See -# architecture/docs/http-client-injection.md for migration path and best practices. -_test_http_client: Optional[httpx.AsyncClient] = None - -def _base_url_from_config() -> str: - return Configuration.get_platform_config().base_url +@dataclass(frozen=True) +class _SDKConnection(Generic[_HTTPClientT]): + base_url: str + http_client: _HTTPClientT + owns_http_client: bool -def resolve_platform_request_url( - url: str, - *, - platform_config: PlatformConfig, - default_resolver: Callable[[str], httpx.URL], -) -> httpx.URL: - """Resolve the destination URL for an SDK request. - - The generated SDK builds requests from relative paths like - ``/apis/entities/v2/...`` and then calls its private ``_prepare_url`` hook. - Keep the routing policy here, not in that private hook: - - - resolve the SDK URL normally against ``platform.base_url``; - - if the path targets ``/apis/{api_name}/...``, replace only the origin - with ``platform.get_service_url(api_name)``; - - preserve the original path and query string. - """ - request_url = default_resolver(url) - service_pattern = platform_config.create_service_pattern() - if service_pattern is None: - return request_url - match = service_pattern.search(request_url.path) - if match is None: - logger.debug( - "Routing URL to original URL", - extra={"service": "unknown", "path": request_url.path, "host": request_url.host, "port": request_url.port}, +def _sync_sdk_connection( + base_url: str | None, + http_client: httpx.Client | None, +) -> _SDKConnection[httpx.Client]: + if http_client is not None: + return _SDKConnection( + base_url=base_url or resolve_platform_endpoint().connect_base_url, + http_client=http_client, + owns_http_client=False, ) - return request_url - - api_name = match.group(1) - svc_endpoint = resolve_service_endpoint(api_name, platform_config) - if svc_endpoint.transport == "uds": - logger.debug( - "Routing URL to UDS service", - extra={ - "service": api_name, - "path": request_url.path, - "transport": svc_endpoint.transport, - }, + if base_url is not None: + return _SDKConnection(base_url=base_url, http_client=ImmutableDefaultHttpxClient(), owns_http_client=True) + endpoint = resolve_platform_endpoint() + return _SDKConnection( + base_url=endpoint.connect_base_url, + http_client=endpoint.sync_sdk_http_client(), + owns_http_client=True, + ) + + +def _async_sdk_connection( + base_url: str | None, + http_client: httpx.AsyncClient | None, +) -> _SDKConnection[httpx.AsyncClient]: + if http_client is not None: + return _SDKConnection( + base_url=base_url or resolve_platform_endpoint().connect_base_url, + http_client=http_client, + owns_http_client=False, ) - return request_url.copy_with(scheme="http", host="nemo-platform.local", port=None) - service_url = httpx.URL(svc_endpoint.connect_base_url) - routed_url = request_url.copy_with( - scheme=service_url.scheme, - host=service_url.host, - port=service_url.port, + if base_url is not None: + return _SDKConnection(base_url=base_url, http_client=ImmutableDefaultAsyncHttpxClient(), owns_http_client=True) + endpoint = resolve_platform_endpoint() + return _SDKConnection( + base_url=endpoint.connect_base_url, + http_client=endpoint.async_sdk_http_client(), + owns_http_client=True, ) - logger.debug( - "Routing URL to service URL", - extra={ - "service": api_name, - "path": request_url.path, - "host": routed_url.host, - "port": routed_url.port, - }, + + +def _async_sdk_connection_for_endpoint( + endpoint: PlatformEndpoint, + http_client: httpx.AsyncClient | None, +) -> _SDKConnection[httpx.AsyncClient]: + if http_client is not None: + return _SDKConnection( + base_url=endpoint.connect_base_url, + http_client=http_client, + owns_http_client=False, + ) + return _SDKConnection( + base_url=endpoint.connect_base_url, + http_client=endpoint.async_sdk_http_client(), + owns_http_client=True, ) - return routed_url + + +def with_options_reusing_http_client(base_sdk: PlatformSDKT, **kwargs: Any) -> PlatformSDKT: + """Return ``base_sdk.with_options(...)`` while reusing its underlying HTTP client.""" + if kwargs.get("http_client") is None: + kwargs["http_client"] = base_sdk.http_client + return base_sdk.with_options(**kwargs) + + +class _ManagedNeMoPlatform(NeMoPlatform): + def __init__( + self, + *, + token_provider: TokenProvider | None = None, + owns_http_client: bool = False, + **kwargs: Any, + ) -> None: + self._owns_http_client = owns_http_client + super().__init__(token_provider=token_provider, **kwargs) + + def close(self) -> None: + if not self._owns_http_client: + return + self._owns_http_client = False + super().close() + + +class _ManagedAsyncNeMoPlatform(AsyncNeMoPlatform): + def __init__( + self, + *, + token_provider: TokenProvider | None = None, + owns_http_client: bool = False, + **kwargs: Any, + ) -> None: + self._owns_http_client = owns_http_client + super().__init__(token_provider=token_provider, **kwargs) + + async def close(self) -> None: + if not self._owns_http_client: + return + self._owns_http_client = False + await super().close() @dataclass(frozen=True) -class PlatformRequestRouter: - """Routes SDK requests to the platform gateway or a service-specific origin.""" +class _ResolvedSDKInitConfig(Generic[_HTTPClientT]): + base_url: str + workspace: str | None + default_headers: Mapping[str, str] | None + http_client: _HTTPClientT + owns_http_client: bool + token_provider: TokenProvider | None - platform_config: PlatformConfig - default_resolver: Callable[[str], httpx.URL] - def resolve(self, url: str) -> httpx.URL: - return resolve_platform_request_url( - url, - platform_config=self.platform_config, - default_resolver=self.default_resolver, - ) +def _should_bootstrap_workload_identity() -> bool: + return is_workload_identity_token_file_set() -def attach_platform_request_router(sdk: PlatformSDKT) -> PlatformSDKT: - """Attach the platform request router to a generated SDK instance. +def _workload_identity_extra_headers(*, internal: bool) -> dict[str, str]: + return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} - Stainless sends every request through ``_prepare_url``. Assigning the hook is - the SDK integration point; the routing policy itself lives in - :class:`PlatformRequestRouter`. - """ - router = PlatformRequestRouter( - platform_config=Configuration.get_platform_config(), - default_resolver=sdk._prepare_url, - ) - setattr(sdk, "_nmp_request_router", router) - sdk._prepare_url = router.resolve - return sdk +def _non_auth_otel_headers() -> dict[str, str]: + headers: dict[str, str] = {} + for name, value in get_otel_headers().items(): + normalized_name = name.lower() + if normalized_name == "x-nmp-internal" or normalized_name.startswith("x-nmp-principal-"): + continue + headers[name] = value + return headers -def with_options_preserving_request_router(base_sdk: PlatformSDKT, **kwargs: Any) -> PlatformSDKT: - """Return ``base_sdk.with_options(...)`` while preserving platform request routing.""" - scoped_sdk = cast(PlatformSDKT, base_sdk.with_options(**kwargs)) - router = getattr(base_sdk, "_nmp_request_router", None) - if isinstance(router, PlatformRequestRouter): - setattr(scoped_sdk, "_nmp_request_router", router) - scoped_sdk._prepare_url = router.resolve - return scoped_sdk +def _workload_identity_token_file() -> Path: + token_file = workload_identity_token_file_from_env() + if token_file is None: + raise RuntimeError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} is not set") + return token_file -def _sync_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.Client | None, -) -> httpx.Client: - if http_client is not None: - return http_client - if endpoint.transport == "uds": - return endpoint.sync_http_client() - return shared_sync_http_client() +def _workload_identity_token_provider(base_url: str) -> TokenProvider: + from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider -def _async_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.AsyncClient | None, -) -> httpx.AsyncClient: - if http_client is not None: - return http_client - if _test_http_client is not None: - return _test_http_client - if endpoint.transport == "uds": - return endpoint.async_http_client() - return shared_async_http_client() + return resolve_workload_exchange_provider( + base_url=base_url, + subject_token_file=_workload_identity_token_file(), + ) -def _should_bootstrap_workload_identity( +def _ensure_no_trusted_headers_for_workload_identity( *, as_service: str | None, on_behalf_of: str | Principal | None, - http_client: httpx.Client | httpx.AsyncClient | None, - endpoint: PlatformEndpoint, -) -> bool: - return ( - as_service is None - and on_behalf_of is None - and http_client is None - and endpoint.transport != "uds" - and is_workload_identity_token_file_set() +) -> None: + if as_service is not None or on_behalf_of is not None: + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} cannot be combined with trusted principal headers") + + +def _resolve_sdk_init_config( + *, + connection: _SDKConnection[_HTTPClientT], + as_service: str | None, + internal: bool, + on_behalf_of: str | Principal | None, +) -> _ResolvedSDKInitConfig[_HTTPClientT]: + if not _should_bootstrap_workload_identity(): + headers = _get_default_headers(as_service, internal, on_behalf_of) + return _ResolvedSDKInitConfig( + base_url=connection.base_url, + workspace=None, + default_headers=headers if headers else None, + http_client=connection.http_client, + owns_http_client=connection.owns_http_client, + token_provider=None, + ) + + require_workload_identity_without_principal_env() + _ensure_no_trusted_headers_for_workload_identity(as_service=as_service, on_behalf_of=on_behalf_of) + extra_headers = _workload_identity_extra_headers(internal=internal) + return _ResolvedSDKInitConfig( + base_url=connection.base_url, + workspace=None, + default_headers=extra_headers or None, + http_client=connection.http_client, + owns_http_client=connection.owns_http_client, + token_provider=_workload_identity_token_provider(connection.base_url), ) -def _workload_identity_extra_headers(*, internal: bool) -> dict[str, str]: - return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} +def _sdk_constructor_kwargs( + *, + max_retries: int | None, + strict_response_validation: bool, +) -> dict[str, Any]: + kwargs: dict[str, Any] = {} + if max_retries is not None: + kwargs["max_retries"] = max_retries + if strict_response_validation: + kwargs["_strict_response_validation"] = True + return kwargs + + +def _sync_sdk_from_init_config( + sdk_config: _ResolvedSDKInitConfig[httpx.Client], + *, + max_retries: int | None = None, + strict_response_validation: bool = False, +) -> NeMoPlatform: + return _ManagedNeMoPlatform( + token_provider=sdk_config.token_provider, + owns_http_client=sdk_config.owns_http_client, + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + **_sdk_constructor_kwargs( + max_retries=max_retries, + strict_response_validation=strict_response_validation, + ), + ) + + +def _async_sdk_from_init_config( + sdk_config: _ResolvedSDKInitConfig[httpx.AsyncClient], + *, + max_retries: int | None = None, + strict_response_validation: bool = False, +) -> AsyncNeMoPlatform: + return _ManagedAsyncNeMoPlatform( + token_provider=sdk_config.token_provider, + owns_http_client=sdk_config.owns_http_client, + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + **_sdk_constructor_kwargs( + max_retries=max_retries, + strict_response_validation=strict_response_validation, + ), + ) + + +def _task_on_behalf_of_principal() -> Principal | None: + principal = principal_from_env() + return principal.effective_principal if principal is not None else None + + +def _warn_missing_task_principal(*, as_service: str, async_sdk: bool) -> None: + qualifier = "async task SDK" if async_sdk else "task SDK" + logger.warning( + "NMP_PRINCIPAL not set; %s will authenticate as service:%s without on-behalf-of delegation", + qualifier, + as_service, + ) def _get_default_headers( @@ -244,9 +354,11 @@ def get_platform_sdk( http_client: httpx.Client | None = None, on_behalf_of: str | Principal | None = None, base_url: str | None = None, + max_retries: int | None = None, + _strict_response_validation: bool = False, ) -> NeMoPlatform: """ - Returns an instance of the NeMoPlatform SDK configured with the platform's base URL. + Returns a NeMoPlatform SDK configured from explicit arguments or the resolved platform endpoint. Args: as_service: If provided, use service principal headers (service:{name}). @@ -257,32 +369,24 @@ def get_platform_sdk( Use this for controllers and background tasks that make internal API calls. http_client: Optional sync HTTP client to use for requests. on_behalf_of: Optional principal ID to use for on-behalf-of authorization. - base_url: Optional platform base URL. Defaults to configured platform base URL. + base_url: Optional platform base URL. When omitted with no explicit http_client, + the resolved platform endpoint supplies both the base URL and HTTP client. Returns: Configured NeMoPlatform SDK instance. """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( + connection = _sync_sdk_connection(base_url, http_client) + sdk_config = _resolve_sdk_init_config( + connection=connection, as_service=as_service, + internal=internal, on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - headers = _workload_identity_extra_headers(internal=internal) - sdk = NeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - default_headers=headers if headers else None, - ) - return attach_platform_request_router(sdk) - - headers = _get_default_headers(as_service, internal, on_behalf_of) - sdk = NeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=_sync_http_client_for_endpoint(endpoint, http_client), - default_headers=headers if headers else None, ) - return attach_platform_request_router(sdk) + return _sync_sdk_from_init_config( + sdk_config, + max_retries=max_retries, + strict_response_validation=_strict_response_validation, + ) def get_task_sdk(as_service: str, http_client: httpx.Client | None = None) -> NeMoPlatform: @@ -299,23 +403,21 @@ def get_task_sdk(as_service: str, http_client: httpx.Client | None = None) -> Ne Returns: Configured NeMoPlatform SDK with internal + on-behalf-of headers. """ - if http_client is None and is_workload_identity_token_file_set(): - return get_platform_sdk(internal=True) - - if http_client is None: - http_client = resolve_platform_endpoint().sync_sdk_http_client() - - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; task SDK will authenticate as service:%s without on-behalf-of delegation", - as_service, + if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() + return get_platform_sdk( + internal=True, + http_client=http_client, ) + + on_behalf_of = _task_on_behalf_of_principal() + if on_behalf_of is None: + _warn_missing_task_principal(as_service=as_service, async_sdk=False) return get_platform_sdk( as_service=as_service, internal=True, http_client=http_client, - on_behalf_of=principal.effective_principal if principal else None, + on_behalf_of=on_behalf_of, ) @@ -333,23 +435,21 @@ def get_async_task_sdk(as_service: str, http_client: Optional[httpx.AsyncClient] Returns: Configured AsyncNeMoPlatform SDK with internal + on-behalf-of headers. """ - if http_client is None and is_workload_identity_token_file_set(): - return get_async_platform_sdk(internal=True) - - if http_client is None: - http_client = resolve_platform_endpoint().async_sdk_http_client() - - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; async task SDK will authenticate as service:%s without on-behalf-of delegation", - as_service, + if is_workload_identity_token_file_set(): + require_workload_identity_without_principal_env() + return get_async_platform_sdk( + internal=True, + http_client=http_client, ) + + on_behalf_of = _task_on_behalf_of_principal() + if on_behalf_of is None: + _warn_missing_task_principal(as_service=as_service, async_sdk=True) return get_async_platform_sdk( as_service=as_service, internal=True, http_client=http_client, - on_behalf_of=principal.effective_principal if principal else None, + on_behalf_of=on_behalf_of, ) @@ -359,9 +459,11 @@ def get_async_platform_sdk( http_client: Optional[httpx.AsyncClient] = None, on_behalf_of: Optional[str | Principal] = None, base_url: str | None = None, + max_retries: int | None = None, + _strict_response_validation: bool = False, ) -> AsyncNeMoPlatform: """ - Returns an instance of the AsyncNeMoPlatform SDK configured with the platform's base URL. + Returns an AsyncNeMoPlatform SDK configured from explicit arguments or the resolved platform endpoint. Args: as_service: If provided, use service principal headers (service:{name}). @@ -370,46 +472,81 @@ def get_async_platform_sdk( If None and auth is enabled, propagates the current user's auth context. internal: If True, mark all requests from this SDK as internal requests. Use this for controllers and background tasks that make internal API calls. - http_client: Optional HTTP client to use for requests. Used for test injection - via DependencyProvider. See architecture/docs/http-client-injection.md. + http_client: Optional async HTTP client to use for requests. on_behalf_of: Optional principal ID to use for on-behalf-of authorization. - base_url: Optional platform base URL. Defaults to configured platform base URL. + base_url: Optional platform base URL. When omitted with no explicit http_client, + the resolved platform endpoint supplies both the base URL and HTTP client. Returns: Configured AsyncNeMoPlatform SDK instance. """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( + connection = _async_sdk_connection(base_url, http_client) + sdk_config = _resolve_sdk_init_config( + connection=connection, as_service=as_service, + internal=internal, on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - headers = _workload_identity_extra_headers(internal=internal) - if _test_http_client is None: - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - default_headers=headers if headers else None, - ) - else: - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=_async_http_client_for_endpoint(endpoint, http_client), - default_headers=headers if headers else None, - ) - return attach_platform_request_router(sdk) - - headers = _get_default_headers(as_service, internal, on_behalf_of) - - # Use explicitly provided http_client (from DependencyProvider) or fall back to - # module-level _test_http_client for backward compatibility with direct callers. - effective_client = _async_http_client_for_endpoint(endpoint, http_client) - - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=effective_client, - default_headers=headers if headers else None, ) - return attach_platform_request_router(sdk) + return _async_sdk_from_init_config( + sdk_config, + max_retries=max_retries, + strict_response_validation=_strict_response_validation, + ) + + +def _get_async_platform_sdk_for_endpoint( + endpoint: PlatformEndpoint, + *, + as_service: str | None = None, + internal: bool = False, + http_client: httpx.AsyncClient | None = None, + on_behalf_of: str | Principal | None = None, + max_retries: int | None = None, + _strict_response_validation: bool = False, +) -> AsyncNeMoPlatform: + """Create an async SDK bound to an already resolved platform endpoint. + + When ``http_client`` is omitted, the endpoint supplies a routed HTTP client + owned by the returned SDK. When ``http_client`` is provided, the returned SDK + borrows it and ``close()`` leaves that client open. + """ + connection = _async_sdk_connection_for_endpoint(endpoint, http_client) + sdk_config = _resolve_sdk_init_config( + connection=connection, + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + ) + return _async_sdk_from_init_config( + sdk_config, + max_retries=max_retries, + strict_response_validation=_strict_response_validation, + ) + + +def get_service_scoped_sdk( + base_sdk: PlatformSDKT, + service_name: str, + *, + internal: bool = True, + on_behalf_of: str | Principal | None = None, +) -> PlatformSDKT: + """Derive a service-principal SDK from an existing SDK. + + The returned SDK reuses the base SDK's HTTP client and endpoint routing. + If the base SDK is using workload identity, keep bearer-token auth and add + only workload-safe internal headers instead of trusted principal headers. + """ + if base_sdk.token_provider is not None: + if on_behalf_of is not None: + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} cannot be combined with trusted principal headers") + headers = _workload_identity_extra_headers(internal=internal) + else: + headers = _get_default_headers( + as_service=service_name, + internal=internal, + on_behalf_of=on_behalf_of, + ) + return with_options_reusing_http_client(base_sdk, set_default_headers=headers) def get_request_scoped_sdk( @@ -417,7 +554,7 @@ def get_request_scoped_sdk( ) -> AsyncNeMoPlatform: """Create a request-scoped SDK with current auth and observability headers. - Takes a base SDK (with shared HTTP client) and returns a new SDK instance + Takes a base SDK and returns a new SDK instance that reuses the same HTTP client with the current request's auth headers applied via .with_options(). This is lightweight - the underlying HTTP client is reused. @@ -433,14 +570,15 @@ def get_request_scoped_sdk( for FastAPI dependency injection. """ - # Combine OTEL headers (tracing) + auth headers (user identity) - headers = get_otel_headers().copy() - headers.update(get_principal_auth_headers()) + # Combine OTEL headers (tracing) + auth headers (user identity). + headers = _non_auth_otel_headers() + if base_sdk.token_provider is None: + headers.update(get_principal_auth_headers()) # If we have headers to add, create a new SDK with them # This reuses the underlying HTTP client (lightweight operation) if headers: - return with_options_preserving_request_router(base_sdk, set_default_headers=headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=headers) return base_sdk @@ -481,6 +619,9 @@ def get_sdk_on_behalf_of( secret = delegated_sdk.secrets.access("my-secret", workspace="workspace-name") ``` """ + if base_sdk.token_provider is not None: + raise ValueError(f"{WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR} cannot be combined with trusted principal headers") + # Merge existing headers with the new on-behalf-of header headers = base_sdk.default_headers or {} if isinstance(on_behalf_of, Principal): @@ -496,7 +637,7 @@ def get_sdk_on_behalf_of( merged_headers = {**headers, "X-NMP-Principal-On-Behalf-Of": on_behalf_of} merged_headers.pop("X-NMP-Principal-On-Behalf-Of-Groups", None) merged_headers.pop("X-NMP-Principal-On-Behalf-Of-Email", None) - return with_options_preserving_request_router(base_sdk, set_default_headers=merged_headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=merged_headers) def get_entity_parts(name: str, default_workspace: str | None = None) -> tuple[str, str]: @@ -518,34 +659,29 @@ def get_entity_parts(name: str, default_workspace: str | None = None) -> tuple[s class PlatformSDKProvider: """Rich :class:`~nemo_platform_plugin.sdk_provider.SDKProvider` that uses - platform internals (shared HTTP clients, URL routing, OTEL headers, auth - context vars). + platform internals (SDK-owned HTTP clients, OTEL headers, auth context vars). Registered as a ``nemo.sdk_provider`` entry-point so it is discovered automatically when ``nmp-common`` is installed. """ - def get_task_sdk(self, service_name: str, http_client: httpx.Client | None = None) -> NeMoPlatform: - return get_task_sdk(service_name, http_client=http_client) + def get_task_sdk(self, service_name: str) -> NeMoPlatform: + return get_task_sdk(service_name) - def get_async_task_sdk(self, service_name: str, http_client: httpx.AsyncClient | None = None) -> AsyncNeMoPlatform: - return get_async_task_sdk(service_name, http_client=http_client) + def get_async_task_sdk(self, service_name: str) -> AsyncNeMoPlatform: + return get_async_task_sdk(service_name) def get_platform_sdk( self, *, as_service: str | None = None, internal: bool = False, - http_client: httpx.Client | None = None, - on_behalf_of: str | Principal | None = None, - base_url: str | None = None, + on_behalf_of: str | None = None, ) -> NeMoPlatform: return get_platform_sdk( as_service=as_service, internal=internal, - http_client=http_client, on_behalf_of=on_behalf_of, - base_url=base_url, ) def get_async_platform_sdk( @@ -553,12 +689,10 @@ def get_async_platform_sdk( *, as_service: str | None = None, internal: bool = False, - on_behalf_of: str | Principal | None = None, - base_url: str | None = None, + on_behalf_of: str | None = None, ) -> AsyncNeMoPlatform: return get_async_platform_sdk( as_service=as_service, internal=internal, on_behalf_of=on_behalf_of, - base_url=base_url, ) diff --git a/packages/nmp_common/src/nmp/common/service/__init__.py b/packages/nmp_common/src/nmp/common/service/__init__.py index 1d8874e47a..084a7a65e8 100644 --- a/packages/nmp_common/src/nmp/common/service/__init__.py +++ b/packages/nmp_common/src/nmp/common/service/__init__.py @@ -6,7 +6,6 @@ from nmp.common.service.base import DependencyProvider, RouterConfig, Service from nmp.common.service.dependencies import ( get_entity_client, - get_nemo_client, get_platform_config, get_sdk_client, get_service_config, @@ -21,7 +20,6 @@ "RouterConfig", "build_downstream_service_headers", "get_entity_client", - "get_nemo_client", "get_platform_config", "get_sdk_client", "get_service_config", diff --git a/packages/nmp_common/src/nmp/common/service/base.py b/packages/nmp_common/src/nmp/common/service/base.py index f314512fa5..cf57ac6ba3 100644 --- a/packages/nmp_common/src/nmp/common/service/base.py +++ b/packages/nmp_common/src/nmp/common/service/base.py @@ -10,23 +10,44 @@ from abc import ABC, abstractmethod from contextlib import asynccontextmanager from dataclasses import dataclass -from threading import RLock -from typing import ClassVar, Dict, Generic, List, Optional, Self, Type, TypeVar, cast, get_args, get_origin +from typing import Any, ClassVar, Dict, Generic, List, Optional, Self, Type, TypeVar, get_args, get_origin import httpx from fastapi import APIRouter, FastAPI, Request from fastapi.openapi.utils import get_openapi +from fastapi.routing import APIRoute from nemo_platform import AsyncNeMoPlatform -from nemo_platform_plugin.client.client import AsyncNemoClient from nmp.common.api.utils import register_query_param_schemas from nmp.common.config import Configuration, PlatformConfig, ServiceConfig +from nmp.common.config import get_platform_config as load_platform_config from nmp.common.controller import Controller from nmp.common.entities.client import EntityClient -from nmp.common.platform_endpoint import resolve_platform_endpoint, resolve_service_endpoint +from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint, resolve_service_endpoint logger = logging.getLogger(__name__) +class _ServiceFastAPI(FastAPI): + """FastAPI app with NeMo Platform OpenAPI post-processing.""" + + nmp_openapi_summary: str | None = None + nmp_openapi_description: str | None = None + + def openapi(self) -> dict[str, Any]: + if self.openapi_schema: + return self.openapi_schema + openapi_schema = get_openapi( + title=self.title, + version=self.version, + summary=self.nmp_openapi_summary or self.summary, + description=self.description if self.nmp_openapi_description is None else self.nmp_openapi_description, + routes=self.routes, + tags=self.openapi_tags, + ) + self.openapi_schema = register_query_param_schemas(openapi_schema) + return self.openapi_schema + + @dataclass class RouterConfig: """Configuration for a router including its OpenAPI tag metadata.""" @@ -40,7 +61,7 @@ class RouterConfig: TConfig = TypeVar("TConfig", bound=ServiceConfig) -def _get_config_class_from_generic(cls: type) -> Type[ServiceConfig] | None: +def _get_config_class_from_generic(cls: type) -> Type[TConfig] | None: """Extract the config class from Service[TConfig] generic parameter. Args: @@ -60,40 +81,77 @@ def _get_config_class_from_generic(cls: type) -> Type[ServiceConfig] | None: class DependencyProvider: """ - Manages SDK, NemoClient, entity client, HTTP client, and config lifecycle for NeMo Platform services. + Manages SDK, entity client, HTTP client, and config lifecycle for NeMo Platform services. - Provides lazy initialization, FastAPI dependency wiring, and cleanup. + Provides explicit lifecycle initialization, FastAPI dependency wiring, and cleanup. - The `_http_client` field supports test injection - when set, it's passed to - `get_async_platform_sdk()` to route requests through ASGI transport in tests. - See architecture/docs/http-client-injection.md for details. + The PlatformEndpoint owns service routing for the current PlatformConfig. + Tests and embedded platform assembly can still inject an explicit endpoint + or HTTP client before initialization. """ - def __init__(self) -> None: - self._client_lock = RLock() - self._http_client: Optional[httpx.AsyncClient] = None + def __init__(self, service_name: str = "platform") -> None: + self._configured_http_client: Optional[httpx.AsyncClient] = None + self._platform_endpoint: Optional[PlatformEndpoint] = None + self._service_http_client: Optional[httpx.AsyncClient] = None self._sdk_client: Optional[AsyncNeMoPlatform] = None self._platform_config: Optional[PlatformConfig] = None - self._service_name: str = "platform" + self._service_name = service_name + + def configure_service_name(self, service_name: str) -> None: + """Set the service name used for downstream service-principal headers.""" + self._service_name = service_name + + def initialize(self) -> None: + """Create service-owned clients before request handling starts.""" + if self._service_http_client is not None and self._sdk_client is not None: + return + + from nmp.common.sdk_factory import _get_async_platform_sdk_for_endpoint + + endpoint = self.get_platform_endpoint() + self._sdk_client = _get_async_platform_sdk_for_endpoint( + endpoint, + http_client=self._configured_http_client, + ) + self._service_http_client = self._sdk_client.http_client + + def _require_initialized(self) -> None: + if self._service_http_client is None or self._sdk_client is None: + raise RuntimeError("DependencyProvider is not initialized. Call initialize() during service startup.") def get_http_client(self) -> httpx.AsyncClient: - """Return the httpx.AsyncClient for this provider, creating it lazily. - - The client is transport-aware: for a ``unix://`` platform endpoint it is - bound to the Unix domain socket, otherwise it is the SDK's default TCP - client. Because this cached client is injected into the SDK and - NemoClient factories (which skip their own transport selection when a - client is supplied), building it endpoint-aware here is what makes - service-to-service requests work over UDS. - - Each DependencyProvider manages its own HTTP client by default. - If you need to share a client across providers (e.g., for connection - pooling), you can inject the same client via _http_client. + """Return the initialized httpx.AsyncClient for this provider. + + The PlatformEndpoint carries resolved routing state. ``initialize()`` + creates the concrete client before requests are handled. + """ + self._require_initialized() + http_client = self._service_http_client + assert http_client is not None + return http_client + + def get_platform_endpoint(self) -> PlatformEndpoint: + """Return the resolved PlatformEndpoint for this provider.""" + if self._platform_endpoint is None: + self._platform_endpoint = resolve_platform_endpoint(self.get_platform_config()) + return self._platform_endpoint + + def configure_platform_endpoint(self, platform_endpoint: PlatformEndpoint) -> None: + """Use an externally resolved PlatformEndpoint before SDK/client creation.""" + if self._service_http_client is not None or self._sdk_client is not None: + raise RuntimeError("Cannot configure DependencyProvider PlatformEndpoint after initialization") + self._platform_endpoint = platform_endpoint + + def configure_http_client(self, http_client: httpx.AsyncClient) -> None: + """Use an externally managed HTTP client before SDK creation. + + This is used by platform assembly and tests to route all service-owned + SDK calls through the same transport and connection pool. """ - with self._client_lock: - if self._http_client is None: - self._http_client = resolve_platform_endpoint().async_sdk_http_client() - return self._http_client + if self._service_http_client is not None or self._sdk_client is not None: + raise RuntimeError("Cannot configure DependencyProvider HTTP client after initialization") + self._configured_http_client = http_client def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: """Return the async platform SDK client. @@ -108,18 +166,19 @@ def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: Returns: SDK client - cached instance if as_service is None, new instance otherwise. """ - from nmp.common.sdk_factory import get_async_platform_sdk + self._require_initialized() # When as_service is specified, return a fresh SDK with service credentials. # This is needed for startup/background code where no user auth context exists. if as_service is not None: - return get_async_platform_sdk(as_service=as_service, internal=True, http_client=self.get_http_client()) + from nmp.common.sdk_factory import get_service_scoped_sdk + + return get_service_scoped_sdk(self.get_sdk_client(), as_service) # For request handling, use cached SDK. EntityClient adds auth headers per-request. - with self._client_lock: - if self._sdk_client is None: - self._sdk_client = get_async_platform_sdk(http_client=self.get_http_client()) - return self._sdk_client + sdk_client = self._sdk_client + assert sdk_client is not None + return sdk_client def get_entity_client(self, as_service: str | None = None) -> Optional[EntityClient]: """Return the EntityClient. @@ -157,19 +216,20 @@ def _get_entity_sdk_on_behalf_of(self) -> AsyncNeMoPlatform: Uses the cached base SDK and applies per-request headers via .with_options() (lightweight — reuses the HTTP connection pool). """ - from nmp.common.sdk_factory import with_options_preserving_request_router + from nmp.common.sdk_factory import with_options_reusing_http_client from nmp.common.service.headers import build_downstream_service_headers base_sdk = self.get_sdk_client() headers = build_downstream_service_headers(self._service_name) - return with_options_preserving_request_router(base_sdk, set_default_headers=headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=headers) def get_platform_config(self) -> PlatformConfig: """Return the PlatformConfig (lazily initialized).""" if self._platform_config is None: - self._platform_config = Configuration.get_platform_config() - return self._platform_config + self._platform_config = load_platform_config() + platform_config = self._platform_config + return platform_config def get_request_scoped_sdk(self) -> AsyncNeMoPlatform: """Return a request-scoped SDK with current auth and OTEL headers. @@ -182,12 +242,6 @@ def get_request_scoped_sdk(self) -> AsyncNeMoPlatform: base_sdk = self.get_sdk_client() # Cached base SDK return get_request_scoped_sdk(base_sdk) - def get_request_scoped_nemo_client(self) -> AsyncNemoClient: - """Return a fresh async NemoClient with request-scoped headers.""" - from nmp.common.client_factory import get_async_nemo_client - - return get_async_nemo_client(http_client=self.get_http_client()) - def get_effective_principal_id(self, request: Request) -> str: """Return the effective principal ID from the current request auth context.""" from nmp.common.auth import get_auth_client @@ -199,14 +253,12 @@ def setup_dependencies(self, app: FastAPI, service: "Service") -> None: from nmp.common.service.dependencies import ( get_effective_principal_id, get_entity_client, - get_nemo_client, get_platform_config, get_sdk_client, get_service_config, ) app.dependency_overrides[get_sdk_client] = self.get_request_scoped_sdk - app.dependency_overrides[get_nemo_client] = self.get_request_scoped_nemo_client app.dependency_overrides[get_entity_client] = self.get_entity_client app.dependency_overrides[get_effective_principal_id] = self.get_effective_principal_id app.dependency_overrides[get_platform_config] = self.get_platform_config @@ -214,14 +266,13 @@ def setup_dependencies(self, app: FastAPI, service: "Service") -> None: app.dependency_overrides[get_service_config] = lambda: service._service_config async def close(self) -> None: - """Close the provider-owned HTTP transport and clear cached wrappers.""" - with self._client_lock: - http_client = self._http_client - self._http_client = None - self._sdk_client = None - - if http_client is not None: - await http_client.aclose() + """Close the cached SDK and detach configured clients.""" + sdk_client = self._sdk_client + self._service_http_client = None + self._configured_http_client = None + self._sdk_client = None + if sdk_client is not None: + await sdk_client.close() class Service(ABC, Generic[TConfig]): @@ -280,7 +331,10 @@ def __init__( self.module_name = module_name self._app: Optional[FastAPI] = None self._startup_background_tasks: list[asyncio.Task] = [] - self._dependency_provider = dependency_provider if dependency_provider is not None else DependencyProvider() + self._dependency_provider = ( + dependency_provider if dependency_provider is not None else DependencyProvider(service_name=name) + ) + self._dependency_provider.configure_service_name(name) if dependencies is not None: self._dependencies = list(dependencies) else: @@ -288,9 +342,7 @@ def __init__( # Extract config class from generic type parameter and load config config_class = _get_config_class_from_generic(type(self)) - self._service_config = ( - cast(TConfig | None, Configuration.get_service_config(config_class)) if config_class else None - ) + self._service_config = Configuration.get_service_config(config_class) if config_class else None @property def dependency_provider(self) -> DependencyProvider: @@ -445,6 +497,7 @@ def create_app(self) -> FastAPI: async def lifespan(app: FastAPI): """Lifespan context manager for the FastAPI app.""" logger.info("Starting service...", extra={"service": self.name}) + self._dependency_provider.initialize() # Run service-specific startup initialization (e.g., database setup) await self.on_startup() @@ -470,13 +523,14 @@ async def lifespan(app: FastAPI): router_configs = self.get_routers() openapi_tags: List[Dict[str, str]] = [{"name": rc.tag, "description": rc.description} for rc in router_configs] - app = FastAPI( + app = _ServiceFastAPI( title=self.title, description=self.description, version=self.version, openapi_tags=openapi_tags, lifespan=lifespan, ) + app.nmp_openapi_summary = f"This is the OpenAPI Schema for the {self.title}." # Store reference to app for use in on_startup self._app = app @@ -497,35 +551,12 @@ async def lifespan(app: FastAPI): # Include service-specific routers, tagging any routes that have no tags yet for rc in router_configs: for route in rc.router.routes: - if hasattr(route, "tags") and not route.tags: + if isinstance(route, APIRoute) and not route.tags: route.tags = [rc.tag] app.include_router(rc.router, prefix=rc.prefix) - # Setup custom OpenAPI schema - self._setup_custom_openapi(app, openapi_tags) - return app - def _setup_custom_openapi(self, app: FastAPI, openapi_tags: List[Dict[str, str]]) -> None: - """Configure custom OpenAPI schema generation.""" - - def custom_openapi(): - if app.openapi_schema: - return app.openapi_schema - openapi_schema = get_openapi( - title=self.title, - version=self.version, - summary=f"This is the OpenAPI Schema for the {self.title}.", - description="", - routes=app.routes, - tags=openapi_tags, - ) - openapi_schema = register_query_param_schemas(openapi_schema) - app.openapi_schema = openapi_schema - return app.openapi_schema - - app.openapi = custom_openapi # type: ignore[method-assign] - # ========================================================================= # Startup and readiness # ========================================================================= diff --git a/packages/nmp_common/src/nmp/common/service/dependencies.py b/packages/nmp_common/src/nmp/common/service/dependencies.py index 2cf329cf03..dd0fe4cc76 100644 --- a/packages/nmp_common/src/nmp/common/service/dependencies.py +++ b/packages/nmp_common/src/nmp/common/service/dependencies.py @@ -9,12 +9,12 @@ from __future__ import annotations -from typing import Callable, TypeVar +from collections.abc import Callable, Mapping +from typing import TypeVar from fastapi import Request from nemo_platform_plugin.dependencies import get_effective_principal_id as get_effective_principal_id from nemo_platform_plugin.dependencies import get_entity_client as get_entity_client -from nemo_platform_plugin.dependencies import get_nemo_client as get_nemo_client from nemo_platform_plugin.dependencies import get_platform_config as get_platform_config from nemo_platform_plugin.dependencies import get_sdk_client as get_sdk_client from nemo_platform_plugin.dependencies import get_service_config as get_service_config @@ -37,12 +37,12 @@ def get_service_config_factory(config_class: type[T]) -> Callable[[Request], T]: """ def _get_config(request: Request) -> T: - registry: dict[type[ServiceConfig], ServiceConfig] = getattr(request.app.state, "service_configs", {}) + registry: Mapping[type[T], T] = getattr(request.app.state, "service_configs", {}) if config_class not in registry: raise RuntimeError( f"Service config {config_class.__name__} not registered. " "Ensure the service is loaded and its config is added to app.state.service_configs." ) - return registry[config_class] # type: ignore[return-value] + return registry[config_class] return _get_config diff --git a/packages/nmp_common/tests/api/test_query_param_schemas.py b/packages/nmp_common/tests/api/test_query_param_schemas.py index 06cbf89c14..25e3b2c3fe 100644 --- a/packages/nmp_common/tests/api/test_query_param_schemas.py +++ b/packages/nmp_common/tests/api/test_query_param_schemas.py @@ -4,17 +4,16 @@ """Tests for register_query_param_schemas / clear_query_param_schemas. These schemas are attached to FastAPI endpoints via ``openapi_extra`` and are -not reachable through Pydantic's response-model walk. The runtime -``custom_openapi`` hook has to call ``register_query_param_schemas`` explicitly -or the live ``/openapi.json`` will contain dangling ``$ref``s to the filter -classes. +not reachable through Pydantic's response-model walk. The runtime OpenAPI path +has to call ``register_query_param_schemas`` or the live ``/openapi.json`` will +contain dangling ``$ref``s to the filter classes. """ from typing import Optional import nemo_platform_plugin.jobs.openapi_utils as job_openapi_utils import pytest -from fastapi import FastAPI, Query, Request +from fastapi import APIRouter, FastAPI, Query, Request from fastapi.testclient import TestClient from nmp.common.api.utils import ( clear_query_param_schemas, @@ -22,6 +21,7 @@ install_query_param_schema_openapi_hook, register_query_param_schemas, ) +from nmp.common.service import RouterConfig, Service from pydantic import BaseModel, create_model @@ -106,23 +106,26 @@ def test_clear_resets_registry_between_services(): assert "_DummyFilter" not in spec["components"]["schemas"] -def test_custom_openapi_hook_resolves_filter_ref(): - """End-to-end: a FastAPI app that wires ``register_query_param_schemas`` - into its ``custom_openapi`` hook emits a spec where the filter $ref - resolves — which is exactly the regression the runtime was missing. - """ - app = FastAPI() +def test_service_openapi_resolves_filter_ref(): + """End-to-end: service OpenAPI generation resolves generated filter refs.""" - @app.get( - "/items", - openapi_extra=generate_openapi_extra_params(filter_schema=_DummyFilter), - ) - async def list_items(request: Request, page: int = Query(default=1)): - return {"data": []} + class _QueryParamService(Service): + def __init__(self): + super().__init__(name="query-param-test", module_name="nmp.test") - install_query_param_schema_openapi_hook(app) + def get_routers(self) -> list[RouterConfig]: + router = APIRouter() + + @router.get( + "/items", + openapi_extra=generate_openapi_extra_params(filter_schema=_DummyFilter), + ) + async def list_items(request: Request, page: int = Query(default=1)): + return {"data": []} + + return [RouterConfig(router, tag="Items", description="Item endpoints")] - spec = TestClient(app).get("/openapi.json").json() + spec = TestClient(_QueryParamService().create_app()).get("/openapi.json").json() assert "_DummyFilter" in spec["components"]["schemas"] param = next(p for p in spec["paths"]["/items"]["get"]["parameters"] if p["name"] == "filter") diff --git a/packages/nmp_common/tests/auth/test_access_key_lifecycle.py b/packages/nmp_common/tests/auth/test_access_key_lifecycle.py index 0bf496c0da..ce5f265682 100644 --- a/packages/nmp_common/tests/auth/test_access_key_lifecycle.py +++ b/packages/nmp_common/tests/auth/test_access_key_lifecycle.py @@ -3,7 +3,7 @@ """Unit tests for auth-service Scoped Access Key lifecycle validation.""" -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -20,7 +20,7 @@ def _config() -> AuthConfig: @pytest.mark.asyncio -async def test_authenticator_uses_sdk_routing_and_returns_trusted_claims() -> None: +async def test_authenticator_uses_injected_client_and_returns_trusted_claims() -> None: requests: list[httpx.Request] = [] def handler(request: httpx.Request) -> httpx.Response: @@ -51,7 +51,39 @@ def handler(request: httpx.Request) -> httpx.Response: assert result.claims.subject == "alice@example.com" assert result.claims.groups == ["team-ml"] assert result.claims.scopes == ["models:read"] - assert requests[0].url == httpx.URL("http://auth.internal:8080/apis/auth/authenticate") + assert requests[0].url == httpx.URL("http://platform.example.com/apis/auth/authenticate") + + +@pytest.mark.asyncio +async def test_aclose_closes_owned_sdk() -> None: + sdk = MagicMock() + sdk.close = AsyncMock() + authenticator = AccessKeyLifecycleAuthenticator(_config()) + authenticator._sdk = sdk + + await authenticator.aclose() + + sdk.close.assert_awaited_once_with() + assert authenticator._sdk is None + + +@pytest.mark.asyncio +async def test_aclose_does_not_close_sdk_with_injected_http_client() -> None: + async with httpx.AsyncClient() as http_client: + authenticator = AccessKeyLifecycleAuthenticator(_config(), http_client=http_client) + with patch.object( + Configuration, + "get_platform_config", + return_value=PlatformConfig(base_url="http://platform.example.com", services=""), + ): + sdk = authenticator._get_sdk() + + assert sdk._client is http_client + + await authenticator.aclose() + + assert not http_client.is_closed + assert authenticator._sdk is None @pytest.mark.asyncio @@ -78,6 +110,39 @@ def handler(request: httpx.Request) -> httpx.Response: assert await authenticator.authenticate("candidate-token") is None +@pytest.mark.asyncio +async def test_authenticator_can_accept_workload_token_response() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={ + "principal": "system:serviceaccount:nemo:job", + "groups": ["team-ml"], + "scopes": ["openid", "email"], + "token_kind": "workload_access_token", + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http_client: + authenticator = AccessKeyLifecycleAuthenticator(_config(), http_client=http_client) + with patch.object( + Configuration, + "get_platform_config", + return_value=PlatformConfig(base_url="http://platform.example.com", services=""), + ): + result = await authenticator.authenticate_token( + "workload-access-token", + token_kinds=("workload_access_token",), + ) + + assert isinstance(result, ResolvedBearerToken) + assert result.token_kind == "workload_access_token" + assert result.claims.subject == "system:serviceaccount:nemo:job" + assert result.claims.groups == ["team-ml"] + assert result.claims.raw_claims["nmp_token_type"] == "workload_access_token" + assert "jti" not in result.claims.raw_claims + + @pytest.mark.asyncio async def test_authenticator_rejects_malformed_sdk_response() -> None: def handler(request: httpx.Request) -> httpx.Response: diff --git a/packages/nmp_common/tests/auth/test_access_keys.py b/packages/nmp_common/tests/auth/test_access_keys.py index b80f8cb60e..ab94d6ca22 100644 --- a/packages/nmp_common/tests/auth/test_access_keys.py +++ b/packages/nmp_common/tests/auth/test_access_keys.py @@ -4,6 +4,7 @@ from datetime import UTC, datetime from pathlib import Path from typing import Any +from unittest.mock import AsyncMock import httpx import jwt @@ -14,17 +15,18 @@ from cryptography.hazmat.primitives.asymmetric import rsa from nemo_platform_plugin.auth.access_keys.issuer import AccessKeyFeatureDisabledError from nemo_platform_plugin.auth.access_keys.types import AccessKeyCreateRequest -from nmp.common import http_clients from nmp.common.auth.access_keys import ( ACCESS_KEY_TOKEN_TYPE, AccessKeyIssuerService, AccessKeyValidationError, access_key_jwks_uri, clear_access_key_signing_key_cache, + close_access_key_jwks_clients, public_jwk_from_private_key_pem, public_jwk_from_private_key_pem_async, validate_access_key_token, ) +from nmp.common.auth.jwks import AsyncJWKSClient from nmp.common.auth.models import Principal from nmp.common.config import AuthConfig from nmp.common.config.base import AccessKeyConfig, TokenSigningConfig @@ -270,6 +272,16 @@ async def counted_load(**kwargs: Any) -> signing_keys_mod.RSASigningKey: assert load_count == 1 +async def test_close_access_key_jwks_clients_closes_cached_clients() -> None: + client = AsyncMock() + access_keys_mod._ACCESS_KEY_JWKS_CLIENTS["https://auth.example.test/jwks"] = client + + await close_access_key_jwks_clients() + + client.aclose.assert_awaited_once_with() + assert access_keys_mod._ACCESS_KEY_JWKS_CLIENTS == {} + + def test_access_key_private_key_uses_cached_private_key_file_for_token_creation( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -565,38 +577,36 @@ async def test_validate_access_key_token_fetches_remote_jwks_once_with_async_cli class ForbiddenPyJWKClient: def __init__(self, *args, **kwargs): - raise AssertionError("Access-key validation should use the shared async HTTP client") + raise AssertionError("Access-key validation should use async JWKS fetching") - class FakeResponse: - def raise_for_status(self) -> None: - pass + calls = 0 - def json(self) -> dict: - return jwks + async def jwks_handler(request: httpx.Request) -> httpx.Response: + nonlocal calls + assert str(request.url) == jwks_uri + timeout = request.extensions["timeout"] + assert all(value == 10.0 for value in timeout.values()) + calls += 1 + return httpx.Response(200, json=jwks, request=request) - class FakeAsyncClient: - def __init__(self) -> None: - self.calls = 0 + async with httpx.AsyncClient(transport=httpx.MockTransport(jwks_handler)) as http_client: - async def get(self, url: str, *, timeout: float) -> FakeResponse: - assert url == jwks_uri - assert timeout == 10.0 - self.calls += 1 - return FakeResponse() + def jwks_client_factory(uri: str, *, lifespan: int) -> AsyncJWKSClient: + assert uri == jwks_uri + return AsyncJWKSClient(uri, lifespan=lifespan, http_client=http_client) - fake_client = FakeAsyncClient() - monkeypatch.setattr(access_keys_mod, "access_key_jwks_uri", lambda config: jwks_uri) - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: fake_client) - monkeypatch.setattr(access_keys_mod.jwt, "PyJWKClient", ForbiddenPyJWKClient) + monkeypatch.setattr(access_keys_mod, "access_key_jwks_uri", lambda config: jwks_uri) + monkeypatch.setattr(access_keys_mod, "AsyncJWKSClient", jwks_client_factory) + monkeypatch.setattr(access_keys_mod.jwt, "PyJWKClient", ForbiddenPyJWKClient) - first_claims = await validate_access_key_token(config, created.token) - second_claims = await validate_access_key_token(config, created.token) + first_claims = await validate_access_key_token(config, created.token) + second_claims = await validate_access_key_token(config, created.token) - assert first_claims is not None - assert second_claims is not None - assert first_claims.subject == "alice@example.com" - assert second_claims.subject == "alice@example.com" - assert fake_client.calls == 1 + assert first_claims is not None + assert second_claims is not None + assert first_claims.subject == "alice@example.com" + assert second_claims.subject == "alice@example.com" + assert calls == 1 async def test_validate_access_key_token_propagates_remote_jwks_fetch_failure(tmp_path, monkeypatch): @@ -606,23 +616,23 @@ async def test_validate_access_key_token_propagates_remote_jwks_fetch_failure(tm created = await issuer.create_async(AccessKeyCreateRequest(name="ci-intake", expires_in_seconds=None)) jwks_uri = "https://auth.example.test/jwks" - class FakeResponse: - def raise_for_status(self) -> None: - request = httpx.Request("GET", jwks_uri) - response = httpx.Response(503, request=request) - raise httpx.HTTPStatusError("JWKS unavailable", request=request, response=response) + async def jwks_handler(request: httpx.Request) -> httpx.Response: + assert str(request.url) == jwks_uri + timeout = request.extensions["timeout"] + assert all(value == 10.0 for value in timeout.values()) + return httpx.Response(503, request=request) + + async with httpx.AsyncClient(transport=httpx.MockTransport(jwks_handler)) as http_client: - class FakeAsyncClient: - async def get(self, url: str, *, timeout: float) -> FakeResponse: - assert url == jwks_uri - assert timeout == 10.0 - return FakeResponse() + def jwks_client_factory(uri: str, *, lifespan: int) -> AsyncJWKSClient: + assert uri == jwks_uri + return AsyncJWKSClient(uri, lifespan=lifespan, http_client=http_client) - monkeypatch.setattr(access_keys_mod, "access_key_jwks_uri", lambda config: jwks_uri) - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: FakeAsyncClient()) + monkeypatch.setattr(access_keys_mod, "access_key_jwks_uri", lambda config: jwks_uri) + monkeypatch.setattr(access_keys_mod, "AsyncJWKSClient", jwks_client_factory) - with pytest.raises(httpx.HTTPStatusError): - await validate_access_key_token(config, created.token) + with pytest.raises(httpx.HTTPStatusError): + await validate_access_key_token(config, created.token) async def test_validate_access_key_token_rejects_wrong_audience(tmp_path): diff --git a/packages/nmp_common/tests/auth/test_jwks.py b/packages/nmp_common/tests/auth/test_jwks.py index 70ea2be96d..4e7907ee6d 100644 --- a/packages/nmp_common/tests/auth/test_jwks.py +++ b/packages/nmp_common/tests/auth/test_jwks.py @@ -3,12 +3,13 @@ import asyncio from typing import Any +from unittest.mock import AsyncMock, patch +import httpx import jwt import pytest from cryptography.hazmat.primitives.asymmetric import rsa from jwt.algorithms import RSAAlgorithm -from nmp.common import http_clients from nmp.common.auth.jwks import AsyncJWKSClient JWKS_URI = "https://auth.example.test/jwks" @@ -26,92 +27,91 @@ def _token(private_key: Any, *, kid: str | None) -> str: return jwt.encode({"sub": "user"}, private_key, algorithm="RS256", headers=headers) -class FakeResponse: - def __init__(self, jwks: dict[str, Any]) -> None: - self._jwks = jwks +class JWKSResponder: + def __init__(self, jwks_responses: list[dict[str, Any]]) -> None: + self._jwks_responses = jwks_responses + self.calls = 0 - def raise_for_status(self) -> None: - pass + async def __call__(self, request: httpx.Request) -> httpx.Response: + assert str(request.url) == JWKS_URI + timeout = request.extensions["timeout"] + assert all(value == 10.0 for value in timeout.values()) + self.calls += 1 + await asyncio.sleep(0.01) + if len(self._jwks_responses) > 1: + return httpx.Response(200, json=self._jwks_responses.pop(0), request=request) + return httpx.Response(200, json=self._jwks_responses[0], request=request) - def json(self) -> dict[str, Any]: - return self._jwks +async def test_aclose_closes_owned_http_client() -> None: + http_client = AsyncMock(spec=httpx.AsyncClient) + with patch("nmp.common.auth.jwks.DefaultAsyncHttpxClient", return_value=http_client): + client = AsyncJWKSClient(JWKS_URI) -class FakeAsyncClient: - def __init__(self, jwks: dict[str, Any]) -> None: - self._jwks = jwks - self.calls = 0 + await client.aclose() - async def get(self, url: str, *, timeout: float) -> FakeResponse: - assert url == JWKS_URI - assert timeout == 10.0 - self.calls += 1 - await asyncio.sleep(0.01) - return FakeResponse(self._jwks) + http_client.aclose.assert_awaited_once_with() -class SequenceFakeAsyncClient: - def __init__(self, jwks_responses: list[dict[str, Any]]) -> None: - self._jwks_responses = jwks_responses - self.calls = 0 +async def test_aclose_does_not_close_injected_http_client(monkeypatch: pytest.MonkeyPatch) -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + close = AsyncMock() + monkeypatch.setattr(http_client, "aclose", close) - async def get(self, url: str, *, timeout: float) -> FakeResponse: - assert url == JWKS_URI - assert timeout == 10.0 - self.calls += 1 - await asyncio.sleep(0.01) - if len(self._jwks_responses) > 1: - return FakeResponse(self._jwks_responses.pop(0)) - return FakeResponse(self._jwks_responses[0]) + client = AsyncJWKSClient(JWKS_URI, http_client=http_client) + await client.aclose() + + close.assert_not_awaited() + await httpx.AsyncClient.aclose(http_client) @pytest.mark.parametrize("bad_token", ["malformed-token", None]) -async def test_cached_jwks_lookup_does_not_refresh_malformed_or_missing_kid_tokens(monkeypatch, bad_token): +async def test_cached_jwks_lookup_does_not_refresh_malformed_or_missing_kid_tokens(bad_token): private_key, jwk = _rsa_key_and_jwk("known-key") - fake_client = FakeAsyncClient({"keys": [jwk]}) - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: fake_client) - client = AsyncJWKSClient(JWKS_URI) - await client.get_signing_key_from_jwt(_token(private_key, kid="known-key")) + responder = JWKSResponder([{"keys": [jwk]}]) + async with httpx.AsyncClient(transport=httpx.MockTransport(responder)) as http_client: + client = AsyncJWKSClient(JWKS_URI, http_client=http_client) + await client.get_signing_key_from_jwt(_token(private_key, kid="known-key")) - token = bad_token if bad_token is not None else _token(private_key, kid=None) - with pytest.raises(jwt.InvalidTokenError): - await client.get_signing_key_from_jwt(token) + token = bad_token if bad_token is not None else _token(private_key, kid=None) + with pytest.raises(jwt.InvalidTokenError): + await client.get_signing_key_from_jwt(token) - assert fake_client.calls == 1 + assert responder.calls == 1 -async def test_cached_unknown_kid_concurrent_refresh_uses_single_refreshed_jwks(monkeypatch): +async def test_cached_unknown_kid_concurrent_refresh_uses_single_refreshed_jwks(): known_private_key, known_jwk = _rsa_key_and_jwk("known-key") rotated_private_key, rotated_jwk = _rsa_key_and_jwk("rotated-key") - fake_client = SequenceFakeAsyncClient([{"keys": [known_jwk]}, {"keys": [rotated_jwk]}]) - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: fake_client) - client = AsyncJWKSClient(JWKS_URI) - await client.get_signing_key_from_jwt(_token(known_private_key, kid="known-key")) - rotated_token = _token(rotated_private_key, kid="rotated-key") + responder = JWKSResponder([{"keys": [known_jwk]}, {"keys": [rotated_jwk]}]) + async with httpx.AsyncClient(transport=httpx.MockTransport(responder)) as http_client: + client = AsyncJWKSClient(JWKS_URI, http_client=http_client) + await client.get_signing_key_from_jwt(_token(known_private_key, kid="known-key")) + rotated_token = _token(rotated_private_key, kid="rotated-key") - results = await asyncio.gather(*(client.get_signing_key_from_jwt(rotated_token) for _ in range(5))) + results = await asyncio.gather(*(client.get_signing_key_from_jwt(rotated_token) for _ in range(5))) - assert [result.key_id for result in results] == ["rotated-key"] * 5 - assert fake_client.calls == 2 + assert [result.key_id for result in results] == ["rotated-key"] * 5 + assert responder.calls == 2 -async def test_cached_unknown_kid_refresh_is_rate_limited(monkeypatch): +async def test_cached_unknown_kid_refresh_is_rate_limited(): private_key, jwk = _rsa_key_and_jwk("known-key") - fake_client = FakeAsyncClient({"keys": [jwk]}) - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: fake_client) - client = AsyncJWKSClient(JWKS_URI) - await client.get_signing_key_from_jwt(_token(private_key, kid="known-key")) - unknown_token = _token(private_key, kid="unknown-key") + responder = JWKSResponder([{"keys": [jwk]}]) + async with httpx.AsyncClient(transport=httpx.MockTransport(responder)) as http_client: + client = AsyncJWKSClient(JWKS_URI, http_client=http_client) + await client.get_signing_key_from_jwt(_token(private_key, kid="known-key")) + unknown_token = _token(private_key, kid="unknown-key") - results = await asyncio.gather( - *(client.get_signing_key_from_jwt(unknown_token) for _ in range(5)), - return_exceptions=True, - ) + results = await asyncio.gather( + *(client.get_signing_key_from_jwt(unknown_token) for _ in range(5)), + return_exceptions=True, + ) - assert all(isinstance(result, jwt.InvalidTokenError) for result in results) - assert fake_client.calls == 2 + assert all(isinstance(result, jwt.InvalidTokenError) for result in results) + assert responder.calls == 2 - with pytest.raises(jwt.InvalidTokenError): - await client.get_signing_key_from_jwt(unknown_token) + with pytest.raises(jwt.InvalidTokenError): + await client.get_signing_key_from_jwt(unknown_token) - assert fake_client.calls == 2 + assert responder.calls == 2 diff --git a/packages/nmp_common/tests/auth/test_jwt.py b/packages/nmp_common/tests/auth/test_jwt.py index 7a90724b4f..b82c2b5f17 100644 --- a/packages/nmp_common/tests/auth/test_jwt.py +++ b/packages/nmp_common/tests/auth/test_jwt.py @@ -12,7 +12,6 @@ import pytest from cryptography.hazmat.primitives.asymmetric import rsa from jwt.algorithms import RSAAlgorithm -from nmp.common import http_clients from nmp.common.auth.jwt import ( JWTValidator, UnsignedJWTRejectedError, @@ -774,9 +773,27 @@ async def test_jwks_uses_async_jwks_client_cache(self, auth_config): jwks_client_class.assert_called_once() assert jwks_client.get_jwks.await_count == 2 + @pytest.mark.asyncio + async def test_aclose_closes_cached_jwks_client(self, auth_config): + auth_config.oidc.jwks_uri = "https://custom.example.com/jwks" + validator = JWTValidator(auth_config) + jwks_client = MagicMock() + jwks_client.aclose = AsyncMock() + + with patch("nmp.common.auth.jwt.AsyncJWKSClient", return_value=jwks_client): + assert await validator._get_jwks_client() is jwks_client + + await validator.aclose() + + jwks_client.aclose.assert_awaited_once_with() + assert validator._jwks_client is None + @pytest.mark.asyncio async def test_validate_token_fetches_jwks_with_async_client(self, auth_config, monkeypatch): """OIDC JWKS lookup must not use sync PyJWKClient in async validation.""" + from nmp.common.auth.jwks import AsyncJWKSClient + from nmp.common.auth.jwt import _JWKS_CACHE_LIFESPAN + auth_config.oidc.jwks_uri = "https://custom.example.com/jwks" validator = JWTValidator(auth_config) private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) @@ -802,32 +819,31 @@ class ForbiddenPyJWKClient: def __init__(self, *args, **kwargs): raise AssertionError("OIDC validation should use async JWKS fetching") - class FakeResponse: - def raise_for_status(self) -> None: - pass + calls = 0 - def json(self) -> dict: - return {"keys": [jwk]} + async def jwks_handler(request: httpx.Request) -> httpx.Response: + nonlocal calls + assert str(request.url) == "https://custom.example.com/jwks" + timeout = request.extensions["timeout"] + assert all(value == 10.0 for value in timeout.values()) + calls += 1 + return httpx.Response(200, json={"keys": [jwk]}, request=request) - class FakeAsyncClient: - def __init__(self) -> None: - self.calls = 0 + async with httpx.AsyncClient(transport=httpx.MockTransport(jwks_handler)) as http_client: - async def get(self, url: str, *, timeout: float) -> FakeResponse: - assert url == "https://custom.example.com/jwks" - assert timeout == 10.0 - self.calls += 1 - return FakeResponse() + def jwks_client_factory(uri: str, *, lifespan: int) -> AsyncJWKSClient: + assert uri == "https://custom.example.com/jwks" + assert lifespan == _JWKS_CACHE_LIFESPAN + return AsyncJWKSClient(uri, lifespan=lifespan, http_client=http_client) - fake_client = FakeAsyncClient() - monkeypatch.setattr(http_clients, "shared_async_http_client", lambda: fake_client) - monkeypatch.setattr("nmp.common.auth.jwt.PyJWKClient", ForbiddenPyJWKClient, raising=False) + monkeypatch.setattr("nmp.common.auth.jwt.AsyncJWKSClient", jwks_client_factory) + monkeypatch.setattr("nmp.common.auth.jwt.PyJWKClient", ForbiddenPyJWKClient, raising=False) - first_claims = await validator.validate_token(token) - second_claims = await validator.validate_token(token) + first_claims = await validator.validate_token(token) + second_claims = await validator.validate_token(token) - assert first_claims is not None - assert second_claims is not None - assert first_claims.subject == "user123" - assert second_claims.subject == "user123" - assert fake_client.calls == 1 + assert first_claims is not None + assert second_claims is not None + assert first_claims.subject == "user123" + assert second_claims.subject == "user123" + assert calls == 1 diff --git a/packages/nmp_common/tests/auth/test_middleware.py b/packages/nmp_common/tests/auth/test_middleware.py index f23d40afae..6c687923fd 100644 --- a/packages/nmp_common/tests/auth/test_middleware.py +++ b/packages/nmp_common/tests/auth/test_middleware.py @@ -36,6 +36,7 @@ def _cleanup_config_overrides(): """Clean up Configuration overrides set by create_test_app to prevent leaking to other tests.""" yield Configuration.clear_override(AuthConfig) + Configuration.clear_override(PlatformConfig) @pytest.fixture @@ -87,6 +88,7 @@ async def make(handler, base_url: str = "http://platform.example.com"): update={"access_keys": auth_config_oidc_disabled.access_keys.model_copy(update={"enabled": True})} ) Configuration.set_override(config) + Configuration.set_override(PlatformConfig(base_url=base_url, services="")) async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http_client: middleware = AuthorizationMiddleware( FastAPI(), @@ -94,11 +96,7 @@ async def make(handler, base_url: str = "http://platform.example.com"): http_client=http_client, access_key_lifecycle_http_client=http_client, ) - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=PlatformConfig(base_url=base_url, services=""), - ): - yield middleware + yield middleware return make @@ -574,7 +572,7 @@ def test_scoped_access_key_middleware_mapping_is_skipped_when_access_keys_are_di mock_validate.assert_not_called() @pytest.mark.asyncio - async def test_access_key_lifecycle_sdk_uses_auth_service_discovery( + async def test_access_key_lifecycle_callout_uses_injected_client_base_url( self, auth_config_oidc_disabled, monkeypatch: pytest.MonkeyPatch, @@ -611,7 +609,7 @@ def handler(request: httpx.Request) -> httpx.Response: response = await middleware._authenticate_access_key_lifecycle("scoped-access-key") assert isinstance(response, ResolvedBearerToken) - assert requests[0].url == httpx.URL("http://auth.internal:8080/apis/auth/authenticate") + assert requests[0].url == httpx.URL("http://platform.example.com/apis/auth/authenticate") @pytest.mark.asyncio async def test_access_key_lifecycle_callout_allows_active_token(self, access_key_lifecycle_middleware): @@ -671,11 +669,8 @@ def authenticate_access_key(request: httpx.Request) -> httpx.Response: http_client=pdp_client, access_key_lifecycle_http_client=lifecycle_client, ) - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=PlatformConfig(base_url="unix:///tmp/nemo-platform.sock", services=""), - ): - response = await middleware._authenticate_access_key_lifecycle("scoped-access-key") + Configuration.set_override(PlatformConfig(base_url="unix:///tmp/nemo-platform.sock", services="")) + response = await middleware._authenticate_access_key_lifecycle("scoped-access-key") assert isinstance(response, ResolvedBearerToken) assert response.claims.subject == "alice@example.com" @@ -878,11 +873,8 @@ def handler(request: httpx.Request) -> httpx.Response: ) try: + Configuration.set_override(PlatformConfig(base_url="http://platform.example.com", services="")) with ( - patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=PlatformConfig(base_url="http://platform.example.com", services=""), - ), patch("nmp.common.auth.middleware.resolve_bearer_token", new=AsyncMock()) as resolver, patch.object(AuthClient, "authorize_request", autospec=True) as mock_authorize, ): @@ -902,6 +894,70 @@ def handler(request: httpx.Request) -> httpx.Response: finally: asyncio.run(http_client.aclose()) + def test_workload_bearer_uses_authenticate_callout_after_local_resolver_rejects(self, auth_config_enabled): + app = FastAPI() + + @app.get("/whoami") + async def whoami(auth_client: AuthClient = Depends(get_auth_client)): + return { + "principal": auth_client.principal.id, + "email": auth_client.principal.email, + "groups": auth_client.principal.groups, + } + + config = auth_config_enabled.model_copy( + update={ + "oidc": auth_config_enabled.oidc.model_copy( + update={"workload_token_exchange_enabled": True}, + ) + } + ) + Configuration.set_override(config) + Configuration.set_override(PlatformConfig(base_url="http://platform.example.com", services="")) + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + json={ + "principal": "system:serviceaccount:nemo:job", + "email": None, + "groups": ["team-ml"], + "scopes": ["openid", "email"], + "token_kind": "workload_access_token", + }, + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + app.add_middleware( + AuthorizationMiddleware, + service_name="test-service", + access_key_lifecycle_http_client=http_client, + ) + client = TestClient(app, raise_server_exceptions=False) + + try: + with ( + patch("nmp.common.auth.middleware.resolve_bearer_token", new=AsyncMock(return_value=None)) as resolver, + patch.object(AuthClient, "authorize_request", autospec=True) as mock_authorize, + ): + mock_authorize.return_value = MagicMock(allowed=True) + response = client.get("/whoami", headers={"Authorization": "Bearer workload-access-token"}) + + assert response.status_code == 200 + assert response.json() == { + "principal": "system:serviceaccount:nemo:job", + "email": None, + "groups": ["team-ml"], + } + assert requests[0].url == httpx.URL("http://platform.example.com/apis/auth/authenticate") + assert requests[0].headers["authorization"] == "Bearer workload-access-token" + resolver.assert_awaited_once() + mock_authorize.assert_called_once() + finally: + asyncio.run(http_client.aclose()) + def test_bearer_token_sets_auth_client_context_for_service_handler(self, auth_config_enabled): app = FastAPI() diff --git a/packages/nmp_common/tests/client_factory/test_client_factory.py b/packages/nmp_common/tests/client_factory/test_client_factory.py deleted file mode 100644 index 7a905cbb38..0000000000 --- a/packages/nmp_common/tests/client_factory/test_client_factory.py +++ /dev/null @@ -1,394 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Tests for :mod:`nmp.common.client_factory` — the rich NemoClient provider. - -Covers what the platform provider adds over the plugin's env-var default: -per-service URL routing, shared HTTP clients, principal/auth + internal + -OTEL headers, workspace defaults, and test-client injection. -""" - -from unittest.mock import patch - -import httpx -import pytest -from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient -from nemo_platform_plugin.client.types import PreparedRequest -from nemo_platform_plugin.client_provider import NemoClientProvider -from nmp.common import client_factory as cf -from nmp.common.config import Configuration -from nmp.common.observability.otel import scoped_otel_headers - - -@pytest.fixture(autouse=True) -def _reset_client_factory_state(): - """Keep tests order-independent: clear the injected test client and config cache.""" - old = cf._test_http_client - cf._test_http_client = None - Configuration.clear_cache() - try: - yield - finally: - cf._test_http_client = old - Configuration.clear_cache() - - -def _get(path_template: str, **path_params: str) -> PreparedRequest: - return PreparedRequest( - method="GET", - path_template=path_template, - path_params=path_params, - content=None, - content_type=None, - response_type=None, - ) - - -def _mock_client(sink: list[httpx.Request]) -> httpx.Client: - def handler(request: httpx.Request) -> httpx.Response: - sink.append(request) - return httpx.Response(200, json={"ok": True}) - - return httpx.Client(transport=httpx.MockTransport(handler)) - - -# --------------------------------------------------------------------------- -# Sync construction -# --------------------------------------------------------------------------- - - -class TestSyncConstruction: - def test_base_url_from_config(self): - client = cf.get_nemo_client() - assert isinstance(client, NemoClient) - assert client.base_url == str(Configuration.get_platform_config().base_url).rstrip("/") - - def test_service_principal_and_internal_headers(self): - client = cf.get_nemo_client(as_service="evaluator", internal=True) - assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" - assert client._default_headers["X-NMP-Internal"] == "true" - - def test_on_behalf_of(self): - client = cf.get_nemo_client(as_service="svc", on_behalf_of="user@example.com") - assert client._default_headers["X-NMP-Principal-On-Behalf-Of"] == "user@example.com" - - def test_workspace_passthrough(self): - client = cf.get_nemo_client(workspace="team-a") - assert client.workspace == "team-a" - - def test_reuses_shared_sync_http_client(self): - client = cf.get_nemo_client() - assert client._http is cf.shared_sync_http_client() - - def test_explicit_http_client_wins(self): - with httpx.Client() as explicit: - client = cf.get_nemo_client(http_client=explicit) - assert client._http is explicit - - -# --------------------------------------------------------------------------- -# Async construction -# --------------------------------------------------------------------------- - - -class TestAsyncConstruction: - def test_base_url_from_config(self): - client = cf.get_async_nemo_client() - assert isinstance(client, AsyncNemoClient) - assert client.base_url == str(Configuration.get_platform_config().base_url).rstrip("/") - - def test_service_principal_and_internal_headers(self): - client = cf.get_async_nemo_client(as_service="evaluator", internal=True) - assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" - assert client._default_headers["X-NMP-Internal"] == "true" - - def test_falls_back_to_shared_async_client(self): - client = cf.get_async_nemo_client() - assert isinstance(client._http, httpx.AsyncClient) - - -# --------------------------------------------------------------------------- -# URL routing -# --------------------------------------------------------------------------- - - -class TestUrlRouting: - def test_routes_service_path_to_discovered_origin(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-svc:9999") - Configuration.clear_cache() - - captured: list[httpx.Request] = [] - client = cf.get_nemo_client(as_service="entities", internal=True, http_client=_mock_client(captured)) - client.send(_get("/apis/entities/v2/foo")) - - assert str(captured[0].url) == "http://entities-svc:9999/apis/entities/v2/foo" - assert captured[0].headers["X-NMP-Principal-Id"] == "service:entities" - assert captured[0].headers["X-NMP-Internal"] == "true" - - def test_preserves_query_string_when_routing(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-svc:9999") - Configuration.clear_cache() - - captured: list[httpx.Request] = [] - client = cf.get_nemo_client(http_client=_mock_client(captured)) - client.send(_get("/apis/entities/v2/models?limit=5")) - - assert str(captured[0].url) == "http://entities-svc:9999/apis/entities/v2/models?limit=5" - - def test_non_discovered_path_stays_on_platform_origin(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.delenv("NMP_MODELS_URL", raising=False) - Configuration.clear_cache() - - captured: list[httpx.Request] = [] - client = cf.get_nemo_client(http_client=_mock_client(captured)) - client.send(_get("/apis/models/v1/bar")) - - assert str(captured[0].url) == "https://nemo-gateway:8080/apis/models/v1/bar" - - def test_workspace_default_fills_path_param(self): - captured: list[httpx.Request] = [] - client = cf.get_nemo_client(workspace="team-a", http_client=_mock_client(captured)) - client.send(_get("/apis/entities/v2/workspaces/{workspace}/models")) - - assert "/workspaces/team-a/models" in str(captured[0].url) - - -# --------------------------------------------------------------------------- -# Headers / auth -# --------------------------------------------------------------------------- - - -class TestHeadersAuth: - def test_propagates_request_principal_when_no_service(self): - auth_headers = {"X-NMP-Principal-Id": "user@example.com", "X-NMP-Principal-Groups": "g1,g2"} - # _get_default_headers reads the request principal via sdk_factory's binding. - with patch("nmp.common.sdk_factory.get_principal_auth_headers", return_value=auth_headers): - client = cf.get_nemo_client() - assert client._default_headers["X-NMP-Principal-Id"] == "user@example.com" - assert client._default_headers["X-NMP-Principal-Groups"] == "g1,g2" - - def test_merges_otel_propagation_headers_without_adding_internal_auth(self): - with scoped_otel_headers({"traceparent": "00-trace-span-01", "X-NMP-Internal": "true"}): - client = cf.get_nemo_client(as_service="svc") - assert client._default_headers["traceparent"] == "00-trace-span-01" - assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" - assert "X-NMP-Internal" not in client._default_headers - - def test_explicit_auth_headers_win_over_conflicting_otel_context(self): - with scoped_otel_headers( - { - "traceparent": "00-trace-span-01", - "x-nmp-principal-id": "attacker@example.com", - "X-NMP-Principal-Groups": "admins", - "x-NMP-Internal": "false", - } - ): - client = cf.get_async_nemo_client(as_service="evaluator", internal=True) - assert client._default_headers["traceparent"] == "00-trace-span-01" - assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" - assert client._default_headers["X-NMP-Internal"] == "true" - assert all(name.lower() != "x-nmp-principal-groups" for name in client._default_headers) - - def test_no_headers_leaves_default_headers_none(self): - # No service, no principal context, no OTEL, no internal → no default headers. - with patch("nmp.common.sdk_factory.get_principal_auth_headers", return_value={}): - with patch("nmp.common.sdk_factory.principal_from_env", return_value=None): - client = cf.get_nemo_client() - assert client._default_headers == {} - - -# --------------------------------------------------------------------------- -# Test-client injection -# --------------------------------------------------------------------------- - - -class TestTestClientInjection: - def test_async_uses_module_level_test_client(self): - test_client = httpx.AsyncClient(base_url="http://testserver") - cf._test_http_client = test_client - try: - client = cf.get_async_nemo_client(as_service="evaluator") - assert client._http is test_client - finally: - cf._test_http_client = None - - def test_async_explicit_http_client_beats_module_level(self): - module_client = httpx.AsyncClient(base_url="http://module") - explicit = httpx.AsyncClient(base_url="http://explicit") - cf._test_http_client = module_client - try: - client = cf.get_async_nemo_client(http_client=explicit) - assert client._http is explicit - finally: - cf._test_http_client = None - - -# --------------------------------------------------------------------------- -# Provider class -# --------------------------------------------------------------------------- - - -class TestPlatformNemoClientProvider: - def test_satisfies_protocol(self): - assert isinstance(cf.PlatformNemoClientProvider(), NemoClientProvider) - - def test_get_nemo_client_returns_routed_sync_client(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - Configuration.clear_cache() - provider = cf.PlatformNemoClientProvider() - client = provider.get_nemo_client(as_service="svc", internal=True, workspace="ws1") - assert isinstance(client, NemoClient) - assert client.base_url == "https://nemo-gateway:8080" - assert client.workspace == "ws1" - assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" - - def test_get_async_nemo_client_returns_async_client(self): - provider = cf.PlatformNemoClientProvider() - client = provider.get_async_nemo_client(as_service="svc") - assert isinstance(client, AsyncNemoClient) - assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" - - -# --------------------------------------------------------------------------- -# Task client: creator delegation (PR-800 claim 1) -# --------------------------------------------------------------------------- - - -class TestTaskClientDelegation: - def test_task_client_delegates_to_job_creator(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv( - "NMP_PRINCIPAL", - '{"id": "user:alice@acme.com", "email": "alice@acme.com", "groups": ["team-a"]}', - ) - client = cf.get_task_nemo_client("evaluator") - headers = client._default_headers - assert headers["X-NMP-Internal"] == "true" - assert headers["X-NMP-Principal-Id"] == "service:evaluator" - assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" - assert headers["X-NMP-Principal-On-Behalf-Of-Email"] == "alice@acme.com" - assert headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "team-a" - - def test_task_client_without_principal_warns(self, monkeypatch: pytest.MonkeyPatch, caplog): - monkeypatch.delenv("NMP_PRINCIPAL", raising=False) - with caplog.at_level("WARNING"): - client = cf.get_task_nemo_client("evaluator") - assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" - assert "X-NMP-Principal-On-Behalf-Of" not in client._default_headers - assert "without on-behalf-of delegation" in caplog.text - - async def test_async_task_client_delegates(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv( - "NMP_PRINCIPAL", - '{"id": "user:alice@acme.com", "email": "alice@acme.com", "groups": ["team-a"]}', - ) - client = cf.get_async_task_nemo_client("evaluator") - assert client._default_headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" - - def test_provider_exposes_task_methods(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') - provider = cf.PlatformNemoClientProvider() - headers = provider.get_task_nemo_client("evaluator")._default_headers - assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" - - -# --------------------------------------------------------------------------- -# Task client: workload identity (PR-800 claim 2) -# --------------------------------------------------------------------------- - - -class _FakeExchangeProvider: - def get_access_token(self) -> str: - return "exchanged-token" - - async def get_access_token_async(self) -> str: - return "exchanged-token" - - -class TestTaskClientWorkloadIdentity: - @pytest.fixture - def _stub_exchange(self, monkeypatch: pytest.MonkeyPatch): - captured: dict[str, str] = {} - - def _fake(*, base_url, subject_token_file): - captured["base_url"] = base_url - captured["subject_token_file"] = str(subject_token_file) - return _FakeExchangeProvider() - - monkeypatch.setattr( - "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", - _fake, - ) - return captured - - def test_task_client_bootstraps_workload_identity(self, monkeypatch, tmp_path, _stub_exchange): - token_file = tmp_path / "token" - token_file.write_text("subject-token") - monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) - monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") - monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') # ignored in WI mode - Configuration.clear_cache() - - client = cf.get_task_nemo_client("evaluator") - assert isinstance(client._auth, _FakeExchangeProvider) - assert _stub_exchange["base_url"] == "http://platform:8080" - # No trusted principal headers in workload-identity mode. - assert "X-NMP-Principal-Id" not in client._default_headers - assert client._default_headers.get("X-NMP-Internal") == "true" - - def test_uds_does_not_bootstrap_workload_identity(self, monkeypatch, tmp_path, _stub_exchange): - # Matches get_task_sdk exactly: with the WI token file set the task path - # delegates to get_nemo_client(internal=True); on UDS transport that skips - # bearer exchange and propagates the env principal as its own identity - # (no service principal, no bearer auth). - token_file = tmp_path / "token" - token_file.write_text("subject-token") - monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') - Configuration.clear_cache() - - client = cf.get_task_nemo_client("evaluator") - assert client._auth is None - assert client._default_headers["X-NMP-Principal-Id"] == "user:alice@acme.com" - assert "X-NMP-Principal-On-Behalf-Of" not in client._default_headers - - -# --------------------------------------------------------------------------- -# UDS endpoint routing + transport (PR-800 claim 3) -# --------------------------------------------------------------------------- - - -class TestUdsTransport: - def test_uds_base_url_is_normalized_not_pathed(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - Configuration.clear_cache() - client = cf.get_nemo_client() - # base_url is the routable host, not the raw unix:// socket path. - assert client.base_url == "http://nemo-platform.local" - # concatenating an API path yields a valid URL, not a broken one. - assert client.base_url + "/apis/entities/v2/foo" == "http://nemo-platform.local/apis/entities/v2/foo" - - def test_uds_sync_client_binds_socket_transport(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - Configuration.clear_cache() - client = cf.get_nemo_client() - transport = client._http._transport - assert isinstance(transport, httpx.HTTPTransport) - assert transport._pool._uds == "/tmp/nemo-platform.sock" - - async def test_uds_async_client_binds_socket_transport(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - Configuration.clear_cache() - client = cf.get_async_nemo_client() - transport = client._http._transport - assert isinstance(transport, httpx.AsyncHTTPTransport) - assert transport._pool._uds == "/tmp/nemo-platform.sock" - - def test_tcp_client_uses_shared_client(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") - Configuration.clear_cache() - client = cf.get_nemo_client() - assert client.base_url == "http://platform:8080" diff --git a/packages/nmp_common/tests/jobs/test_log_client.py b/packages/nmp_common/tests/jobs/test_log_client.py index 3e4fd695f0..42a55c8e48 100644 --- a/packages/nmp_common/tests/jobs/test_log_client.py +++ b/packages/nmp_common/tests/jobs/test_log_client.py @@ -3,15 +3,17 @@ """Tests for JobLogsClient SDK wrapper and PageCursor.""" +from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from nemo_platform_plugin.client.errors import NotFoundError -from nmp.common.jobs.log_client import JobLogsClient +from nmp.common.jobs.log_client import JobLogsClient, dep_job_logs_client from nmp.common.jobs.schemas import ( PageCursor, PaginationDirection, + PlatformJobLog, PlatformJobLogPage, ) @@ -78,13 +80,13 @@ async def test_query_logs_success(log_client): page = PlatformJobLogPage( data=[ - { - "timestamp": "2024-01-01T12:00:00", - "job": "job-123", - "job_step": "step1", - "job_task": "task1", - "message": "Test log message", - } + PlatformJobLog( + timestamp=datetime.fromisoformat("2024-01-01T12:00:00"), + job="job-123", + job_step="step1", + job_task="task1", + message="Test log message", + ) ], total=1, next_page=None, @@ -186,3 +188,33 @@ def test_sdk_created_in_constructor(): client = JobLogsClient(sdk=sdk) assert client._sdk is sdk mock_adapter.assert_called_once() + + +@pytest.mark.asyncio +async def test_aclose_does_not_close_injected_sdk() -> None: + sdk = MagicMock() + sdk.close = AsyncMock() + with patch("nmp.common.jobs.log_client.client_from_platform"): + client = JobLogsClient(sdk=sdk) + + await client.aclose() + + sdk.close.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_dependency_closes_owned_sdk() -> None: + sdk = MagicMock() + sdk.close = AsyncMock() + with ( + patch("nmp.common.jobs.log_client.get_async_platform_sdk", return_value=sdk), + patch("nmp.common.jobs.log_client.client_from_platform"), + ): + dependency = dep_job_logs_client() + client = await anext(dependency) + + assert isinstance(client, JobLogsClient) + with pytest.raises(StopAsyncIteration): + await anext(dependency) + + sdk.close.assert_awaited_once_with() diff --git a/packages/nmp_common/tests/jobs/test_result_manager.py b/packages/nmp_common/tests/jobs/test_result_manager.py index 39a3626506..a8c25defab 100644 --- a/packages/nmp_common/tests/jobs/test_result_manager.py +++ b/packages/nmp_common/tests/jobs/test_result_manager.py @@ -67,6 +67,70 @@ def test_result_manager_factory_fileset_async(mock_platform_config, mock_fs_clas assert mgr.workspace == "my-workspace" +@patch("nmp.common.jobs.result_manager.get_async_platform_sdk") +def test_result_manager_factory_defaults_to_one_owned_async_sdk(mock_get_sdk): + mock_sdk = MagicMock() + mock_get_sdk.return_value = mock_sdk + + mgr = rm.result_manager_factory(job_name="test-job", workspace="my-workspace") + + mock_get_sdk.assert_called_once_with() + assert mgr.files_sdk is mock_sdk + assert mgr.jobs_sdk is mock_sdk + assert mgr.owns_files_sdk is True + assert mgr.owns_jobs_sdk is False + + +@patch("nmp.common.jobs.result_manager.get_platform_sdk") +def test_result_manager_factory_defaults_to_one_owned_sync_sdk(mock_get_sdk): + mock_sdk = MagicMock() + mock_get_sdk.return_value = mock_sdk + + mgr = rm.result_manager_factory(job_name="test-job", workspace="my-workspace", is_async=False) + + mock_get_sdk.assert_called_once_with() + assert mgr.files_sdk is mock_sdk + assert mgr.jobs_sdk is mock_sdk + assert mgr.owns_files_sdk is True + assert mgr.owns_jobs_sdk is False + + +def test_result_manager_close_closes_owned_sync_sdk_once(): + mock_sdk = MagicMock() + mgr = rm.ResultManager( + job_name="test-job", + workspace="my-workspace", + file_manager_cls=FilesetFileManager, + files_sdk=mock_sdk, + jobs_sdk=mock_sdk, + owns_files_sdk=True, + owns_jobs_sdk=True, + ) + + mgr.close() + + mock_sdk.close.assert_called_once_with() + + +@pytest.mark.asyncio +async def test_result_manager_aclose_closes_owned_async_sdk_once(): + mock_sdk = MagicMock() + mock_sdk.close = AsyncMock() + mgr = rm.AsyncResultManager( + job_name="test-job", + workspace="my-workspace", + file_manager_cls=AsyncFilesetFileManager, + files_sdk=mock_sdk, + jobs_sdk=mock_sdk, + owns_files_sdk=True, + owns_jobs_sdk=True, + ) + + await mgr.aclose() + + mock_sdk.close.assert_awaited_once_with() + + @pytest.mark.asyncio @patch("nmp.common.jobs.result_manager.result_manager_factory") async def test_download_from_result_info(mock_factory, tmp_path, mock_sdk): @@ -76,6 +140,7 @@ async def test_download_from_result_info(mock_factory, tmp_path, mock_sdk): mock_result_manager = MagicMock() mock_result_manager.download_artifact = AsyncMock(return_value=TmpDirPath(tmp_dir=tmp_path, path=test_file)) + mock_result_manager.aclose = AsyncMock() mock_factory.return_value = mock_result_manager await rm.download_from_result_info( @@ -91,6 +156,7 @@ async def test_download_from_result_info(mock_factory, tmp_path, mock_sdk): assert call_kwargs["job_name"] == "test-job" assert call_kwargs["workspace"] == "my-workspace" assert call_kwargs["files_sdk"] is mock_sdk + mock_result_manager.aclose.assert_awaited_once_with() def test_result_remote_path_nests_under_base(mock_sdk, mock_nmp_sdk): @@ -263,16 +329,13 @@ async def test_create_result_wraps_transport_errors_async( @pytest.mark.asyncio @patch("nmp.common.jobs.result_manager.result_manager_factory") -@patch("nmp.common.jobs.result_manager.get_async_platform_sdk") -async def test_download_from_result_info_defaults_sdk(mock_get_sdk, mock_factory, tmp_path): - """Test that download_from_result_info auto-creates SDK when files_sdk is None.""" - mock_sdk = MagicMock() - mock_get_sdk.return_value = mock_sdk - +async def test_download_from_result_info_defaults_sdk(mock_factory, tmp_path): + """Test that download_from_result_info lets the manager own a default SDK.""" test_file = tmp_path / "artifact.bin" test_file.write_bytes(b"test content") mock_mgr = MagicMock() mock_mgr.download_artifact = AsyncMock(return_value=TmpDirPath(tmp_dir=tmp_path, path=test_file)) + mock_mgr.aclose = AsyncMock() mock_factory.return_value = mock_mgr await rm.download_from_result_info( @@ -282,9 +345,9 @@ async def test_download_from_result_info_defaults_sdk(mock_get_sdk, mock_factory workspace="workspace", ) - mock_get_sdk.assert_called_once_with() mock_factory.assert_called_once_with( job_name="test-job", workspace="workspace", - files_sdk=mock_sdk, + files_sdk=None, ) + mock_mgr.aclose.assert_awaited_once_with() diff --git a/packages/nmp_common/tests/nmp_common/test_common_config.py b/packages/nmp_common/tests/nmp_common/test_common_config.py index 047a419827..afec6d031d 100644 --- a/packages/nmp_common/tests/nmp_common/test_common_config.py +++ b/packages/nmp_common/tests/nmp_common/test_common_config.py @@ -42,16 +42,16 @@ def test_platform_config_to_shared_envvars(self): assert "NMP_BASE_URL" in envvars assert "NMP_JOBS_URL" in envvars - def test_service_url_env_var_populates_service_discovery(self, monkeypatch): - """NMP__URL env vars are merged into service_discovery.""" + def test_service_url_env_var_does_not_populate_service_discovery(self, monkeypatch): + """NMP__URL env vars are resolved by endpoint factories, not config.""" monkeypatch.setenv("NMP_FILES_URL", "http://files:8000") config = PlatformConfig() - assert config.get_service_url("files") == "http://files:8000" - assert config.service_discovery["files"] == "http://files:8000" + assert config.get_service_url("files") == "http://localhost:8080" + assert "files" not in config.service_discovery - def test_service_url_env_var_overrides_config_file(self, monkeypatch): - """Env var NMP_*_URL overrides or adds to file service_discovery.""" + def test_service_url_env_var_does_not_override_config_file(self, monkeypatch): + """Env var NMP_*_URL is not merged into file service_discovery.""" monkeypatch.setenv("NMP_JOBS_URL", "http://jobs-from-env:9000") settings = { "platform": { @@ -61,9 +61,9 @@ def test_service_url_env_var_overrides_config_file(self, monkeypatch): config = Configuration.global_settings_to_service_config(settings, PlatformConfig) assert config.get_service_url("files") == "http://files-from-file:8080" - assert config.get_service_url("jobs") == "http://jobs-from-env:9000" + assert config.get_service_url("jobs") == "http://localhost:8080" assert config.service_discovery["files"] == "http://files-from-file:8080" - assert config.service_discovery["jobs"] == "http://jobs-from-env:9000" + assert "jobs" not in config.service_discovery def test_base_url_env_var_not_in_service_discovery(self, monkeypatch): """NMP_BASE_URL sets base_url and is not added to service_discovery.""" diff --git a/packages/nmp_common/tests/nmp_common/test_common_service.py b/packages/nmp_common/tests/nmp_common/test_common_service.py index 4339f7eb3f..2c55e6a5ce 100644 --- a/packages/nmp_common/tests/nmp_common/test_common_service.py +++ b/packages/nmp_common/tests/nmp_common/test_common_service.py @@ -5,26 +5,23 @@ import asyncio import threading -import time from concurrent.futures import ThreadPoolExecutor -from threading import Barrier, Event, Lock +from threading import Barrier from types import SimpleNamespace from typing import List from unittest.mock import AsyncMock, patch import httpx +import nmp.common.service as common_service import pytest from fastapi import APIRouter, Depends, FastAPI from fastapi.testclient import TestClient from nemo_platform import AsyncNeMoPlatform -from nemo_platform_plugin.client.client import AsyncNemoClient -from nemo_platform_plugin.dependencies import get_nemo_client as plugin_get_nemo_client from nmp.common.config import PlatformConfig from nmp.common.observability.otel import scoped_otel_headers from nmp.common.service import DependencyProvider, RouterConfig, Service from nmp.common.service import __all__ as service_exports -from nmp.common.service import get_nemo_client as facade_get_nemo_client -from nmp.common.service.dependencies import get_nemo_client +from nmp.common.service.dependencies import get_sdk_client def _route_paths(app: FastAPI) -> set[str]: @@ -33,8 +30,9 @@ def _route_paths(app: FastAPI) -> set[str]: queue = list(app.routes) while queue: route = queue.pop() - if hasattr(route, "path"): - paths.add(route.path) + path = getattr(route, "path", None) + if isinstance(path, str): + paths.add(path) fn = getattr(route, "effective_candidates", None) if callable(fn): queue.extend(fn()) # type: ignore[arg-type] @@ -124,6 +122,7 @@ def test_service_create_app(self): assert app is not None assert app.title == service.title assert app.version == service.version + assert app.openapi()["info"]["description"] == service.description def test_service_app_property_caches(self): """Test app property returns cached instance.""" @@ -192,7 +191,8 @@ def handler(request: httpx.Request) -> httpx.Response: provider = DependencyProvider() provider._platform_config = PlatformConfig(base_url="http://platform.local") async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: - provider._http_client = client + provider.configure_http_client(client) + provider.initialize() service = MockService(dependency_provider=provider) ready = await service.wait_for_service_ready("entities", timeout=1.0, poll_interval=0) @@ -212,7 +212,8 @@ def handler(request: httpx.Request) -> httpx.Response: provider = DependencyProvider() provider._platform_config = PlatformConfig(base_url="http://platform.local") async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: - provider._http_client = client + provider.configure_http_client(client) + provider.initialize() service = MockService(dependency_provider=provider) ready = await service.wait_for_service_ready("models", timeout=1.0, poll_interval=0) @@ -240,42 +241,17 @@ def test_init(self): """Test DependencyProvider initialization.""" provider = DependencyProvider() assert provider._sdk_client is None - assert provider._http_client is None + assert provider._configured_http_client is None - def test_nemo_client_dependency_is_exported_with_exact_plugin_identity(self): - assert get_nemo_client is plugin_get_nemo_client - assert facade_get_nemo_client is plugin_get_nemo_client - assert "get_nemo_client" in service_exports + def test_generic_nemo_client_dependency_is_not_exported(self): + assert not hasattr(common_service, "get_nemo_client") + assert "get_nemo_client" not in service_exports - def test_setup_dependencies_registers_nemo_client_override(self): - provider = DependencyProvider() - app = FastAPI() - - provider.setup_dependencies(app, MockService()) - - assert app.dependency_overrides[get_nemo_client] == provider.get_request_scoped_nemo_client - - @pytest.mark.asyncio - @pytest.mark.parametrize("first_client", ["sdk", "nemo"], ids=["sdk-first", "nemo-first"]) - async def test_sdk_and_nemo_clients_share_provider_transport_regardless_of_order(self, first_client: str): - provider = DependencyProvider() - - if first_client == "sdk": - sdk = provider.get_request_scoped_sdk() - nemo = provider.get_request_scoped_nemo_client() - else: - nemo = provider.get_request_scoped_nemo_client() - sdk = provider.get_request_scoped_sdk() - - assert sdk._client is provider.get_http_client() - assert nemo._http is provider.get_http_client() - - await provider.close() - - def test_request_scoped_nemo_clients_are_distinct_and_share_transport(self): + def test_request_scoped_sdks_are_distinct_and_share_transport(self): provider = DependencyProvider() transport = AsyncMock(spec=httpx.AsyncClient) - provider._http_client = transport + provider.configure_http_client(transport) + provider.initialize() with patch( "nmp.common.sdk_factory.get_principal_auth_headers", @@ -285,42 +261,55 @@ def test_request_scoped_nemo_clients_are_distinct_and_share_transport(self): }, ): with scoped_otel_headers({"traceparent": "00-trace-one-span-one-01"}): - first = provider.get_request_scoped_nemo_client() + first = provider.get_request_scoped_sdk() with patch( "nmp.common.sdk_factory.get_principal_auth_headers", return_value={"X-NMP-Principal-Id": "user-two@example.com"}, ): with scoped_otel_headers({"traceparent": "00-trace-two-span-two-01"}): - second = provider.get_request_scoped_nemo_client() + second = provider.get_request_scoped_sdk() assert first is not second - assert first._http is transport - assert second._http is transport - assert first._default_headers["X-NMP-Principal-Id"] == "user-one@example.com" - assert first._default_headers["X-NMP-Principal-On-Behalf-Of"] == "delegate-one@example.com" - assert first._default_headers["traceparent"] == "00-trace-one-span-one-01" - assert second._default_headers["X-NMP-Principal-Id"] == "user-two@example.com" - assert second._default_headers["traceparent"] == "00-trace-two-span-two-01" + assert first._client is transport + assert second._client is transport + assert first.default_headers["X-NMP-Principal-Id"] == "user-one@example.com" + assert first.default_headers["X-NMP-Principal-On-Behalf-Of"] == "delegate-one@example.com" + assert first.default_headers["traceparent"] == "00-trace-one-span-one-01" + assert second.default_headers["X-NMP-Principal-Id"] == "user-two@example.com" + assert second.default_headers["traceparent"] == "00-trace-two-span-two-01" @pytest.mark.asyncio - async def test_close_closes_shared_sdk_and_nemo_transport_exactly_once(self): + async def test_close_closes_owned_sdk_transport_exactly_once(self, monkeypatch: pytest.MonkeyPatch): + from nmp.common.service import base as service_base + provider = DependencyProvider() - transport = CloseCountingAsyncClient() - provider._http_client = transport + created: list[CloseCountingAsyncClient] = [] + + def create_transport() -> CloseCountingAsyncClient: + transport = CloseCountingAsyncClient() + created.append(transport) + return transport + + endpoint = SimpleNamespace( + connect_base_url="http://platform.test", + async_sdk_http_client=lambda timeout=None: create_transport(), + ) + monkeypatch.setattr(service_base, "resolve_platform_endpoint", lambda _platform_config=None: endpoint) + + provider.initialize() sdk = provider.get_sdk_client() - nemo = provider.get_request_scoped_nemo_client() await provider.close() await provider.close() - assert sdk._client is transport - assert nemo._http is transport - assert transport.close_count == 1 - assert provider._http_client is None + assert created == [sdk._client] + assert created[0].close_count == 1 + assert provider._configured_http_client is None + assert provider._service_http_client is None assert provider._sdk_client is None @pytest.mark.asyncio - async def test_concurrent_first_dependency_resolution_creates_one_transport_and_sdk( + async def test_concurrent_dependency_resolution_reuses_initialized_transport_and_sdk( self, monkeypatch: pytest.MonkeyPatch ): from nmp.common import sdk_factory @@ -328,83 +317,72 @@ async def test_concurrent_first_dependency_resolution_creates_one_transport_and_ provider = DependencyProvider() resolution_ready = Barrier(13) - factory_started = Event() - release_factory = Event() created: list[CloseCountingAsyncClient] = [] - created_lock = Lock() - def resolve_dependency(index: int) -> AsyncNeMoPlatform | AsyncNemoClient: + def resolve_dependency() -> AsyncNeMoPlatform: resolution_ready.wait(timeout=5) - factory = provider.get_request_scoped_sdk if index % 2 == 0 else provider.get_request_scoped_nemo_client - return factory() + return provider.get_request_scoped_sdk() def create_transport() -> CloseCountingAsyncClient: transport = CloseCountingAsyncClient() - with created_lock: - created.append(transport) - factory_started.set() - assert release_factory.wait(timeout=5) + created.append(transport) return transport - endpoint = SimpleNamespace(async_sdk_http_client=lambda: create_transport()) - monkeypatch.setattr(service_base, "resolve_platform_endpoint", lambda: endpoint) + endpoint = SimpleNamespace( + connect_base_url="http://platform.test", + async_sdk_http_client=lambda timeout=None: create_transport(), + ) + monkeypatch.setattr(service_base, "resolve_platform_endpoint", lambda _platform_config=None: endpoint) with patch.object( - sdk_factory, "get_async_platform_sdk", wraps=sdk_factory.get_async_platform_sdk + sdk_factory, "_get_async_platform_sdk_for_endpoint", wraps=sdk_factory._get_async_platform_sdk_for_endpoint ) as sdk_factory_call: + provider.initialize() with ThreadPoolExecutor(max_workers=12) as executor: - futures = [executor.submit(resolve_dependency, index) for index in range(12)] + futures = [executor.submit(resolve_dependency) for _index in range(12)] resolution_ready.wait(timeout=5) - assert factory_started.wait(timeout=5) - time.sleep(0.05) - release_factory.set() clients = [future.result(timeout=5) for future in futures] assert sdk_factory_call.call_count == 1 transport = provider.get_http_client() - sdk_clients = [client for client in clients if isinstance(client, AsyncNeMoPlatform)] - nemo_clients = [client for client in clients if isinstance(client, AsyncNemoClient)] assert created == [transport] - assert len({id(client) for client in sdk_clients}) == 1 - assert all(client._client is transport for client in sdk_clients) - assert all(client._http is transport for client in nemo_clients) + assert len({id(client) for client in clients}) == 1 + assert all(client._client is transport for client in clients) await provider.close() assert created[0].close_count == 1 @pytest.mark.asyncio - async def test_fastapi_caches_nemo_client_within_request_and_isolates_requests(self): + async def test_fastapi_caches_sdk_dependency_within_request(self): provider = DependencyProvider() app = FastAPI() provider.setup_dependencies(app, MockService()) - resolved: list[tuple[AsyncNemoClient, AsyncNemoClient]] = [] + provider.initialize() + resolved: list[tuple[AsyncNeMoPlatform, AsyncNeMoPlatform]] = [] @app.get("/clients") async def clients( - first: AsyncNemoClient = Depends(get_nemo_client), - second: AsyncNemoClient = Depends(get_nemo_client), + first: AsyncNeMoPlatform = Depends(get_sdk_client), + second: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> dict[str, bool]: resolved.append((first, second)) return {"same": first is second} async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: - first_response = await client.get("/clients") - second_response = await client.get("/clients") + response = await client.get("/clients") - assert first_response.json() == {"same": True} - assert second_response.json() == {"same": True} + assert response.json() == {"same": True} assert resolved[0][0] is resolved[0][1] - assert resolved[1][0] is resolved[1][1] - assert resolved[0][0] is not resolved[1][0] - assert resolved[0][0]._http is resolved[1][0]._http + assert resolved[0][0]._client is provider.get_http_client() await provider.close() @pytest.mark.asyncio async def test_service_principal_sdk_shares_provider_transport(self): provider = DependencyProvider() + provider.initialize() cached_sdk = provider.get_sdk_client() service_sdk = provider.get_sdk_client(as_service="entities") @@ -430,6 +408,20 @@ def test_service_has_provider(self): assert service.dependency_provider is not None assert isinstance(service.dependency_provider, DependencyProvider) + def test_service_configures_provider_service_name_for_entity_client_headers(self): + """Test Service configures the provider before entity SDK creation.""" + provider = DependencyProvider() + service = MockService(dependency_provider=provider) + + with patch.object(service.dependency_provider, "get_sdk_client") as mock_sdk: + mock_base_sdk = mock_sdk.return_value + + service.dependency_provider._get_entity_sdk_on_behalf_of() + + call_kwargs = mock_base_sdk.with_options.call_args + headers = call_kwargs.kwargs.get("set_default_headers") or call_kwargs[1].get("set_default_headers") + assert headers["X-NMP-Principal-Id"] == "service:test-service" + class LifecycleService(MockService): def __init__(self): diff --git a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py index eafb54651e..a3d2800af5 100644 --- a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py +++ b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py @@ -10,7 +10,7 @@ from fastapi import FastAPI from nemo_platform_plugin.entities.client import AsyncEntitiesClient from nmp.common.auth import AuthClient, Principal -from nmp.common.config import Configuration +from nmp.common.config import PlatformConfig from nmp.common.service import DependencyProvider from nmp.common.service.dependencies import ( get_effective_principal_id, @@ -20,84 +20,127 @@ ) -def test_get_http_client_caches_endpoint_client() -> None: +def test_get_http_client_requires_initialization() -> None: provider = DependencyProvider() - client = MagicMock() - - with patch( - "nmp.common.service.base.resolve_platform_endpoint", - ) as resolve: - resolve.return_value.async_sdk_http_client.return_value = client - first = provider.get_http_client() - second = provider.get_http_client() - - assert first is client - assert second is client - resolve.assert_called_once_with() - resolve.return_value.async_sdk_http_client.assert_called_once_with() - - -def _uds_of(client: httpx.AsyncClient) -> str | None: - """Socket path bound to the client's transport pool, or None for TCP.""" - return getattr(client._transport._pool, "_uds", None) - - -def test_get_http_client_binds_uds_transport(monkeypatch: pytest.MonkeyPatch) -> None: - """Under a unix:// endpoint the provider-owned client must be socket-bound. - - Regression: the provider used to build a plain TCP DefaultAsyncHttpxClient - and inject it into the SDK/NemoClient factories, which then skipped their - own UDS selection, so service-to-service calls went to http://nemo-platform - .local over TCP (_uds=None) and failed. - """ - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - Configuration.clear_cache() - try: - provider = DependencyProvider() - assert _uds_of(provider.get_http_client()) == "/tmp/nemo-platform.sock" - finally: - Configuration.clear_cache() - - -def test_request_scoped_nemo_client_binds_uds_transport(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") - Configuration.clear_cache() - try: - provider = DependencyProvider() - client = provider.get_request_scoped_nemo_client() - assert _uds_of(client._http) == "/tmp/nemo-platform.sock" - finally: - Configuration.clear_cache() - - -def test_tcp_endpoint_client_is_not_socket_bound(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") - Configuration.clear_cache() - try: - provider = DependencyProvider() - assert _uds_of(provider.get_http_client()) is None - finally: - Configuration.clear_cache() + with pytest.raises(RuntimeError, match="DependencyProvider is not initialized"): + provider.get_http_client() + + +def test_initialize_creates_cached_sdk_from_endpoint() -> None: + provider = DependencyProvider() + platform_config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + provider._platform_config = platform_config + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + http_client = MagicMock(name="http_client") + platform_endpoint.async_sdk_http_client.return_value = http_client + sdk = MagicMock(name="sdk") + sdk.http_client = http_client + + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint) as resolve_endpoint, + patch("nmp.common.sdk_factory._get_async_platform_sdk_for_endpoint", return_value=sdk) as sdk_factory, + ): + provider.initialize() + provider.initialize() + + assert provider.get_http_client() is http_client + assert provider.get_sdk_client() is sdk + resolve_endpoint.assert_called_once_with(platform_config) + platform_endpoint.async_sdk_http_client.assert_not_called() + sdk_factory.assert_called_once_with(platform_endpoint, http_client=None) + + +def test_configure_http_client_is_used_for_cached_sdk_client() -> None: + provider = DependencyProvider() + http_client = MagicMock(name="http_client") + sdk = MagicMock(name="sdk") + platform_config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + provider._platform_config = platform_config + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + sdk.http_client = http_client + + provider.configure_http_client(http_client) + + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint), + patch("nmp.common.sdk_factory._get_async_platform_sdk_for_endpoint", return_value=sdk) as sdk_factory, + ): + provider.initialize() + assert provider.get_sdk_client() is sdk + assert provider.get_sdk_client() is sdk + + platform_endpoint.async_sdk_http_client.assert_not_called() + sdk_factory.assert_called_once_with(platform_endpoint, http_client=http_client) + + +def test_configure_platform_endpoint_is_used_during_initialization() -> None: + provider = DependencyProvider() + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + http_client = MagicMock(name="http_client") + platform_endpoint.async_sdk_http_client.return_value = http_client + sdk = MagicMock(name="sdk") + sdk.http_client = http_client + + provider.configure_platform_endpoint(platform_endpoint) + + with patch("nmp.common.sdk_factory._get_async_platform_sdk_for_endpoint", return_value=sdk): + provider.initialize() + + assert provider.get_http_client() is http_client + assert provider.get_sdk_client() is sdk + + +def test_configure_http_client_rejects_changes_after_sdk_created() -> None: + provider = DependencyProvider() + sdk = MagicMock(name="sdk") + http_client = MagicMock(name="http_client") + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + platform_endpoint.async_sdk_http_client.return_value = http_client + sdk.http_client = http_client + + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint), + patch("nmp.common.sdk_factory._get_async_platform_sdk_for_endpoint", return_value=sdk), + ): + provider.initialize() + + with pytest.raises(RuntimeError, match="Cannot configure DependencyProvider HTTP client after initialization"): + provider.configure_http_client(MagicMock(name="late_http_client")) def test_get_sdk_client_caches_request_sdk_and_creates_fresh_service_sdk() -> None: provider = DependencyProvider() request_sdk = MagicMock(name="request_sdk") service_sdk = MagicMock(name="service_sdk") + request_sdk.token_provider = None + request_sdk.with_options.return_value = service_sdk + http_client = MagicMock(name="http_client") + request_sdk.http_client = http_client + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + platform_endpoint.async_sdk_http_client.return_value = http_client - with patch("nmp.common.sdk_factory.get_async_platform_sdk", side_effect=[request_sdk, service_sdk]) as factory: + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint), + patch("nmp.common.sdk_factory._get_async_platform_sdk_for_endpoint", return_value=request_sdk) as sdk_factory, + ): + provider.initialize() assert provider.get_sdk_client() is request_sdk assert provider.get_sdk_client() is request_sdk assert provider.get_sdk_client(as_service="jobs") is service_sdk - # The provider now shares its pooled HTTP client with every SDK it builds. - http_client = provider.get_http_client() - assert factory.call_args_list[0].kwargs == {"http_client": http_client} - assert factory.call_args_list[1].kwargs == { - "as_service": "jobs", - "internal": True, - "http_client": http_client, - } + sdk_factory.assert_called_once_with(platform_endpoint, http_client=None) + request_sdk.with_options.assert_called_once_with( + set_default_headers={ + "X-NMP-Internal": "true", + "X-NMP-Principal-Id": "service:jobs", + }, + http_client=http_client, + ) def test_setup_dependencies_registers_fastapi_overrides() -> None: @@ -132,20 +175,68 @@ def test_get_effective_principal_id_uses_delegated_identity() -> None: @pytest.mark.asyncio async def test_close_closes_managed_clients_and_clears_references() -> None: provider = DependencyProvider() - http_client = MagicMock() - http_client.aclose = AsyncMock() sdk = MagicMock() sdk.close = AsyncMock() - provider._http_client = http_client + provider._service_http_client = AsyncMock(spec=httpx.AsyncClient) + provider._sdk_client = sdk + + await provider.close() + + sdk.close.assert_awaited_once_with() + assert provider._service_http_client is None + assert provider._configured_http_client is None + assert provider._sdk_client is None + + +@pytest.mark.asyncio +async def test_close_closes_endpoint_owned_http_client() -> None: + provider = DependencyProvider() + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + platform_endpoint.async_sdk_http_client.return_value = http_client + + provider.configure_platform_endpoint(platform_endpoint) + provider.initialize() + + await provider.close() + + assert http_client.is_closed + assert provider._service_http_client is None + assert provider._configured_http_client is None + assert provider._sdk_client is None + + +@pytest.mark.asyncio +async def test_close_leaves_configured_http_client_open() -> None: + provider = DependencyProvider() + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + + provider.configure_platform_endpoint(platform_endpoint) + provider.configure_http_client(http_client) + provider.initialize() + + await provider.close() + + assert not http_client.is_closed + assert provider._service_http_client is None + assert provider._configured_http_client is None + assert provider._sdk_client is None + await http_client.aclose() + + +@pytest.mark.asyncio +async def test_close_closes_cached_sdk_when_no_factory_exists() -> None: + provider = DependencyProvider() + sdk = MagicMock() + sdk.close = AsyncMock() provider._sdk_client = sdk await provider.close() - # close() owns and closes the shared HTTP transport; the SDK borrows that - # transport, so it is dropped without a separate sdk.close(). - http_client.aclose.assert_awaited_once_with() - sdk.close.assert_not_awaited() - assert provider._http_client is None + sdk.close.assert_awaited_once_with() assert provider._sdk_client is None diff --git a/packages/nmp_common/tests/sdk_factory/test_sdk.py b/packages/nmp_common/tests/sdk_factory/test_sdk.py index a2f0bda99c..932264a528 100644 --- a/packages/nmp_common/tests/sdk_factory/test_sdk.py +++ b/packages/nmp_common/tests/sdk_factory/test_sdk.py @@ -7,22 +7,20 @@ import httpx import pytest -from nemo_platform_ext.auth.helpers import NMPOIDCConfig -from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform +from nemo_platform_plugin.client.auth import TokenProviderAuth from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR -from nemo_platform_plugin.jobs.client import JobsClient -from nmp.common.config import Configuration, PlatformConfig -from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client +from nemo_platform_plugin.client.oidc import NMPOIDCConfig +from nmp.common.config import Configuration from nmp.common.sdk_factory import ( - PlatformRequestRouter, get_async_platform_sdk, get_async_task_sdk, get_entity_parts, get_platform_sdk, get_request_scoped_sdk, get_sdk_on_behalf_of, + get_service_scoped_sdk, get_task_sdk, - resolve_platform_request_url, ) @@ -37,24 +35,25 @@ def _workload_oidc_config() -> NMPOIDCConfig: ) -@pytest.fixture(autouse=True) -def _clear_sdk_factory_test_client(): - """Clear SDK factory state before each test so config-based SDK behavior is asserted. +def _apply_sync_auth(sdk: NeMoPlatform, request: httpx.Request) -> None: + assert sdk.token_provider is not None + auth_flow = TokenProviderAuth(sdk.token_provider).sync_auth_flow(request) + assert next(auth_flow) is request - When _test_http_client is set (e.g. by another test's create_test_client), the SDK - is created with base_url='http://testserver' and no request router, which breaks tests - that assert on base_url or service routing. Clearing it keeps tests order-independent - and ensures sdk_factory tests always exercise the config path. - """ - import nmp.common.sdk_factory as sdk_factory_module - old = sdk_factory_module._test_http_client - sdk_factory_module._test_http_client = None +async def _apply_async_auth(sdk: AsyncNeMoPlatform, request: httpx.Request) -> None: + assert sdk.token_provider is not None + auth_flow = TokenProviderAuth(sdk.token_provider).async_auth_flow(request) + assert await anext(auth_flow) is request + + +@pytest.fixture(autouse=True) +def _clear_sdk_factory_config(): + """Clear SDK factory config state before each test.""" Configuration.clear_cache() try: yield finally: - sdk_factory_module._test_http_client = old Configuration.clear_cache() @@ -70,7 +69,7 @@ def test_get_platform_sdk(): def test_get_platform_sdk_keeps_platform_base_url_for_local_services(monkeypatch: pytest.MonkeyPatch): - """The SDK base URL remains the platform entrypoint; per-service routing handles local APIs.""" + """The SDK base URL remains the platform entrypoint.""" monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICES", "auth") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") @@ -113,14 +112,14 @@ def capture_request(request: httpx.Request) -> httpx.Response: sdk = get_platform_sdk(http_client=http_client) assert str(sdk.base_url).rstrip("/") == "http://nemo-platform-api:8080" - client_from_platform(sdk, JobsClient).list_jobs(workspace="default") + sdk.jobs.list(workspace="default") assert len(captured_requests) == 1 assert str(captured_requests[0].url) == "http://nemo-platform-api:8080/apis/jobs/v2/workspaces/default/jobs" -def test_get_platform_sdk_routes_local_service_path_to_process_listener(monkeypatch: pytest.MonkeyPatch): - """Requests for APIs hosted in this process bypass the platform entrypoint.""" +def test_get_platform_sdk_does_not_route_local_service_paths(monkeypatch: pytest.MonkeyPatch): + """SDK instances build URLs only from their configured base URL.""" monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICES", "auth") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") @@ -130,21 +129,12 @@ def test_get_platform_sdk_routes_local_service_path_to_process_listener(monkeypa sdk = get_platform_sdk() prepared = sdk._prepare_url("https://nemo-gateway:8080/apis/auth/v2/authz/allow") - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" + assert prepared.scheme == "https" + assert prepared.host == "nemo-gateway" assert prepared.port == 8080 assert prepared.path == "/apis/auth/v2/authz/allow" -def test_get_platform_sdk_uses_uds_endpoint_from_base_url(): - config = PlatformConfig(base_url="unix:///tmp/nemo-platform.sock") # type: ignore[abstract] - - with patch("nmp.common.sdk_factory.Configuration.get_platform_config", return_value=config): - sdk = get_platform_sdk() - - assert sdk.base_url == "http://nemo-platform.local" - - def test_get_platform_sdk_with_service_principal(): """Test get_platform_sdk with as_service parameter.""" sdk = get_platform_sdk(as_service="my-service") @@ -174,6 +164,37 @@ def test_get_platform_sdk_internal_flag(): assert sdk.default_headers["X-NMP-Internal"] == "true" +def test_get_platform_sdk_closes_factory_owned_http_client(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + sdk = get_platform_sdk() + http_client = sdk._client + + sdk.close() + + assert http_client.is_closed + + +def test_get_platform_sdk_does_not_close_injected_http_client() -> None: + with httpx.Client(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) as http_client: + sdk = get_platform_sdk(base_url="http://nmp.example.test", http_client=http_client) + + sdk.close() + + assert not http_client.is_closed + + +def test_service_scoped_sdk_close_does_not_close_base_http_client() -> None: + base_sdk = get_platform_sdk(base_url="http://nmp.example.test") + scoped_sdk = get_service_scoped_sdk(base_sdk, "jobs") + + try: + scoped_sdk.close() + + assert not base_sdk._client.is_closed + finally: + base_sdk.close() + + def test_get_platform_sdk_uses_workload_identity_when_token_file_configured(monkeypatch: pytest.MonkeyPatch, tmp_path): """Workload-token task environments should let the generated SDK inject Bearer auth.""" subject_token_file = tmp_path / "workload-token" @@ -189,17 +210,18 @@ def token_exchange_grant(**kwargs): monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) monkeypatch.setattr( - "nemo_platform_ext.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config() + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), ) - monkeypatch.setattr("nemo_platform_ext.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + monkeypatch.setattr("nemo_platform_plugin.client.oidc.token_exchange_grant", token_exchange_grant) sdk = get_platform_sdk() try: request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") - sdk._client._event_hooks["request"][0](request) + _apply_sync_auth(sdk, request) finally: sdk.close() @@ -208,47 +230,176 @@ def token_exchange_grant(**kwargs): assert exchange_requests[0]["subject_token"] == "subject-token-from-file" -def test_get_async_platform_sdk(): - """Test get_async_platform_sdk basic functionality (config path: SDK base_url matches platform config).""" - sdk = get_async_platform_sdk() +def test_get_platform_sdk_rejects_workload_identity_with_principal_env(monkeypatch: pytest.MonkeyPatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") - assert sdk is not None - assert hasattr(sdk, "base_url") - # Normalize to str: SDK may expose URL object, config may be str; both environments - expected = Configuration.get_platform_config().base_url - assert str(sdk.base_url).rstrip("/") == str(expected).rstrip("/") + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + with pytest.raises(ValueError, match="mutually exclusive"): + get_platform_sdk() -def test_get_async_platform_sdk_uses_uds_endpoint_from_base_url(): - config = PlatformConfig(base_url="unix:///tmp/nemo-platform.sock") # type: ignore[abstract] - with patch("nmp.common.sdk_factory.Configuration.get_platform_config", return_value=config): - sdk = get_async_platform_sdk() +def test_get_platform_sdk_rejects_workload_identity_with_service_headers(monkeypatch: pytest.MonkeyPatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") - assert str(sdk.base_url).rstrip("/") == "http://nemo-platform.local" + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + with pytest.raises(ValueError, match="trusted principal headers"): + get_platform_sdk(as_service="jobs", internal=True) -@pytest.mark.asyncio -async def test_get_async_platform_sdk_workload_identity_reuses_test_http_client( + +def test_get_platform_sdk_uses_workload_identity_with_explicit_sync_http_client( monkeypatch: pytest.MonkeyPatch, tmp_path ): - import nmp.common.sdk_factory as sdk_factory_module + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + captured_requests: list[httpx.Request] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "injected-client-access-token", "expires_in": 300} + + def capture_request(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "data": [], + "pagination": { + "current_page_size": 0, + "page": 1, + "page_size": 0, + "total_pages": 1, + "total_results": 0, + }, + }, + ) + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + monkeypatch.setattr("nemo_platform_plugin.client.oidc.token_exchange_grant", token_exchange_grant) + + with httpx.Client(transport=httpx.MockTransport(capture_request)) as http_client: + sdk = get_platform_sdk(http_client=http_client) + sdk.jobs.list(workspace="default") + + assert sdk._client is http_client + + assert captured_requests[0].headers["Authorization"] == "Bearer injected-client-access-token" + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + + +def test_get_service_scoped_sdk_reuses_base_http_client() -> None: + with httpx.Client() as http_client: + base_sdk = get_platform_sdk(http_client=http_client, base_url="http://nmp.example.test") + service_sdk = get_service_scoped_sdk(base_sdk, "jobs") + assert service_sdk is not base_sdk + assert service_sdk._client is base_sdk._client + assert service_sdk.default_headers["X-NMP-Principal-Id"] == "service:jobs" + assert service_sdk.default_headers["X-NMP-Internal"] == "true" + + +def test_get_service_scoped_sdk_preserves_workload_identity_auth(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + + with httpx.Client() as http_client: + base_sdk = get_platform_sdk(http_client=http_client) + service_sdk = get_service_scoped_sdk(base_sdk, "jobs") + scoped_sdk = base_sdk.with_options(set_default_headers={"X-Test": "true"}) + + assert service_sdk is not base_sdk + assert service_sdk._client is base_sdk._client + assert service_sdk.token_provider is base_sdk.token_provider + assert service_sdk.default_headers["X-NMP-Internal"] == "true" + assert "X-NMP-Principal-Id" not in service_sdk.default_headers + assert scoped_sdk.token_provider is base_sdk.token_provider + + +def test_get_service_scoped_sdk_rejects_workload_identity_on_behalf_of( + monkeypatch: pytest.MonkeyPatch, tmp_path +) -> None: + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + + with httpx.Client() as http_client: + base_sdk = get_platform_sdk(http_client=http_client) + + with pytest.raises(ValueError, match="trusted principal headers"): + get_service_scoped_sdk(base_sdk, "jobs", on_behalf_of="user@example.com") + + +def test_get_sdk_on_behalf_of_rejects_workload_identity_sdk(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + + with httpx.Client() as http_client: + base_sdk = get_platform_sdk(http_client=http_client) - transport = httpx.MockTransport(lambda _request: httpx.Response(200, json={})) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as http_client: - sdk_factory_module._test_http_client = http_client - try: - sdk = get_async_platform_sdk() + with pytest.raises(ValueError, match="trusted principal headers"): + get_sdk_on_behalf_of(base_sdk, "user@example.com") - assert sdk._client is http_client - assert str(sdk.base_url).rstrip("/") == "http://nmp.example.test" - finally: - sdk_factory_module._test_http_client = None + +def test_get_async_platform_sdk(): + """Test get_async_platform_sdk basic functionality (config path: SDK base_url matches platform config).""" + sdk = get_async_platform_sdk() + + assert sdk is not None + assert hasattr(sdk, "base_url") + # Normalize to str: SDK may expose URL object, config may be str; both environments + expected = Configuration.get_platform_config().base_url + assert str(sdk.base_url).rstrip("/") == str(expected).rstrip("/") def test_get_async_platform_sdk_with_service_principal(): @@ -280,6 +431,110 @@ def test_get_async_platform_sdk_internal_flag(): assert sdk.default_headers["X-NMP-Internal"] == "true" +@pytest.mark.asyncio +async def test_get_async_platform_sdk_closes_factory_owned_http_client(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + sdk = get_async_platform_sdk() + http_client = sdk._client + + await sdk.close() + + assert http_client.is_closed + + +@pytest.mark.asyncio +async def test_get_async_platform_sdk_does_not_close_injected_http_client() -> None: + async with httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) as http_client: + sdk = get_async_platform_sdk(base_url="http://nmp.example.test", http_client=http_client) + + await sdk.close() + + assert not http_client.is_closed + + +@pytest.mark.asyncio +async def test_request_scoped_sdk_close_does_not_close_base_http_client() -> None: + base_sdk = get_async_platform_sdk(base_url="http://nmp.example.test") + with patch( + "nmp.common.sdk_factory.get_principal_auth_headers", + return_value={"X-NMP-Principal-Id": "user@example.com"}, + ): + scoped_sdk = get_request_scoped_sdk(base_sdk) + + try: + await scoped_sdk.close() + + assert not base_sdk._client.is_closed + finally: + await base_sdk.close() + + +@pytest.mark.asyncio +async def test_get_async_platform_sdk_rejects_workload_identity_with_service_headers( + monkeypatch: pytest.MonkeyPatch, tmp_path +): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with pytest.raises(ValueError, match="trusted principal headers"): + get_async_platform_sdk(as_service="jobs", internal=True) + + +@pytest.mark.asyncio +async def test_get_async_platform_sdk_uses_workload_identity_with_explicit_async_http_client( + monkeypatch: pytest.MonkeyPatch, tmp_path +): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + captured_requests: list[httpx.Request] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "async-injected-client-access-token", "expires_in": 300} + + def capture_request(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "data": [], + "pagination": { + "current_page_size": 0, + "page": 1, + "page_size": 0, + "total_pages": 1, + "total_results": 0, + }, + }, + ) + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + monkeypatch.setattr("nemo_platform_plugin.client.oidc.token_exchange_grant", token_exchange_grant) + + async with httpx.AsyncClient(transport=httpx.MockTransport(capture_request)) as http_client: + sdk = get_async_platform_sdk(http_client=http_client) + await sdk.jobs.list(workspace="default") + + assert sdk._client is http_client + + assert captured_requests[0].headers["Authorization"] == "Bearer async-injected-client-access-token" + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + + def test_on_behalf_of_without_service_principal(): """Test that on_behalf_of works without as_service (propagates user context).""" # When auth is enabled but no context is set, on_behalf_of should still be added @@ -322,60 +577,62 @@ def test_get_task_sdk_without_principal(monkeypatch: pytest.MonkeyPatch): assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers -def test_get_task_sdk_does_not_inherit_shared_client_authorization(monkeypatch: pytest.MonkeyPatch): - """Service-principal task SDK auth must come from SDK headers, not stale shared-client auth.""" +def test_get_task_sdk_creates_fresh_immutable_sdk_client(monkeypatch: pytest.MonkeyPatch): + """Service-principal SDK auth should be SDK-scoped, not stored on the SDK HTTP client.""" monkeypatch.setenv( "NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"}), ) - client = shared_sync_http_client() - old_headers = dict(client.headers) - client.headers["Authorization"] = "Bearer service:jobs" - sdk = None + + sdk = get_task_sdk(as_service="jobs") + other_sdk = get_task_sdk(as_service="jobs") + scoped_sdk = sdk.with_options(set_default_headers={"X-Test": "true"}) try: - sdk = get_task_sdk(as_service="jobs") - assert sdk._client is not client + assert sdk._client is not other_sdk._client + assert scoped_sdk._client is sdk._client assert sdk.default_headers["X-NMP-Principal-Id"] == "service:jobs" assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "creator@example.com" assert "Authorization" not in sdk.default_headers assert "Authorization" not in sdk._client.headers + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + sdk._client.headers["Authorization"] = "Bearer stale" finally: - if sdk is not None: - sdk.close() - client.headers.clear() - client.headers.update(old_headers) - client.headers.pop("Authorization", None) + sdk.close() + other_sdk.close() + scoped_sdk.close() @pytest.mark.asyncio -async def test_get_async_task_sdk_does_not_inherit_shared_client_authorization(monkeypatch: pytest.MonkeyPatch): - """Async task SDK auth must also avoid stale shared-client auth.""" +async def test_get_async_task_sdk_creates_fresh_immutable_sdk_client(monkeypatch: pytest.MonkeyPatch): + """Async service-principal SDK auth should also stay off the SDK HTTP client.""" monkeypatch.setenv( "NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"}), ) - client = shared_async_http_client() - old_headers = dict(client.headers) - client.headers["Authorization"] = "Bearer service:jobs" - sdk = None + + sdk = get_async_task_sdk(as_service="jobs") + other_sdk = get_async_task_sdk(as_service="jobs") + scoped_sdk = sdk.with_options(set_default_headers={"X-Test": "true"}) try: - sdk = get_async_task_sdk(as_service="jobs") - assert sdk._client is not client + assert sdk._client is not other_sdk._client + assert scoped_sdk._client is sdk._client assert sdk.default_headers["X-NMP-Principal-Id"] == "service:jobs" assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "creator@example.com" assert "Authorization" not in sdk.default_headers assert "Authorization" not in sdk._client.headers + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + sdk._client.headers["Authorization"] = "Bearer stale" finally: - if sdk is not None: - await sdk.close() - client.headers.clear() - client.headers.update(old_headers) - client.headers.pop("Authorization", None) + await sdk.close() + await other_sdk.close() + await scoped_sdk.close() -def test_get_task_sdk_uses_workload_identity_when_token_file_configured(monkeypatch: pytest.MonkeyPatch, tmp_path): +def test_get_task_sdk_uses_workload_identity_when_token_file_configured( + monkeypatch: pytest.MonkeyPatch, tmp_path, caplog: pytest.LogCaptureFixture +): """Task SDKs should centralize the workload-token-vs-service-header choice.""" subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") @@ -390,17 +647,19 @@ def token_exchange_grant(**kwargs): monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) monkeypatch.setattr( - "nemo_platform_ext.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config() + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), ) - monkeypatch.setattr("nemo_platform_ext.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + monkeypatch.setattr("nemo_platform_plugin.client.oidc.token_exchange_grant", token_exchange_grant) + caplog.set_level(logging.WARNING, logger="nmp.common.sdk_factory") sdk = get_task_sdk(as_service="customizer") try: request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") - sdk._client._event_hooks["request"][0](request) + _apply_sync_auth(sdk, request) finally: sdk.close() @@ -409,6 +668,7 @@ def token_exchange_grant(**kwargs): assert "X-NMP-Principal-Id" not in request.headers assert "X-NMP-Principal-On-Behalf-Of" not in request.headers assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + assert "will authenticate as service:customizer without on-behalf-of delegation" not in caplog.text def test_get_task_sdk_uses_explicit_sync_http_client(monkeypatch: pytest.MonkeyPatch): @@ -421,6 +681,16 @@ def test_get_task_sdk_uses_explicit_sync_http_client(monkeypatch: pytest.MonkeyP assert sdk._client is client +@pytest.mark.asyncio +async def test_get_async_task_sdk_uses_explicit_async_http_client(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + async with httpx.AsyncClient() as client: + sdk = get_async_task_sdk(as_service="customizer", http_client=client) + + assert sdk._client is client + + def test_get_request_scoped_sdk_merges_otel_and_auth_headers(): """Test that get_request_scoped_sdk merges OTEL and auth headers.""" base_sdk = get_async_platform_sdk() @@ -448,54 +718,28 @@ def test_get_request_scoped_sdk_merges_otel_and_auth_headers(): assert scoped_sdk.default_headers["X-NMP-Principal-Groups"] == "group1,group2" -def test_get_request_scoped_sdk_preserves_request_router(monkeypatch: pytest.MonkeyPatch): - """Derived request SDKs must keep the base SDK's path-aware platform request router.""" - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_SERVICES", "entities") - monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") - monkeypatch.setenv("NMP_SERVICE_PORT", "8080") - Configuration.clear_cache() - - try: - base_sdk = get_async_platform_sdk() - - with patch("nmp.common.sdk_factory.get_otel_headers", return_value={}): - with patch( - "nmp.common.sdk_factory.get_principal_auth_headers", - return_value={"X-NMP-Principal-Id": "service:models"}, - ): - scoped_sdk = get_request_scoped_sdk(base_sdk) - - prepared = scoped_sdk._prepare_url("https://nemo-gateway:8080/apis/entities/v2/workspaces") - - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - finally: - Configuration.clear_cache() +def test_get_request_scoped_sdk_reuses_base_sdk_http_client(): + """Derived request SDKs must keep the base SDK's lifecycle-owned HTTP client.""" + base_sdk = get_async_platform_sdk() + with patch("nmp.common.sdk_factory.get_otel_headers", return_value={}): + with patch( + "nmp.common.sdk_factory.get_principal_auth_headers", + return_value={"X-NMP-Principal-Id": "service:models"}, + ): + scoped_sdk = get_request_scoped_sdk(base_sdk) -def test_get_sdk_on_behalf_of_preserves_request_router(monkeypatch: pytest.MonkeyPatch): - """SDKs derived with on-behalf-of headers must still keep platform request routing.""" - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_SERVICES", "entities") - monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") - monkeypatch.setenv("NMP_SERVICE_PORT", "8080") - Configuration.clear_cache() + assert scoped_sdk is not base_sdk + assert scoped_sdk._client is base_sdk._client - try: - base_sdk = get_async_platform_sdk(as_service="models", internal=True) - scoped_sdk = get_sdk_on_behalf_of(base_sdk, "user@example.com") - prepared = scoped_sdk._prepare_url("https://nemo-gateway:8080/apis/entities/v2/workspaces") +def test_get_sdk_on_behalf_of_reuses_base_sdk_http_client(): + """SDKs derived with on-behalf-of headers must keep the base SDK's HTTP client.""" + base_sdk = get_async_platform_sdk(as_service="models", internal=True) + scoped_sdk = get_sdk_on_behalf_of(base_sdk, "user@example.com") - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - finally: - Configuration.clear_cache() + assert scoped_sdk is not base_sdk + assert scoped_sdk._client is base_sdk._client def test_get_request_scoped_sdk_returns_base_sdk_when_no_headers(): @@ -662,221 +906,6 @@ def test_get_request_scoped_sdk_service_principal_with_on_behalf_of(): assert scoped_sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "user@example.com" -# --- Dynamic routing (service discovery map) tests --- - - -@pytest.fixture -def platform_config_with_service_discovery(): - """Platform config with service_discovery map for entities and jobs.""" - return PlatformConfig( # type: ignore[abstract] - base_url="http://platform:8080", - service_discovery={ - "entities": "http://entities-service:8080", - "jobs": "http://jobs-service:8080", - }, - ) - - -def test_resolve_platform_request_url_routes_api_path_to_service_url(platform_config_with_service_discovery): - """The named request router policy owns per-service routing.""" - - def default_resolver(url: str) -> httpx.URL: - if url.startswith("/"): - return httpx.URL(f"http://platform:8080{url}") - return httpx.URL(url) - - prepared = resolve_platform_request_url( - "/apis/entities/v2/workspaces?limit=10", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - - assert prepared.scheme == "http" - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - assert prepared.query == b"limit=10" - - -def test_resolve_platform_request_url_logs_path_without_raw_url( - caplog: pytest.LogCaptureFixture, - platform_config_with_service_discovery, -): - """Routing logs expose the resolved path without query parameters.""" - - def default_resolver(url: str) -> httpx.URL: - if url.startswith("/"): - return httpx.URL(f"http://platform:8080{url}") - return httpx.URL(url) - - caplog.set_level(logging.DEBUG, logger="nmp.common.sdk_factory") - - resolve_platform_request_url( - "/health/ready?token=secret", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - resolve_platform_request_url( - "/apis/entities/v2/workspaces?token=secret", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - - original_record = next(record for record in caplog.records if record.message == "Routing URL to original URL") - service_record = next(record for record in caplog.records if record.message == "Routing URL to service URL") - - assert not hasattr(original_record, "url") - assert original_record.service == "unknown" - assert original_record.path == "/health/ready" - assert original_record.host == "platform" - assert original_record.port == 8080 - - assert not hasattr(service_record, "url") - assert service_record.service == "entities" - assert service_record.path == "/apis/entities/v2/workspaces" - assert service_record.host == "entities-service" - assert service_record.port == 8080 - - for record in (original_record, service_record): - assert "token=secret" not in str(record.__dict__) - - -def test_platform_request_router_uses_default_resolver_for_non_api_paths(platform_config_with_service_discovery): - """Non-API paths follow the SDK's normal URL preparation.""" - router = PlatformRequestRouter( - platform_config=platform_config_with_service_discovery, - default_resolver=lambda url: httpx.URL(f"http://platform:8080{url}"), - ) - - prepared = router.resolve("/health/ready") - - assert str(prepared) == "http://platform:8080/health/ready" - - -def test_get_platform_sdk_routes_entities_path_to_entities_service( - platform_config_with_service_discovery, -): - """Routes /apis/entities/v2/workspaces to the entities service URL.""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/entities/v2/workspaces" in str(prepared.path) - - -def test_get_platform_sdk_routes_service_path_to_env_override( - monkeypatch: pytest.MonkeyPatch, - platform_config_with_service_discovery, -): - monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-env:9090") - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-env" - assert prepared.port == 9090 - assert prepared.scheme == "http" - - -def test_get_platform_sdk_routes_jobs_path_to_jobs_service( - platform_config_with_service_discovery, -): - """Routes /apis/jobs/v2/workspaces/jobs to the jobs service URL.""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/jobs/v2/workspaces/jobs" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "jobs-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/jobs/v2/workspaces/jobs" in str(prepared.path) - - -def test_get_platform_sdk_routing_fallback_to_base_url_when_no_match( - platform_config_with_service_discovery, -): - """When the path does not match /apis/{service-name}/ (lowercase+dashes), use the original URL (base).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - # Path that does not match /apis/{service-name}/ (e.g. /api/ singular, or no such prefix) - request_url = "http://platform:8080/api/other/v1/thing" - prepared = sdk._prepare_url(request_url) - - # Should pass through to original behavior: same host as request - assert prepared.host == "platform" - assert prepared.port == 8080 - - -def test_get_async_platform_sdk_routes_entities_path_to_entities_service( - platform_config_with_service_discovery, -): - """Routes /apis/entities/v2/workspaces to the entities service URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/entities/v2/workspaces" in str(prepared.path) - - -def test_get_async_platform_sdk_routes_jobs_path_to_jobs_service( - platform_config_with_service_discovery, -): - """Routes /apis/jobs/v2/workspaces/jobs to the jobs service URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/apis/jobs/v2/workspaces/jobs" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "jobs-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/jobs/v2/workspaces/jobs" in str(prepared.path) - - -def test_get_async_platform_sdk_routing_fallback_to_base_url_when_no_match( - platform_config_with_service_discovery, -): - """When the path does not match /apis/{service-name}/, use the original URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/api/other/v1/thing" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "platform" - assert prepared.port == 8080 - - # --- get_entity_parts tests --- diff --git a/packages/nmp_common/tests/test_immutable_http_client.py b/packages/nmp_common/tests/test_immutable_http_client.py new file mode 100644 index 0000000000..143d2365cf --- /dev/null +++ b/packages/nmp_common/tests/test_immutable_http_client.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import httpx +import pytest +from nmp.common.immutable_http_client import ImmutableHttpClientMixin + + +def _noop_request_hook(request: httpx.Request) -> None: + pass + + +class _FrozenClient(ImmutableHttpClientMixin, httpx.Client): + def __init__(self) -> None: + super().__init__( + headers={"X-Initial": "true"}, + event_hooks={"request": [_noop_request_hook]}, + ) + self._freeze_http_client() + + +def test_immutable_sdk_client_still_builds_requests() -> None: + with _FrozenClient() as client: + request = client.build_request("GET", "http://nmp.example.test/health") + + assert request.url == "http://nmp.example.test/health" + assert request.headers["X-Initial"] == "true" + + +def test_immutable_sdk_client_blocks_client_configuration_assignment() -> None: + with _FrozenClient() as client: + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.headers = httpx.Headers({"Authorization": "Bearer stale"}) + + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.base_url = httpx.URL("http://other.example.test") + + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.params = {"debug": "true"} + + +def test_immutable_sdk_client_blocks_header_mutation() -> None: + with _FrozenClient() as client: + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers["Authorization"] = "Bearer stale" + + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers.update({"Authorization": "Bearer stale"}) + + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers.pop("X-Initial") + + +def test_immutable_sdk_client_blocks_cookie_mutation_and_ignores_response_cookies() -> None: + with _FrozenClient() as client: + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.cookies["session"] = "stale" + + request = client.build_request("GET", "http://nmp.example.test/health") + response = httpx.Response(200, headers={"Set-Cookie": "session=stale"}, request=request) + + client.cookies.extract_cookies(response) + + assert "session" not in client.cookies + + +def test_immutable_sdk_client_blocks_event_hook_mutation() -> None: + with _FrozenClient() as client: + with pytest.raises(AttributeError): + client.event_hooks["request"].append(_noop_request_hook) + + with pytest.raises(TypeError): + client.event_hooks["request"] = [_noop_request_hook] diff --git a/packages/nmp_common/tests/test_platform_endpoint.py b/packages/nmp_common/tests/test_platform_endpoint.py index 2b4911b53e..38e0ded177 100644 --- a/packages/nmp_common/tests/test_platform_endpoint.py +++ b/packages/nmp_common/tests/test_platform_endpoint.py @@ -1,12 +1,32 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import logging +import os from pathlib import Path from unittest.mock import patch +import httpx import pytest -from nmp.common.config import PlatformConfig -from nmp.common.platform_endpoint import UDS_BASE_URL, parse_platform_endpoint, resolve_service_endpoint +from nmp.common.config import PlatformConfig, get_common_service_config +from nmp.common.platform_endpoint import ( + UDS_BASE_URL, + PlatformEndpoint, + _AsyncPlatformEndpointRoutingTransport, + _SyncPlatformEndpointRoutingTransport, + parse_platform_endpoint, + resolve_platform_endpoint, + resolve_service_endpoint, +) + + +@pytest.fixture(autouse=True) +def clear_service_url_env_vars(monkeypatch: pytest.MonkeyPatch) -> None: + for key in tuple(os.environ): + if key == "NMP_BASE_URL": + continue + if key.startswith("NMP_") and key.endswith("_URL"): + monkeypatch.delenv(key, raising=False) def test_parse_tcp_endpoint() -> None: @@ -80,6 +100,199 @@ def test_endpoint_env_family_is_not_part_of_contract(monkeypatch: pytest.MonkeyP assert endpoint.connect_base_url == "http://platform:8080" +@pytest.fixture +def platform_config_with_service_discovery() -> PlatformConfig: + return PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={ + "entities": "http://entities-service:8080", + "jobs": "http://jobs-service:8080", + }, + ) + + +def test_resolve_platform_endpoint_carries_service_routes( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + assert endpoint.connect_base_url == "http://platform:8080" + assert endpoint.service_endpoints["entities"].connect_base_url == "http://entities-service:8080" + assert endpoint.service_endpoints["jobs"].connect_base_url == "http://jobs-service:8080" + + +def test_resolve_platform_endpoint_rejects_malformed_service_route() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={ + "entities": "http://entities-service:8080", + "jobs": "not-a-url", + }, + ) + + with pytest.raises(ValueError, match="Unsupported platform endpoint URL 'not-a-url'"): + resolve_platform_endpoint(config) + + +def test_resolve_platform_endpoint_rejects_malformed_base_url() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="not-a-url", + service_discovery={"entities": "http://entities-service:8080"}, + ) + + with pytest.raises(ValueError, match="Unsupported platform endpoint URL 'not-a-url'"): + resolve_platform_endpoint(config) + + +def test_platform_endpoint_routes_api_path_to_service_url( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?limit=10") + + assert routed.endpoint.connect_base_url == "http://entities-service:8080" + assert routed.url.scheme == "http" + assert routed.url.host == "entities-service" + assert routed.url.port == 8080 + assert routed.url.path == "/apis/entities/v2/workspaces" + assert routed.url.query == b"limit=10" + + +def test_platform_endpoint_preserves_service_url_path_prefix() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={"entities": "http://entities-service:8080/entities-prefix"}, + ) + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?limit=10") + + assert routed.endpoint.connect_base_url == "http://entities-service:8080/entities-prefix" + assert routed.url.scheme == "http" + assert routed.url.host == "entities-service" + assert routed.url.port == 8080 + assert routed.url.path == "/entities-prefix/apis/entities/v2/workspaces" + assert routed.url.query == b"limit=10" + + +def test_platform_endpoint_uses_env_override( + monkeypatch: pytest.MonkeyPatch, + platform_config_with_service_discovery: PlatformConfig, +) -> None: + monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-env:9090") + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces") + + assert routed.endpoint.connect_base_url == "http://entities-env:9090" + assert routed.url.host == "entities-env" + assert routed.url.port == 9090 + + +def test_platform_endpoint_uses_matching_env_route_without_mutating_config(monkeypatch: pytest.MonkeyPatch) -> None: + config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-env:9090") + + endpoint = resolve_platform_endpoint(config) + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces") + + assert "entities" not in endpoint.service_endpoints + assert "entities" not in config.service_discovery + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert routed.url.host == "platform" + + +def test_platform_endpoint_keeps_configured_local_service_on_local_url(monkeypatch: pytest.MonkeyPatch) -> None: + config = PlatformConfig(base_url="http://platform:8080", services="hello-world") # type: ignore[abstract] + monkeypatch.setenv("NMP_HELLO_WORLD_URL", "http://hello-world-service:8080") + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/hello-world/v2/workspaces/default/hello") + + assert routed.endpoint.connect_base_url == get_common_service_config().get_host_url() + assert routed.url.host == "127.0.0.1" + + +def test_platform_endpoint_ignores_unknown_service_url_env_var(monkeypatch: pytest.MonkeyPatch) -> None: + config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + monkeypatch.setenv("NMP_NOT_A_SERVICE_URL", "http://not-a-service:8080") + + endpoint = resolve_platform_endpoint(config) + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces/default") + + assert "not-a-service" not in endpoint.service_endpoints + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert routed.url.host == "platform" + + +def test_platform_endpoint_routes_uds_service() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={"entities": "unix:///tmp/entities.sock"}, + ) + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces") + + assert routed.endpoint.transport == "uds" + assert routed.endpoint.socket_path == Path("/tmp/entities.sock") + assert str(routed.url) == "http://nemo-platform.local/apis/entities/v2/workspaces" + + +def test_platform_endpoint_keeps_non_api_path_on_default_endpoint( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/health/ready") + + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert str(routed.url) == "http://platform:8080/health/ready" + + +def test_platform_endpoint_keeps_non_api_service_url( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://entities-service:8080/status") + + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert str(routed.url) == "http://entities-service:8080/status" + + +def test_platform_endpoint_logs_path_without_raw_url( + caplog: pytest.LogCaptureFixture, + platform_config_with_service_discovery: PlatformConfig, +) -> None: + caplog.set_level(logging.DEBUG, logger="nmp.common.platform_endpoint") + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + endpoint.route_request_url("http://platform:8080/health/ready?token=secret") + endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?token=secret") + + default_record = next( + record for record in caplog.records if record.message == "Routing SDK URL to default endpoint" + ) + service_record = next( + record for record in caplog.records if record.message == "Routing SDK URL to service endpoint" + ) + + assert not hasattr(default_record, "url") + assert getattr(default_record, "service") == "unknown" + assert getattr(default_record, "path") == "/health/ready" + + assert not hasattr(service_record, "url") + assert getattr(service_record, "service") == "entities" + assert getattr(service_record, "path") == "/apis/entities/v2/workspaces" + assert getattr(service_record, "host") == "entities-service" + assert getattr(service_record, "port") == 8080 + + for record in (default_record, service_record): + assert "token=secret" not in str(record.__dict__) + + def test_sync_http_client_omits_unset_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") @@ -101,7 +314,7 @@ def test_sync_http_client_passes_explicit_timeout() -> None: def test_sync_sdk_http_client_uses_sdk_default_for_tcp() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as client: endpoint.sync_sdk_http_client() client.assert_called_once_with() @@ -110,7 +323,7 @@ def test_sync_sdk_http_client_uses_sdk_default_for_tcp() -> None: def test_sync_sdk_http_client_passes_explicit_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as client: endpoint.sync_sdk_http_client(timeout=2.0) client.assert_called_once_with(timeout=2.0) @@ -132,8 +345,8 @@ def test_uds_sync_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> N endpoint = parse_platform_endpoint("unix:///tmp/nemo-platform.sock") with ( - patch("nmp.common.platform_endpoint.DefaultHttpxClient") as sdk_client, - patch("nmp.common.platform_endpoint.httpx.Client") as client, + patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as sdk_client, + patch("nmp.common.platform_endpoint.ImmutableHttpxClient") as client, ): endpoint.sync_sdk_http_client() @@ -144,6 +357,102 @@ def test_uds_sync_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> N assert "timeout" not in kwargs +def test_sync_routing_transport_prebuilds_uds_service_transports() -> None: + service_endpoint = parse_platform_endpoint("unix:///tmp/entities.sock") + constructed_transports: list[object] = [] + closed_transports: list[object] = [] + + class FakeHTTPTransport: + def __init__(self, *, uds: str | None = None) -> None: + self.uds = uds + if uds is not None: + constructed_transports.append(self) + + def close(self) -> None: + if self.uds is not None: + closed_transports.append(self) + + with patch("nmp.common.platform_endpoint.httpx.HTTPTransport", FakeHTTPTransport): + routing_transport = _SyncPlatformEndpointRoutingTransport( + endpoint=PlatformEndpoint( + connect_base_url="http://platform:8080", + socket_path=None, + transport="tcp", + service_endpoints={"entities": service_endpoint}, + ) + ) + returned_transport = routing_transport._transport_for_endpoint(service_endpoint) + + routing_transport.close() + + assert len(constructed_transports) == 1 + assert returned_transport is constructed_transports[0] + assert closed_transports == constructed_transports + + +def test_sync_routing_transport_preserves_non_api_service_status_url( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + captured_urls: list[str] = [] + + class FakeHTTPTransport: + def __init__(self, **kwargs) -> None: + pass + + def handle_request(self, request) -> httpx.Response: + captured_urls.append(str(request.url)) + return httpx.Response(200, request=request) + + def close(self) -> None: + pass + + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + with patch("nmp.common.platform_endpoint.httpx.HTTPTransport", FakeHTTPTransport): + routing_transport = _SyncPlatformEndpointRoutingTransport(endpoint=endpoint) + request = httpx.Request("GET", "http://entities-service:8080/status") + routing_transport.handle_request(request) + + assert captured_urls == ["http://entities-service:8080/status"] + + +def test_routing_transports_pass_custom_ca_to_tcp_transport(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + cert_file = tmp_path / "ca.pem" + cert_file.write_text("certificate", encoding="utf-8") + monkeypatch.setenv("NMP_CLIENT_SSL_CERT_FILE", str(cert_file)) + sync_kwargs: list[dict] = [] + async_kwargs: list[dict] = [] + + class FakeHTTPTransport: + def __init__(self, **kwargs) -> None: + sync_kwargs.append(kwargs) + + def close(self) -> None: + pass + + class FakeAsyncHTTPTransport: + def __init__(self, **kwargs) -> None: + async_kwargs.append(kwargs) + + async def aclose(self) -> None: + pass + + endpoint = resolve_platform_endpoint( + PlatformConfig( + base_url="https://platform:8443", + service_discovery={"entities": "https://entities-service:9443"}, + ) + ) + with ( + patch("nmp.common.platform_endpoint.httpx.HTTPTransport", FakeHTTPTransport), + patch("nmp.common.platform_endpoint.httpx.AsyncHTTPTransport", FakeAsyncHTTPTransport), + ): + _SyncPlatformEndpointRoutingTransport(endpoint=endpoint) + _AsyncPlatformEndpointRoutingTransport(endpoint=endpoint) + + assert sync_kwargs[0] == {"verify": str(cert_file)} + assert async_kwargs[0] == {"verify": str(cert_file)} + + def test_async_http_client_omits_unset_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") @@ -165,7 +474,7 @@ def test_async_http_client_passes_explicit_timeout() -> None: def test_async_sdk_http_client_uses_sdk_default_for_tcp() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as client: endpoint.async_sdk_http_client() client.assert_called_once_with() @@ -174,7 +483,7 @@ def test_async_sdk_http_client_uses_sdk_default_for_tcp() -> None: def test_async_sdk_http_client_passes_explicit_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as client: endpoint.async_sdk_http_client(timeout=2.0) client.assert_called_once_with(timeout=2.0) @@ -196,8 +505,8 @@ def test_uds_async_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> endpoint = parse_platform_endpoint("unix:///tmp/nemo-platform.sock") with ( - patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as sdk_client, - patch("nmp.common.platform_endpoint.httpx.AsyncClient") as client, + patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as sdk_client, + patch("nmp.common.platform_endpoint.ImmutableAsyncHttpxClient") as client, ): endpoint.async_sdk_http_client() diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py b/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py index ed6ea42c37..e9a01e32d0 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py @@ -9,13 +9,13 @@ import logging import threading from collections.abc import Callable -from typing import cast from nmp.common.service import Service from nmp.common.service.deptree import resolve_service_loading_order logger = logging.getLogger(__name__) + ControllerRunFunc = Callable[[threading.Event], object] @@ -66,4 +66,4 @@ def load_controller_run_func(controller_name: str, import_path: str) -> Controll raise TypeError(f"Controller {controller_name} must be a callable, got {type(run_func)}") logger.debug("Loaded controller %s", controller_name) - return cast(ControllerRunFunc, run_func) + return run_func diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py b/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py index 1d074b64cb..e004a9e3b7 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py @@ -11,7 +11,8 @@ Wraps a :class:`~nemo_platform_plugin.controller.NemoController` (async) as a platform-native :class:`~nmp.common.controller.Controller` (sync / thread-based). Use :func:`make_controller_run_func` to create the - ``run(stop_signal)`` callable expected by :func:`~nmp.platform_runner.server.create_app`. + ``run(stop_signal)`` callable expected by + :func:`~nmp.platform_runner.server.create_app`. """ from __future__ import annotations @@ -166,7 +167,6 @@ def make_controller_run_func(controller_cls: type[NemoController]) -> Callable[[ The returned callable matches the signature expected by :func:`~nmp.platform_runner.server.create_app` (same as core controller run functions registered in ``AVAILABLE_CONTROLLERS``). - Lifecycle inside the returned function: 1. Instantiate *controller_cls*. diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/run.py b/packages/nmp_platform_runner/src/nmp/platform_runner/run.py index 98da70493e..492bc97dd2 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/run.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/run.py @@ -11,7 +11,6 @@ import threading import time from collections.abc import Callable, Mapping -from typing import cast from nmp.common.config import get_auth_config, get_common_service_config, get_platform_config, get_service_config from nmp.common.observability import initialize_obs, setup_global_instrumentations @@ -69,7 +68,12 @@ def run_controllers_in_threads( """Start controller run functions in daemon threads.""" threads = [] for name, run_func in controller_run_funcs.items(): - thread = threading.Thread(target=run_func, args=(stop_signal,), name=f"controller-{name}", daemon=True) + thread = threading.Thread( + target=run_func, + args=(stop_signal,), + name=f"controller-{name}", + daemon=True, + ) thread.start() threads.append(thread) return threads @@ -222,10 +226,10 @@ def _load_run_functions( t0 = time.perf_counter() value = registry[name] try: - if callable(value): - run_funcs[name] = cast(ControllerRunFunc, value) - else: + if isinstance(value, str): run_funcs[name] = load_controller_run_func(name, value) + else: + run_funcs[name] = value except (ImportError, TypeError, AttributeError, ValueError) as error: logger.error("Failed to load %s %s: %s", kind, name, error) raise SystemExit(1) from error diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/server.py b/packages/nmp_platform_runner/src/nmp/platform_runner/server.py index a2494a7241..6994d9e940 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/server.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/server.py @@ -10,9 +10,8 @@ import logging import os import threading -from collections.abc import Callable, Mapping, MutableMapping +from collections.abc import Mapping, MutableMapping from contextlib import asynccontextmanager -from typing import cast import httpx import uvicorn @@ -22,7 +21,6 @@ from nmp.common.api.utils import install_query_param_schema_openapi_hook from nmp.common.auth import AuthorizationMiddleware from nmp.common.config import get_auth_config, get_platform_config -from nmp.common.http_clients import close_shared_http_clients from nmp.common.observability import initialize_obs, setup_fastapi_instrumentations, setup_global_instrumentations from nmp.common.observability.context import create_app_context_dependency from nmp.common.pyleak import detect_blocking @@ -160,11 +158,17 @@ def create_app( ) ) + if http_client is not None: + for service_instance in services: + service_instance.dependency_provider.configure_http_client(http_client) + @asynccontextmanager async def lifespan(app: FastAPI): logger.info("Starting Nemo Platform server") controller_threads = [] platform_seed_task: asyncio.Task[None] | None = None + for service_instance in services: + service_instance.dependency_provider.initialize() if controller_run_funcs: logger.info("Starting controllers in lifespan: %s", list(controller_run_funcs)) for name, run_func in controller_run_funcs.items(): @@ -217,8 +221,9 @@ async def run_platform_seed_and_update_readiness() -> None: for thread in controller_threads: thread.join(timeout=5) - await close_shared_http_clients() logger.info("Shutting down Nemo Platform API server") + for service_instance in reversed(services): + await service_instance.on_shutdown() app = FastAPI( title="Nemo Platform API", @@ -289,9 +294,9 @@ async def root_handler() -> Response: def _load_run_functions( names: list[str], - registry: Mapping[str, str | Callable[[threading.Event], object]], -) -> dict[str, Callable[[threading.Event], object]]: - run_funcs: dict[str, Callable[[threading.Event], object]] = {} + registry: Mapping[str, str | ControllerRunFunc], +) -> dict[str, ControllerRunFunc]: + run_funcs: dict[str, ControllerRunFunc] = {} for name in names: value = registry[name] if isinstance(value, str): @@ -485,9 +490,9 @@ def create_default_app() -> FastAPI: "Unknown controller %r requested via NMP_CONTROLLERS=%r. Available controllers: %s" % (controller_name, controller_names_env, available) ) - if callable(controller_value): - controller_run_funcs[controller_name] = cast(ControllerRunFunc, controller_value) - else: + if isinstance(controller_value, str): controller_run_funcs[controller_name] = load_controller_run_func(controller_name, controller_value) + else: + controller_run_funcs[controller_name] = controller_value return create_app(services, controller_run_funcs=controller_run_funcs) diff --git a/packages/nmp_platform_runner/tests/test_run.py b/packages/nmp_platform_runner/tests/test_run.py index 787f7dfcf3..2496320bcc 100644 --- a/packages/nmp_platform_runner/tests/test_run.py +++ b/packages/nmp_platform_runner/tests/test_run.py @@ -74,7 +74,7 @@ def test_run_platform_marks_loaded_services_local_before_starting_controllers(mo monkeypatch.setattr( runner, "_load_run_functions", - lambda names, registry, kind: {"jobs": lambda stop_signal: None} if kind == "controller" else {}, + lambda names, registry, kind: {"jobs": lambda _stop_signal: None} if kind == "controller" else {}, ) monkeypatch.setattr(runner, "_display_banner", lambda **_: None) monkeypatch.setattr( diff --git a/packages/nmp_platform_runner/tests/test_server.py b/packages/nmp_platform_runner/tests/test_server.py index 4696194cf4..849314e1dd 100644 --- a/packages/nmp_platform_runner/tests/test_server.py +++ b/packages/nmp_platform_runner/tests/test_server.py @@ -11,11 +11,13 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +import httpx import pytest from fastapi import APIRouter, FastAPI from fastapi.testclient import TestClient from nemo_platform_plugin.jobs.openapi_utils import clear_query_param_schemas, generate_openapi_extra_params from nmp.common.config import AuthConfig, Configuration +from nmp.common.config import get_platform_config as load_platform_config from nmp.common.config.base import OIDCConfig from nmp.common.service import RouterConfig, Service from nmp.platform_runner import config as runner_config @@ -64,6 +66,16 @@ def get_routers(self): return [] +class ShutdownTrackingService(PluginService): + def __init__(self): + super().__init__() + self.shutdown_called = False + + async def on_shutdown(self) -> None: + self.shutdown_called = True + await super().on_shutdown() + + class _DateFilter(BaseModel): gte: str | None = None @@ -240,13 +252,66 @@ def test_create_app_uses_separate_http_clients_for_auth_callouts(monkeypatch): assert auth_middleware.kwargs["access_key_lifecycle_http_client"] is lifecycle_http_client -def test_create_app_mounted_services_drive_sdk_local_routing_without_services_env(monkeypatch): +@pytest.mark.asyncio +async def test_create_app_configures_service_dependency_provider_http_client(monkeypatch): + _patch_platform_app_config(monkeypatch, seed_on_startup=False) + service = PluginService() + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + + try: + app = server.create_app([service], http_client=http_client) + with TestClient(app): + assert service.dependency_provider.get_http_client() is http_client + finally: + await http_client.aclose() + + +@pytest.mark.asyncio +async def test_create_app_runs_service_shutdown_without_closing_injected_http_client(monkeypatch): + _patch_platform_app_config(monkeypatch, seed_on_startup=False) + service = ShutdownTrackingService() + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + + try: + app = server.create_app([service], http_client=http_client) + with TestClient(app): + pass + + assert service.shutdown_called is True + assert http_client.is_closed is False + finally: + await http_client.aclose() + + +@pytest.mark.asyncio +async def test_create_app_starts_controller_with_stop_signal(monkeypatch): + _patch_platform_app_config(monkeypatch, seed_on_startup=False) + captured: dict[str, object] = {} + started = threading.Event() + + def run_controller(stop_signal: threading.Event) -> None: + captured["stop_signal"] = stop_signal + started.set() + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + app = server.create_app(controller_run_funcs={"test": run_controller}, http_client=http_client) + + try: + with TestClient(app): + assert started.wait(timeout=2) + finally: + await http_client.aclose() + + assert isinstance(captured["stop_signal"], threading.Event) + + +def test_create_app_mounted_services_drive_lifecycle_routing_without_services_env(monkeypatch): monkeypatch.delenv("NMP_SERVICES", raising=False) monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") monkeypatch.setenv("NMP_SERVICE_PORT", "8080") Configuration.clear_cache() - platform_cfg = Configuration.get_platform_config() + platform_cfg = load_platform_config() try: auth_cfg = _make_auth_config(enabled=False) @@ -254,19 +319,19 @@ def test_create_app_mounted_services_drive_sdk_local_routing_without_services_en monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) import nmp.common.auth.middleware as auth_middleware - from nmp.common.sdk_factory import get_platform_sdk + from nmp.common.platform_endpoint import resolve_platform_endpoint monkeypatch.setattr(auth_middleware, "get_auth_config", lambda: auth_cfg) server.create_app(services=[PluginService()]) - sdk = get_platform_sdk() - prepared = sdk._prepare_url("https://nemo-gateway:8080/apis/agents/v2/example") + endpoint = resolve_platform_endpoint(platform_cfg) + routed = endpoint.route_request_url("https://nemo-gateway:8080/apis/agents/v2/example") assert platform_cfg.services == "agents" - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 + assert routed.url.scheme == "http" + assert routed.url.host == "127.0.0.1" + assert routed.url.port == 8080 finally: Configuration.clear_cache() @@ -511,6 +576,7 @@ def test_create_default_app_raises_for_unknown_controller_from_env(monkeypatch): def _make_platform_config_mock(*, redirect_root_to_studio: bool = True) -> MagicMock: cfg = MagicMock() + cfg.base_url = "http://platform.local" cfg.seed_on_startup = False cfg.redirect_root_to_studio = redirect_root_to_studio return cfg diff --git a/packages/nmp_platform_runner/tests/test_sidecars.py b/packages/nmp_platform_runner/tests/test_sidecars.py index e3950b3a93..4dcfddc54c 100644 --- a/packages/nmp_platform_runner/tests/test_sidecars.py +++ b/packages/nmp_platform_runner/tests/test_sidecars.py @@ -105,6 +105,7 @@ def test_create_app_starts_and_stops_dummy_sidecar_with_lifespan() -> None: patch("nmp.platform_runner.server.get_auth_config", return_value=_auth_config(False)), patch("nmp.common.auth.middleware.get_auth_config", return_value=_auth_config(False)), ): + platform_config.return_value.base_url = "http://platform.local" platform_config.return_value.seed_on_startup = False platform_config.return_value.redirect_root_to_studio = False app = server.create_app( @@ -130,6 +131,7 @@ def test_build_platform_app_loads_dependent_sidecar_into_lifespan(monkeypatch: p patch("nmp.platform_runner.server.get_auth_config", return_value=_auth_config(False)), patch("nmp.common.auth.middleware.get_auth_config", return_value=_auth_config(False)), ): + platform_config.return_value.base_url = "http://platform.local" platform_config.return_value.seed_on_startup = False platform_config.return_value.redirect_root_to_studio = False app = server.build_platform_app(runner_config.PlatformAppConfig(services=["models"], controllers=[]), env={}) @@ -187,14 +189,18 @@ def join(self) -> None: monkeypatch.delenv("VLLM_ENDPOINT", raising=False) monkeypatch.setattr(adapters_main, "get_platform_config", lambda: MagicMock(base_url="http://platform.local")) - monkeypatch.setattr(adapters_main, "get_platform_sdk", lambda **_kwargs: MagicMock()) monkeypatch.setattr(adapters_main.asyncio, "new_event_loop", lambda: MagicMock()) monkeypatch.setattr(adapters_main, "Loop", FakeLoop) monkeypatch.setattr(adapters_main, "TimedLoopWaiter", lambda *_args, **_kwargs: object()) monkeypatch.setattr(adapters_main.ControllerManager, "get_instance", classmethod(lambda _cls: manager)) + monkeypatch.setattr(adapters_main, "get_platform_sdk", MagicMock(return_value=MagicMock())) stop_signal = threading.Event() - thread = threading.Thread(target=adapters_main.run, args=(stop_signal,), daemon=True) + thread = threading.Thread( + target=adapters_main.run, + args=(stop_signal,), + daemon=True, + ) try: thread.start() diff --git a/packages/nmp_testing/src/nmp/testing/__init__.py b/packages/nmp_testing/src/nmp/testing/__init__.py index 551e817d37..3cd11a61d9 100644 --- a/packages/nmp_testing/src/nmp/testing/__init__.py +++ b/packages/nmp_testing/src/nmp/testing/__init__.py @@ -15,6 +15,7 @@ - create_test_client: Helper for creating FastAPI test clients with in-memory storage - ClientContext: Container for all client types returned by create_test_client - TEST_USER_EMAIL, TEST_ADMIN_EMAIL: Constants for test principals +- routed_mock_client / routed_async_mock_client: Mock httpx clients that apply platform endpoint routing - subprocess_job_executor_patch: Opt into cpu/default to subprocess/default translation Utilities: @@ -60,6 +61,7 @@ ensure_mock_sidecar_image, get_worker_port_range, ) +from .http_clients import RoutedMockTransport, routed_async_mock_client, routed_mock_client, routed_mock_transport from .jobs import subprocess_job_executor_patch from .notebooks import ( cleanup_temp_venv_and_kernel, @@ -96,6 +98,10 @@ "SDKTestClientAdapter", "TEST_USER_EMAIL", "TEST_ADMIN_EMAIL", + "RoutedMockTransport", + "routed_mock_transport", + "routed_mock_client", + "routed_async_mock_client", "subprocess_job_executor_patch", # Utilities "short_unique_name", diff --git a/packages/nmp_testing/src/nmp/testing/client.py b/packages/nmp_testing/src/nmp/testing/client.py index 1b1b0f1b49..97139a3afb 100644 --- a/packages/nmp_testing/src/nmp/testing/client.py +++ b/packages/nmp_testing/src/nmp/testing/client.py @@ -445,12 +445,6 @@ def _add_service( services_to_start = [_create_svc(svc, configs) for svc in services_to_create] services_to_start = order_services_by_dependencies(services_to_start) - # Clear any stale SDK client from previous tests BEFORE creating app. - # This prevents service startup code from using a previous test's http transport. - from nmp.common import sdk_factory as sdk_factory_module - - sdk_factory_module._test_http_client = None - # Create transport and http_client BEFORE the app, so we can inject the client # into create_app() for middleware (AuthorizationMiddleware). We set transport.app # after app creation - this works because no requests are made until setup completes. @@ -484,27 +478,11 @@ async def _pending_asgi_app(scope: Scope, receive: Receive, send: Send) -> None: # Store on app.state so tests can access it via test_client.app.state.access_log app.state.access_log = access_log_instance - # Configure module-level http client as FALLBACK for direct callers of - # get_async_platform_sdk()/get_platform_sdk() that don't use DependencyProvider. - # The primary injection path is through DependencyProvider (see below). These - # module-level variables will be removed once all direct callers are migrated. - # See architecture/docs/http-client-injection.md for details. - sdk_factory_module._test_http_client = async_http_client - stack.callback(lambda: setattr(sdk_factory_module, "_test_http_client", None)) - async_sdk = AsyncNeMoPlatform(base_url="http://testserver", http_client=async_http_client, workspace=workspace) # Create the EntityClient (used for DI and optionally yielded) entity_client = EntityClient(client_from_platform(async_sdk, AsyncEntitiesClient)) - # Inject ASGI-transport clients into each service's DependencyProvider. - # This is critical for services that call dependency_provider.get_sdk_client() - # directly (e.g., in on_startup for background tasks like auth policy refresh). - # See architecture/docs/http-client-injection.md for details. - for svc in services_to_start: - svc.dependency_provider._http_client = async_http_client - svc.dependency_provider._sdk_client = async_sdk - # Merge dependency overrides all_overrides = {} if dependency_overrides: diff --git a/packages/nmp_testing/src/nmp/testing/http_clients.py b/packages/nmp_testing/src/nmp/testing/http_clients.py new file mode 100644 index 0000000000..b05a010383 --- /dev/null +++ b/packages/nmp_testing/src/nmp/testing/http_clients.py @@ -0,0 +1,66 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""HTTP client helpers for tests that need platform endpoint routing.""" + +from __future__ import annotations + +from collections.abc import Callable, Coroutine +from typing import TypeAlias + +import httpx +from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint + +SyncMockHandler: TypeAlias = Callable[[httpx.Request], httpx.Response] +AsyncMockHandler: TypeAlias = Callable[[httpx.Request], Coroutine[None, None, httpx.Response]] +MockHandler: TypeAlias = SyncMockHandler | AsyncMockHandler + + +def _set_request_url(request: httpx.Request, url: httpx.URL) -> None: + if request.url == url: + return + request.url = url + if url.host: + request.headers["Host"] = url.netloc.decode("ascii") + + +class RoutedMockTransport(httpx.AsyncBaseTransport, httpx.BaseTransport): + """Mock transport that applies platform endpoint routing before handling requests.""" + + def __init__(self, handler: MockHandler, *, endpoint: PlatformEndpoint | None = None) -> None: + self._endpoint = endpoint or resolve_platform_endpoint() + self._mock_transport = httpx.MockTransport(handler) + + def handle_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return self._mock_transport.handle_request(request) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return await self._mock_transport.handle_async_request(request) + + +def routed_mock_transport( + handler: MockHandler, + *, + endpoint: PlatformEndpoint | None = None, +) -> RoutedMockTransport: + return RoutedMockTransport(handler, endpoint=endpoint) + + +def routed_mock_client( + handler: MockHandler, + *, + endpoint: PlatformEndpoint | None = None, +) -> httpx.Client: + return httpx.Client(transport=routed_mock_transport(handler, endpoint=endpoint)) + + +def routed_async_mock_client( + handler: MockHandler, + *, + endpoint: PlatformEndpoint | None = None, +) -> httpx.AsyncClient: + return httpx.AsyncClient(transport=routed_mock_transport(handler, endpoint=endpoint)) diff --git a/packages/nmp_testing/tests/unit/test_http_clients.py b/packages/nmp_testing/tests/unit/test_http_clients.py new file mode 100644 index 0000000000..8c8284e385 --- /dev/null +++ b/packages/nmp_testing/tests/unit/test_http_clients.py @@ -0,0 +1,45 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import httpx +import pytest +from nmp.common.config import PlatformConfig +from nmp.common.platform_endpoint import resolve_platform_endpoint +from nmp.testing.http_clients import routed_async_mock_client, routed_mock_client + + +@pytest.fixture +def service_endpoint(): + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={"entities": "http://entities:9090/entities-prefix"}, + ) + return resolve_platform_endpoint(config) + + +def test_routed_mock_client_applies_platform_endpoint_routes(service_endpoint): + captured: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request) + return httpx.Response(200, json={"ok": True}) + + with routed_mock_client(handler, endpoint=service_endpoint) as client: + client.get("http://platform:8080/apis/entities/v2/workspaces?limit=5") + + assert str(captured[0].url) == "http://entities:9090/entities-prefix/apis/entities/v2/workspaces?limit=5" + assert captured[0].headers["host"] == "entities:9090" + + +async def test_routed_async_mock_client_applies_platform_endpoint_routes(service_endpoint): + captured: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request) + return httpx.Response(200, json={"ok": True}) + + async with routed_async_mock_client(handler, endpoint=service_endpoint) as client: + await client.get("http://platform:8080/apis/entities/v2/workspaces?limit=5") + + assert str(captured[0].url) == "http://entities:9090/entities-prefix/apis/entities/v2/workspaces?limit=5" + assert captured[0].headers["host"] == "entities:9090" diff --git a/plugins/nemo-iron-swarm/tests/unit/test_synth_hitl.py b/plugins/nemo-iron-swarm/tests/unit/test_synth_hitl.py index 09070dc6f1..904dcf4e00 100644 --- a/plugins/nemo-iron-swarm/tests/unit/test_synth_hitl.py +++ b/plugins/nemo-iron-swarm/tests/unit/test_synth_hitl.py @@ -6,12 +6,12 @@ from __future__ import annotations import json -from types import SimpleNamespace from typing import Any, cast import httpx from nemo_iron_swarm_plugin.jobs import hitl from nemo_iron_swarm_plugin.jobs.synth_client import SynthClient +from nemo_platform import NeMoPlatform def test_synth_client_maps_endpoints() -> None: @@ -98,20 +98,19 @@ def handler(request: httpx.Request) -> httpx.Response: # The channel wraps the SDK via ``client_from_platform`` into a typed JobsClient; # model that with a real JobsClient over a mocked transport. - http_client = httpx.Client(transport=httpx.MockTransport(handler), base_url="http://platform:8080") - platform = SimpleNamespace( - base_url="http://platform:8080", - workspace="default", - _custom_headers={"Authorization": "Bearer x"}, - _client=http_client, - timeout=None, - max_retries=2, - _prepare_url=lambda url: url, - ) - - channel = hitl.StatusDetailsChannel(platform, name="job1", workspace="default", poll_interval=0.0) - channel.publish("interview", {"questions": [{"gap": "g"}]}) - assert published["interview"]["round"] == 1 - assert channel.await_response("interview") == [{"gap": "g", "answer": "a"}] - # Interview answers are accumulated so the run can persist the Q&A on the manifest for display. - assert channel.interview == [{"gap": "g", "answer": "a"}] + with httpx.Client(transport=httpx.MockTransport(handler), base_url="http://platform:8080") as http_client: + platform = NeMoPlatform( + base_url="http://platform:8080", + workspace="default", + default_headers={"Authorization": "Bearer x"}, + timeout=None, + max_retries=2, + http_client=http_client, + ) + + channel = hitl.StatusDetailsChannel(platform, name="job1", workspace="default", poll_interval=0.0) + channel.publish("interview", {"questions": [{"gap": "g"}]}) + assert published["interview"]["round"] == 1 + assert channel.await_response("interview") == [{"gap": "g", "answer": "a"}] + # Interview answers are accumulated so the run can persist the Q&A on the manifest for display. + assert channel.interview == [{"gap": "g", "answer": "a"}] diff --git a/plugins/nemo-safe-synthesizer/tests/unit/test_sdk.py b/plugins/nemo-safe-synthesizer/tests/unit/test_sdk.py index f1bd927edd..15570cead4 100644 --- a/plugins/nemo-safe-synthesizer/tests/unit/test_sdk.py +++ b/plugins/nemo-safe-synthesizer/tests/unit/test_sdk.py @@ -177,7 +177,8 @@ def test_job_builder_uploads_dataframe_and_creates_job() -> None: .with_hf_token_secret("hf-token") ) - job = builder.create_job(name="safe-synth-job") + with patch("nemo_safe_synthesizer_plugin.sdk.job.client_from_platform", return_value=MagicMock()): + job = builder.create_job(name="safe-synth-job") assert job.job_name == "safe-synth-job" client.files.upload.assert_called_once() @@ -205,7 +206,8 @@ def test_job_builder_creates_pretrained_model_job_for_adapter_reuse() -> None: .with_generate(num_records=25) ) - job = builder.create_job(name="adapter-reuse-job") + with patch("nemo_safe_synthesizer_plugin.sdk.job.client_from_platform", return_value=MagicMock()): + job = builder.create_job(name="adapter-reuse-job") assert job.job_name == "adapter-reuse-job" create_kwargs = client.safe_synthesizer.jobs.create.call_args.kwargs diff --git a/sdk/python/nemo-platform/src/nemo_platform/_client.py b/sdk/python/nemo-platform/src/nemo_platform/_client.py index 333bcf8386..bbbc6b38ac 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/_client.py +++ b/sdk/python/nemo-platform/src/nemo_platform/_client.py @@ -19,9 +19,17 @@ import os from typing import TYPE_CHECKING, Any, Mapping +from pathlib import Path from typing_extensions import Self, override import httpx +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env +from nemo_platform_plugin.client.auth import TokenProvider, TokenProviderAuth, AsyncTokenProvider +from nemo_platform_plugin.client.types import RetryPolicy +from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.platform_options import SyncPlatformClientOptions, AsyncPlatformClientOptions + +from nemo_platform._base_client import DefaultHttpxClient, DefaultAsyncHttpxClient from . import _exceptions from ._qs import Querystring @@ -35,8 +43,6 @@ not_given, ) from ._utils import ( - is_given, - is_mapping_t, get_async_library, ) from ._compat import cached_property @@ -48,12 +54,9 @@ SyncAPIClient, AsyncAPIClient, ) -from nemo_platform._base_client import DefaultAsyncHttpxClient, DefaultHttpxClient -from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR -from nemo_platform_plugin.client.tls import client_verify_from_env -from pathlib import Path if TYPE_CHECKING: + from .models import ModelsResource, AsyncModelsResource from .resources import ( iam, auth, @@ -73,11 +76,10 @@ experiments, ) from .resources.iam.iam import IamResource, AsyncIamResource + from .filesets.resources import FilesResource, AsyncFilesResource from .resources.auth.auth import AuthResource, AsyncAuthResource from .resources.jobs.jobs import JobsResource, AsyncJobsResource - from .filesets.resources import FilesResource, AsyncFilesResource from .resources.intake.intake import IntakeResource, AsyncIntakeResource - from .models import ModelsResource, AsyncModelsResource from .resources.secrets.secrets import SecretsResource, AsyncSecretsResource from .resources.adapters.adapters import AdaptersResource, AsyncAdaptersResource from .resources.entities.entities import EntitiesResource, AsyncEntitiesResource @@ -130,14 +132,42 @@ def _copy_requires_bootstrap( context_name: str | None, access_token: str | None, ) -> bool: - return ( - config_path is not None - or context_name is not None - or access_token is not None - or bool(os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)) + return config_path is not None or context_name is not None or access_token is not None + + +def _typed_client_default_headers( + custom_headers: Mapping[str, object], + transport_headers: Mapping[str, str], +) -> Mapping[str, str] | None: + headers = {str(key): value for key, value in custom_headers.items() if isinstance(value, str)} + if headers: + return headers + + default_transport_headers = {"accept", "accept-encoding", "connection", "user-agent", "host"} + headers = { + str(key): value + for key, value in transport_headers.items() + if str(key).lower() not in default_transport_headers and isinstance(value, str) + } + return headers or None + + +def _typed_client_retry(max_retries: int) -> RetryPolicy: + return RetryPolicy( + max_retries=max_retries, + retryable_status_codes=(408, 409, 429), + retry_all_server_errors=True, + respect_retry_decision_headers=True, + respect_retry_after_headers=True, ) +def _typed_client_timeout(timeout: float | Timeout | None) -> float | Timeout | None: + if timeout is None: + return httpx.Timeout(None) + return timeout + + class NeMoPlatform(SyncAPIClient): # client options workspace: str | None @@ -154,6 +184,7 @@ def __init__( max_retries: int = DEFAULT_MAX_RETRIES, default_headers: Mapping[str, str] | None = None, default_query: Mapping[str, object] | None = None, + token_provider: TokenProvider | None = None, # Configure a custom httpx client. # We provide a `DefaultHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`. # See the [httpx documentation](https://www.python-httpx.org/api/#client) for more details. @@ -249,6 +280,7 @@ def __init__( if workspace is None: workspace = client_init_kwargs.workspace default_headers = client_init_kwargs.default_headers + token_provider = client_init_kwargs.token_provider if client_init_kwargs.http_client is not None and not isinstance( client_init_kwargs.http_client, httpx.Client ): @@ -263,11 +295,13 @@ def __init__( if base_url is None: raise RuntimeError("NeMoPlatform client initialization failed: base_url is required") - client_verify = client_verify_from_env() - if http_client is None and client_verify is not True: - http_client = DefaultHttpxClient(verify=client_verify) + tls_config = httpx_tls_config_from_env() + if http_client is None and tls_config: + http_client = DefaultHttpxClient(**tls_config) self.workspace = workspace + self._token_provider = token_provider + self._token_provider_auth = TokenProviderAuth(token_provider) if token_provider is not None else None super().__init__( version=__version__, @@ -440,12 +474,16 @@ def copy( elif set_default_query is not None: params = set_default_query - if http_client is None and not _copy_requires_bootstrap( + requires_bootstrap = _copy_requires_bootstrap( config_path=config_path, context_name=context_name, access_token=access_token, - ): + ) + if http_client is None and not requires_bootstrap: http_client = self._client + extra_kwargs = dict(_extra_kwargs) + if not requires_bootstrap: + extra_kwargs.setdefault("token_provider", self._token_provider) return self.__class__( workspace=workspace or self.workspace, base_url=base_url or self.base_url, @@ -458,7 +496,7 @@ def copy( max_retries=self.max_retries if isinstance(max_retries, NotGiven) else max_retries, default_headers=headers, default_query=params, - **_extra_kwargs, + **extra_kwargs, ) # Alias for `copy` for nicer inline usage, e.g. @@ -522,6 +560,30 @@ def __getattr__(self, name: str) -> Any: self.__dict__[name] = instance return instance + @property + def custom_auth(self) -> TokenProviderAuth | None: + return self._token_provider_auth + + @property + def http_client(self) -> httpx.Client: + return self._client + + @property + def token_provider(self) -> TokenProvider | None: + return self._token_provider + + def typed_client_options(self) -> SyncPlatformClientOptions: + return SyncPlatformClientOptions( + base_url=str(self.base_url).rstrip("/"), + workspace=self.workspace, + default_headers=_typed_client_default_headers(self._custom_headers, self._client.headers), + timeout=_typed_client_timeout(self.timeout), + retry=_typed_client_retry(self.max_retries), + http_client=self._client, + url_resolver=self._prepare_url, + auth=self._token_provider, + ) + class AsyncNeMoPlatform(AsyncAPIClient): # client options @@ -540,6 +602,7 @@ def __init__( max_retries: int = DEFAULT_MAX_RETRIES, default_headers: Mapping[str, str] | None = None, default_query: Mapping[str, object] | None = None, + token_provider: TokenProvider | AsyncTokenProvider | None = None, # Configure a custom httpx client. # We provide a `DefaultAsyncHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`. # See the [httpx documentation](https://www.python-httpx.org/api/#asyncclient) for more details. @@ -654,6 +717,7 @@ async def main() -> None: if workspace is None: workspace = client_init_kwargs.workspace default_headers = client_init_kwargs.default_headers + token_provider = client_init_kwargs.token_provider if client_init_kwargs.http_client is not None and not isinstance( client_init_kwargs.http_client, httpx.AsyncClient ): @@ -668,11 +732,13 @@ async def main() -> None: if base_url is None: raise RuntimeError("NeMoPlatform client initialization failed: base_url is required") - client_verify = client_verify_from_env() - if http_client is None and client_verify is not True: - http_client = DefaultAsyncHttpxClient(verify=client_verify) + tls_config = httpx_tls_config_from_env() + if http_client is None and tls_config: + http_client = DefaultAsyncHttpxClient(**tls_config) self.workspace = workspace + self._token_provider = token_provider + self._token_provider_auth = TokenProviderAuth(token_provider) if token_provider is not None else None super().__init__( version=__version__, @@ -848,12 +914,16 @@ def copy( elif set_default_query is not None: params = set_default_query - if http_client is None and not _copy_requires_bootstrap( + requires_bootstrap = _copy_requires_bootstrap( config_path=config_path, context_name=context_name, access_token=access_token, - ): + ) + if http_client is None and not requires_bootstrap: http_client = self._client + extra_kwargs = dict(_extra_kwargs) + if not requires_bootstrap: + extra_kwargs.setdefault("token_provider", self._token_provider) return self.__class__( workspace=workspace or self.workspace, base_url=base_url or self.base_url, @@ -866,7 +936,7 @@ def copy( max_retries=self.max_retries if isinstance(max_retries, NotGiven) else max_retries, default_headers=headers, default_query=params, - **_extra_kwargs, + **extra_kwargs, ) # Alias for `copy` for nicer inline usage, e.g. @@ -930,6 +1000,30 @@ def __getattr__(self, name: str) -> Any: self.__dict__[name] = instance return instance + @property + def custom_auth(self) -> TokenProviderAuth | None: + return self._token_provider_auth + + @property + def http_client(self) -> httpx.AsyncClient: + return self._client + + @property + def token_provider(self) -> TokenProvider | AsyncTokenProvider | None: + return self._token_provider + + def typed_client_options(self) -> AsyncPlatformClientOptions: + return AsyncPlatformClientOptions( + base_url=str(self.base_url).rstrip("/"), + workspace=self.workspace, + default_headers=_typed_client_default_headers(self._custom_headers, self._client.headers), + timeout=_typed_client_timeout(self.timeout), + retry=_typed_client_retry(self.max_retries), + http_client=self._client, + url_resolver=self._prepare_url, + auth=self._token_provider, + ) + class NeMoPlatformWithRawResponse: _client: NeMoPlatform diff --git a/services/core/entities/src/nmp/core/entities/controllers/main.py b/services/core/entities/src/nmp/core/entities/controllers/main.py index a5a7416980..566afc3574 100644 --- a/services/core/entities/src/nmp/core/entities/controllers/main.py +++ b/services/core/entities/src/nmp/core/entities/controllers/main.py @@ -28,7 +28,7 @@ def handle_signal(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: platform_config = get_platform_config() logger.info("Starting entities controller") @@ -39,56 +39,56 @@ def run(parent_stop_signal: threading.Event | None = None): else: local_stop_signal = parent_stop_signal - nmp_sdk = get_async_platform_sdk( - as_service="entities", - internal=True, - ) - # Create a single event loop that will be shared for DB init and the cleanup controller, # so SQLAlchemy's async pool is bound to the same loop that later runs queries. - entities_config = get_service_config(EntitiesConfig) loop = asyncio.new_event_loop() - loop.run_until_complete(initialize_async_engine(entities_config)) - session_maker = loop.run_until_complete(get_async_session_maker()) - - workspace_repository = SQLAlchemyWorkspaceRepository(session_maker) - - if not wait_for_service_ready(platform_config, "entities", local_stop_signal): - if local_stop_signal.is_set(): - logger.info("Shutdown requested before server became ready") - return - logger.warning("Server did not become ready in time, starting loops anyway") - - cleanup_controller = WorkspaceCleanup( - nmp_sdk=nmp_sdk, - workspace_repository=workspace_repository, - stop_signal=local_stop_signal, - loop=loop, - ) - cleanup_controller_monitored = TrackLastExecutionTime(cleanup_controller) - - cleanup_loop = Loop( - TimedLoopWaiter(entities_config.workspace_cleanup_interval, stop_signal=local_stop_signal), - cleanup_controller_monitored, - stop_signal=local_stop_signal, - ) - - controller_manager = ControllerManager.get_instance() - controller_manager.register("workspace_cleanup", cleanup_loop) - - cleanup_loop.start() - logger.info("Entities controller started successfully") - + nmp_sdk = get_async_platform_sdk(as_service="entities", internal=True) + cleanup_loop: Loop | None = None try: + entities_config = get_service_config(EntitiesConfig) + loop.run_until_complete(initialize_async_engine(entities_config)) + session_maker = loop.run_until_complete(get_async_session_maker()) + + workspace_repository = SQLAlchemyWorkspaceRepository(session_maker) + + if not wait_for_service_ready(platform_config, "entities", local_stop_signal): + if local_stop_signal.is_set(): + logger.info("Shutdown requested before server became ready") + return + logger.warning("Server did not become ready in time, starting loops anyway") + + cleanup_controller = WorkspaceCleanup( + nmp_sdk=nmp_sdk, + workspace_repository=workspace_repository, + stop_signal=local_stop_signal, + loop=loop, + ) + cleanup_controller_monitored = TrackLastExecutionTime(cleanup_controller) + + cleanup_loop = Loop( + TimedLoopWaiter(entities_config.workspace_cleanup_interval, stop_signal=local_stop_signal), + cleanup_controller_monitored, + stop_signal=local_stop_signal, + ) + + controller_manager = ControllerManager.get_instance() + controller_manager.register("workspace_cleanup", cleanup_loop) + + cleanup_loop.start() + logger.info("Entities controller started successfully") + while not local_stop_signal.is_set(): local_stop_signal.wait(timeout=1) except KeyboardInterrupt: logger.info("Received keyboard interrupt, stopping entities controller") finally: - cleanup_loop.stop() - cleanup_loop.join(timeout=10) - if cleanup_loop.is_alive(): - logger.warning("Workspace cleanup loop did not stop in time") + if cleanup_loop is not None: + cleanup_loop.stop() + cleanup_loop.join(timeout=10) + if cleanup_loop.is_alive(): + logger.warning("Workspace cleanup loop did not stop in time") + loop.run_until_complete(nmp_sdk.close()) + loop.close() logger.info("Entities controller stopped") diff --git a/services/core/files/src/nmp/core/files/api/endpoint_helpers.py b/services/core/files/src/nmp/core/files/api/endpoint_helpers.py index 7bdc8b7685..2a8f5fcb80 100644 --- a/services/core/files/src/nmp/core/files/api/endpoint_helpers.py +++ b/services/core/files/src/nmp/core/files/api/endpoint_helpers.py @@ -302,7 +302,13 @@ async def resolve_storage_secrets_for_user( auth_client: AuthClient, ) -> dict[str, str]: """Resolve storage secrets using delegated headers on request-scoped SDK.""" - service_sdk = get_async_platform_sdk(as_service="files", internal=True, on_behalf_of=auth_client.principal.id) + service_sdk = get_async_platform_sdk( + as_service="files", + internal=True, + http_client=sdk._client, + on_behalf_of=auth_client.principal.id, + base_url=str(sdk.base_url).rstrip("/"), + ) return await resolve_storage_secrets(storage, workspace, service_sdk) diff --git a/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py b/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py index dcea2c16e3..3af085d358 100644 --- a/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py +++ b/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py @@ -10,7 +10,6 @@ import aiohttp from fastapi import Request from fastapi.exceptions import RequestValidationError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service import RouterConfig, Service from starlette import status from starlette.responses import JSONResponse @@ -96,7 +95,7 @@ async def on_startup(self) -> None: from nmp.core.inference_gateway.api.virtual_model_cache import VirtualModelCache from nmp.core.inference_gateway.config import config as inference_gateway_config - sdk = get_async_platform_sdk(as_service="inference-gateway", internal=True) + sdk = self.dependency_provider.get_sdk_client(as_service="inference-gateway") # Initialize caches model_cache = set_global_model_cache(ModelCache(secret_value_ttl=inference_gateway_config.secrets_ttl_sec)) diff --git a/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py b/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py index 3c6cda1236..c8c318b2ed 100644 --- a/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py +++ b/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py @@ -60,29 +60,38 @@ def _enable_post_response_task_tracking(client_context: ClientContext) -> None: _app_from(client_context).state.pending_post_response_tasks = [] +@contextmanager +def _loopback_plugin_sdk_provider(base_url: str) -> Generator[None, None, None]: + """Make plugin-created SDKs target the loopback test app.""" + from unittest.mock import patch + + from nemo_platform_plugin import sdk_provider as sdk_provider_module + from nemo_platform_plugin.sdk_provider import DefaultSDKProvider + + previous_provider = getattr(sdk_provider_module, "_cached_provider", None) + with patch.dict("os.environ", {"NMP_BASE_URL": base_url}): + sdk_provider_module.set_sdk_provider(DefaultSDKProvider()) + try: + yield + finally: + sdk_provider_module.set_sdk_provider(previous_provider) + + @contextmanager def _build_app_context( *extra_services: ServiceFactory, ) -> Generator[ClientContext, None, None]: """Yield an IGW + Models + extras :class:`ClientContext` (module-lived). - Two module-scope hazards are neutralised here: - - 1. The 3-second background ``refresh_model_cache_task``. ``on_startup`` - reads ``refresh_model_cache_interval_sec`` from the module-level - config snapshot (captured at first import), so a ``service_configs`` - override is too late. Patch the snapshot field to 0 *before* - entering ``create_test_client`` and ``on_startup`` never schedules - the loop. - 2. The shared SDK HTTP client's ``aclose``. Plugins like - ``nemo-guardrails`` call ``await sdk.close()`` in ``on_shutdown``, - which would close the shared client for every later test in the - module. Patch ``aclose`` to a no-op for the module's lifetime; - ``ASGITransport`` is in-process so nothing actually leaks. + The 3-second background ``refresh_model_cache_task`` is neutralised here. + ``on_startup`` reads ``refresh_model_cache_interval_sec`` from the + module-level config snapshot (captured at first import), so a + ``service_configs`` override is too late. Patch the snapshot field to 0 + *before* entering ``create_test_client`` and ``on_startup`` never schedules + the loop. """ from unittest.mock import patch - from nmp.common import sdk_factory as sdk_factory_module from nmp.core.inference_gateway import config as igw_config_module from nmp.core.inference_gateway.service import InferenceGatewayService from nmp.core.models.service import ModelsService @@ -95,21 +104,7 @@ def _build_app_context( client_type=ClientContext, igw_mock_provider_mode=False, ) as client_context: - shared_async_client = sdk_factory_module._test_http_client - if shared_async_client is None: - yield client_context - return - - original_aclose = shared_async_client.aclose - - async def _noop_aclose() -> None: - return None - - shared_async_client.aclose = _noop_aclose # type: ignore[method-assign] - try: - yield client_context - finally: - shared_async_client.aclose = original_aclose # type: ignore[method-assign] + yield client_context @pytest.fixture(scope="module") @@ -226,6 +221,8 @@ def _build_loopback_harness( * ``get_platform_config`` is patched at IGW's middleware-registry import site so :meth:`get_openai_compatible_inference_url_and_model` returns URLs reachable from the test process. + * The plugin SDK provider is temporarily rebound to the loopback URL + so middleware can fetch platform entities created by the test SDK. Both patches roll back before the next test runs, so a plain ``igw_plugin_harness`` test sharing the module doesn't observe them. @@ -269,6 +266,7 @@ def _restore_http_client_override() -> None: stack.callback(_restore_http_client_override) stack.enter_context(override_platform_base_url(igw_loopback_base_url)) + stack.enter_context(_loopback_plugin_sdk_provider(igw_loopback_base_url)) harness = cast( IGWLoopbackHarness, diff --git a/services/core/inference-gateway/tests/integration/conftest.py b/services/core/inference-gateway/tests/integration/conftest.py index baeabe10a3..99160567c2 100644 --- a/services/core/inference-gateway/tests/integration/conftest.py +++ b/services/core/inference-gateway/tests/integration/conftest.py @@ -90,6 +90,10 @@ def init(self) -> None: """No-op init for mock backend.""" pass + def shutdown(self) -> None: + """No-op shutdown for mock backend.""" + pass + async def create_model_deployment(self, ctx: Any) -> DeploymentStatusUpdate: """Record call and return configured response.""" self.create_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) @@ -105,9 +109,16 @@ async def get_model_deployment_status(self, ctx: Any) -> DeploymentStatusUpdate: self.status_calls.append(ctx.model_deployment) return self.status_response - async def delete_model_deployment(self, deployment: Any) -> DeploymentStatusUpdate: + async def delete_model_deployment( + self, + workspace: str, + name: str, + *, + deleting_elapsed_seconds: float | None = None, + ) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.delete_calls.append(deployment) + del deleting_elapsed_seconds + self.delete_calls.append((workspace, name)) return self.delete_response @@ -352,7 +363,10 @@ def patched_get_qualified_image(name: str, tag=None, registry=None): return_value=mock_platform_config, ), patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk") as mock_sdk_factory, - patch("nemo_platform_plugin.sdk_provider.get_async_platform_sdk") as mock_sdk, + patch( + "nmp.core.models.controllers.backends.deployments_plugin.backend.get_async_platform_sdk" + ) as mock_backend_sdk, + patch("nemo_deployments_plugin.controller.get_async_platform_sdk") as mock_deployments_sdk, patch("nemo_deployments_plugin.config.DeploymentsConfig.get", return_value=deployments_config), patch( "nemo_platform_plugin.jobs.image.get_qualified_image", @@ -360,23 +374,21 @@ def patched_get_qualified_image(name: str, tag=None, registry=None): ), ): mock_sdk_factory.return_value = test_clients.async_sdk - mock_sdk.return_value = test_clients.async_sdk + mock_backend_sdk.return_value = test_clients.async_sdk + mock_deployments_sdk.return_value = test_clients.async_sdk - controller = ModelsController( + class ModelsControllerWithDeploymentsPlugin(ModelsController): + def step(self) -> None: + super().step() + self._loop.run_until_complete(deployments_controller.reconcile()) + + controller = ModelsControllerWithDeploymentsPlugin( backend_registry=backend_registry, stop_signal=None, ) controller._provider_reconciler.reconcile_model_providers = AsyncMock(return_value=None) controller._loop.run_until_complete(deployments_controller.on_startup()) - original_step = controller.step - - def step_with_deployments_plugin() -> None: - original_step() - controller._loop.run_until_complete(deployments_controller.reconcile()) - - controller.step = step_with_deployments_plugin - model_cache = global_model_cache() yield controller, model_cache, test_clients.sdk, mock_nim_image, docker_test_context, test_clients.async_sdk @@ -416,7 +428,7 @@ async def trigger_cache_refresh( @pytest.hookimpl(tryfirst=True, hookwrapper=True) -def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, None, None]: +def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, Any, None]: """Store test results on the item for fixture access.""" outcome = yield rep = outcome.get_result() diff --git a/services/core/inference-gateway/tests/unit/conftest.py b/services/core/inference-gateway/tests/unit/conftest.py index 3aeabfd212..aec6a54fac 100644 --- a/services/core/inference-gateway/tests/unit/conftest.py +++ b/services/core/inference-gateway/tests/unit/conftest.py @@ -170,9 +170,8 @@ def app_and_client( ], ) - mocker.patch("nmp.core.inference_gateway.service.get_async_platform_sdk", return_value=mock_nmp_sdk) - service = InferenceGatewayService() + mocker.patch.object(service.dependency_provider, "get_sdk_client", return_value=mock_nmp_sdk) app = service.app app.dependency_overrides[global_http_client] = lambda: mock_proxy_client app.dependency_overrides[global_model_cache] = lambda: model_cache diff --git a/services/core/inference-gateway/tests/unit/test_service.py b/services/core/inference-gateway/tests/unit/test_service.py index 5c1cf8b583..6dbba82cbd 100644 --- a/services/core/inference-gateway/tests/unit/test_service.py +++ b/services/core/inference-gateway/tests/unit/test_service.py @@ -17,7 +17,6 @@ async def test_debug_startup_hydrates_model_entity_metadata(mocker): http_client = Mock() http_client.close = AsyncMock() - mocker.patch("nmp.core.inference_gateway.service.get_async_platform_sdk", return_value=sdk) mocker.patch( "nmp.core.inference_gateway.api.middleware_registry.load_middleware_plugins", AsyncMock(return_value=MiddlewareRegistry()), @@ -38,6 +37,7 @@ async def test_debug_startup_hydrates_model_entity_metadata(mocker): ) service = InferenceGatewayService() + mocker.patch.object(service.dependency_provider, "get_sdk_client", return_value=sdk) await service.on_startup() await service.on_shutdown() diff --git a/services/core/jobs/src/nmp/core/jobs/api/dependencies.py b/services/core/jobs/src/nmp/core/jobs/api/dependencies.py index 1a1a7049df..0e89bf77a0 100644 --- a/services/core/jobs/src/nmp/core/jobs/api/dependencies.py +++ b/services/core/jobs/src/nmp/core/jobs/api/dependencies.py @@ -3,27 +3,22 @@ """FastAPI dependencies for the Jobs API.""" -from fastapi import Depends, Request +from fastapi import Depends from nemo_platform import AsyncNeMoPlatform from nmp.common.entities.client import EntityClient -from nmp.common.sdk_factory import get_async_platform_sdk -from nmp.common.service.dependencies import get_entity_client +from nmp.common.service.dependencies import get_entity_client, get_sdk_client from nmp.core.jobs.app.dispatcher import JobDispatcher -async def get_sdk_with_auth(request: Request) -> AsyncNeMoPlatform: +async def get_sdk_with_auth( + sdk: AsyncNeMoPlatform = Depends(get_sdk_client), +) -> AsyncNeMoPlatform: """Get SDK client with current request's auth headers. - This dependency creates a new SDK instance with the current user's - auth context propagated. This is needed for internal service calls - that require authorization (e.g., creating filesets). - - Args: - request: The FastAPI request object (needed to ensure we're in request context) + The platform overrides get_sdk_client with a request-scoped SDK that + preserves the service provider's HTTP transport. """ - # By taking request as parameter, we ensure this runs in request context - # where auth headers are available via context vars - return get_async_platform_sdk() + return sdk async def dep_dispatcher( diff --git a/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py b/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py index 2fad1befba..907b7bf489 100644 --- a/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py +++ b/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py @@ -31,7 +31,6 @@ PlatformJobStatusResponse, ) from nmp.common.observability import scoped_app_ctx -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service.dependencies import get_sdk_client from nmp.core.jobs.api.dependencies import dep_dispatcher from nmp.core.jobs.api.v2.jobs.schemas import ( @@ -694,7 +693,7 @@ async def download_job_result( job_name=job, workspace=workspace, artifact_url=result.artifact_url, - files_sdk=get_async_platform_sdk(), + files_sdk=dispatcher.sdk, ) background_tasks.add_task(lambda: tmp_dir_path.cleanup_tmp_dir()) return FileResponse(path=tmp_dir_path.path, filename=filename, background=background_tasks) diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/main.py b/services/core/jobs/src/nmp/core/jobs/controllers/main.py index 0a76525a3c..88485f0483 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/main.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/main.py @@ -24,7 +24,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: # Create logger after configuration is set up logger.info("Starting jobs controller") @@ -42,57 +42,59 @@ def run(parent_stop_signal: threading.Event | None = None): nmp_sdk = get_platform_sdk(as_service="jobs", internal=True) logger.debug("Platform SDK initialized successfully.") - # from_config also prunes the shared ``profiles`` list so advertised executors - # match backends that actually registered (e.g. Docker skipped when unavailable). - backend_registry = BackendRegistry.from_config(nmp_sdk=nmp_sdk, profiles=profiles) - logger.info("Executor backends registry initialized successfully.") - - # Wait for the jobs service to be ready before starting control loops (polls /status so we can start once jobs is ready) - platform_config = get_platform_config() - if not wait_for_service_ready(platform_config, "jobs", local_stop_signal): - if local_stop_signal.is_set(): - logger.info("Shutdown requested before server became ready") - return - logger.warning("Server did not become ready in time, starting loops anyway") - - # Job scheduling loop - job_scheduler = JobScheduler(backend_registry, nmp_sdk, stop_signal=local_stop_signal) - job_scheduler_monitored = TrackLastExecutionTime(job_scheduler) - job_scheduler_loop = Loop( - TimedLoopWaiter(jobs_config.schedule_interval_seconds, stop_signal=local_stop_signal), - job_scheduler_monitored, - stop_signal=local_stop_signal, - ) - - # Job reconciler loop - job_reconciler = JobReconciler(backend_registry, nmp_sdk, stop_signal=local_stop_signal) - job_reconciler_monitored = TrackLastExecutionTime(job_reconciler) - job_reconciler_loop = Loop( - TimedLoopWaiter(jobs_config.reconcile_interval_seconds, stop_signal=local_stop_signal), - job_reconciler_monitored, - shutdown_func=backend_registry.shutdown_all_backends, - stop_signal=local_stop_signal, - ) - - # Register loops with ControllerManager - controller_manager = ControllerManager.get_instance() - controller_manager.register("job_scheduler", job_scheduler_loop) - controller_manager.register("job_reconciler", job_reconciler_loop) - - # Start control loops - job_scheduler_loop.start() - job_reconciler_loop.start() - logger.info("Jobs controller started successfully") - - # Main loop - while not local_stop_signal.is_set(): - time.sleep(1) - - logger.info("Shutting down control loops...") - job_scheduler_loop.stop() - job_reconciler_loop.stop() - - for loop in [job_scheduler_loop, job_reconciler_loop]: - loop.join() - - logger.info("All jobs control loops have been shut down.") + job_scheduler_loop: Loop | None = None + job_reconciler_loop: Loop | None = None + try: + # from_config also prunes the shared ``profiles`` list so advertised executors + # match backends that actually registered (e.g. Docker skipped when unavailable). + backend_registry = BackendRegistry.from_config(nmp_sdk=nmp_sdk, profiles=profiles) + logger.info("Executor backends registry initialized successfully.") + + # Wait for the jobs service to be ready before starting control loops (polls /status so we can start once jobs is ready) + platform_config = get_platform_config() + if not wait_for_service_ready(platform_config, "jobs", local_stop_signal): + if local_stop_signal.is_set(): + logger.info("Shutdown requested before server became ready") + return + logger.warning("Server did not become ready in time, starting loops anyway") + + # Job scheduling loop + job_scheduler = JobScheduler(backend_registry, nmp_sdk, stop_signal=local_stop_signal) + job_scheduler_monitored = TrackLastExecutionTime(job_scheduler) + job_scheduler_loop = Loop( + TimedLoopWaiter(jobs_config.schedule_interval_seconds, stop_signal=local_stop_signal), + job_scheduler_monitored, + stop_signal=local_stop_signal, + ) + + # Job reconciler loop + job_reconciler = JobReconciler(backend_registry, nmp_sdk, stop_signal=local_stop_signal) + job_reconciler_monitored = TrackLastExecutionTime(job_reconciler) + job_reconciler_loop = Loop( + TimedLoopWaiter(jobs_config.reconcile_interval_seconds, stop_signal=local_stop_signal), + job_reconciler_monitored, + shutdown_func=backend_registry.shutdown_all_backends, + stop_signal=local_stop_signal, + ) + + # Register loops with ControllerManager + controller_manager = ControllerManager.get_instance() + controller_manager.register("job_scheduler", job_scheduler_loop) + controller_manager.register("job_reconciler", job_reconciler_loop) + + # Start control loops + job_scheduler_loop.start() + job_reconciler_loop.start() + logger.info("Jobs controller started successfully") + + # Main loop + while not local_stop_signal.is_set(): + time.sleep(1) + finally: + logger.info("Shutting down control loops...") + for loop in (job_scheduler_loop, job_reconciler_loop): + if loop is not None: + loop.stop() + loop.join() + nmp_sdk.close() + logger.info("All jobs control loops have been shut down.") diff --git a/services/core/jobs/tests/conftest.py b/services/core/jobs/tests/conftest.py index 213e2046d2..3c9923b099 100644 --- a/services/core/jobs/tests/conftest.py +++ b/services/core/jobs/tests/conftest.py @@ -9,12 +9,14 @@ from typing import AsyncGenerator from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest import pytest_asyncio from fastapi import FastAPI from httpx import ASGITransport, AsyncClient from nemo_platform import AsyncNeMoPlatform from nemo_platform_plugin.capabilities import reset_capability_cache +from nemo_platform_plugin.entities import EntityClient as PluginEntityClient from nemo_platform_plugin.jobs.api_factory import ContainerSpec as FactoryContainerSpec from nemo_platform_plugin.jobs.api_factory import CPUExecutionProviderSpec as FactoryCPUExecutionProviderSpec from nemo_platform_plugin.jobs.api_factory import PlatformJobEnvironmentVariableParam, job_route_factory @@ -316,9 +318,21 @@ def mock_nmp_client(_mock_files_client, mock_jobs_client): patches dispatch on the requested client type: ``JobsClient`` requests resolve to ``mock_jobs_client``; anything else falls back to the files client. """ + http_client = httpx.Client( + transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request)), + base_url="http://test", + ) mock_client = MagicMock() mock_client.beta = MagicMock() mock_client.jobs = MagicMock() + mock_client.base_url = "http://test" + mock_client.workspace = "default" + mock_client._custom_headers = {} + mock_client._client = http_client + mock_client.timeout = httpx.Timeout(60.0) + mock_client.max_retries = 0 + mock_client._prepare_url = lambda url: url + mock_client.token_provider = None from nemo_platform_plugin.jobs.client import JobsClient @@ -332,6 +346,7 @@ def _dispatch(_sdk, client_type): patch(f"{module}.client_from_platform", side_effect=_dispatch) for module in _JOBS_CLIENT_CONTROLLER_MODULES ] with ExitStack() as stack: + stack.enter_context(http_client) for patcher in patchers: stack.enter_context(patcher) yield mock_client @@ -349,6 +364,7 @@ async def download_artifact(artifact_url: str, local_dir: str | Path | None = No return TmpDirPath(path=m._path, tmp_dir=m._tmp_dir) m.download_artifact = download_artifact + m.aclose = AsyncMock() return m @@ -429,7 +445,7 @@ def mock_platform_config() -> PlatformConfig: """Real PlatformConfig for controller tests (get_service_url, loopback_address, to_shared_envvars work correctly).""" return PlatformConfig( # type: ignore[abstract] base_url="http://localhost:8080", - files_url="http://localhost:8080", + service_discovery={"files": "http://localhost:8080"}, image_pull_secrets=[ImagePullSecret(name="global-pull-secret")], loopback_address=None, ) @@ -554,9 +570,9 @@ def hello_world_job_config( workspace: str, input_spec: HelloWorldJobConfig, output_spec: HelloWorldJobConfig, - entity_client: EntityClient, + entity_client: PluginEntityClient, job_name: str | None, - sdk, + sdk: AsyncNeMoPlatform, ) -> FactoryPlatformJobSpec: return FactoryPlatformJobSpec( steps=[ diff --git a/services/core/jobs/tests/controllers/test_base.py b/services/core/jobs/tests/controllers/test_base.py index 8ccefa73d3..a3c5cbb2a0 100644 --- a/services/core/jobs/tests/controllers/test_base.py +++ b/services/core/jobs/tests/controllers/test_base.py @@ -9,7 +9,9 @@ from unittest.mock import MagicMock, patch from urllib.parse import urlunsplit +import httpx import pytest +from nemo_platform import NeMoPlatform from nmp.common.config import PlatformConfig from nmp.common.jobs.constants import EPHEMERAL_TASK_STORAGE_PATH_ENVVAR, PERSISTENT_JOB_STORAGE_PATH_ENVVAR from nmp.common.jobs.schemas import PlatformJobStatus @@ -509,15 +511,13 @@ def test_loopback_replacement_applies_to_base_url_fallback(self): def _make_step( staleness_timeout: int = 0, created_at: datetime.datetime | None = None, - step_spec: PlatformJobStepSpec | None = ..., # type: ignore[assignment] ) -> PlatformJobStepWithContext: - if step_spec is ...: - step_spec = PlatformJobStepSpec( - name="test-step", - executor=CPUExecutionProvider(provider="cpu", profile="default", container=ContainerSpec(image="img")), - config={}, - lifecycle=StepLifecycle(staleness_timeout_seconds=staleness_timeout), - ) + step_spec = PlatformJobStepSpec( + name="test-step", + executor=CPUExecutionProvider(provider="cpu", profile="default", container=ContainerSpec(image="img")), + config={}, + lifecycle=StepLifecycle(staleness_timeout_seconds=staleness_timeout), + ) return PlatformJobStepWithContext( id="step-1", job="job-1", @@ -538,8 +538,9 @@ def _make_task(status: str = "active", updated_at: datetime.datetime | None = No return task -def _make_backend(mock_sdk: MagicMock | None = None) -> MockKubernetesCPUJobBackend: - sdk = mock_sdk or MagicMock() +def _make_backend() -> MockKubernetesCPUJobBackend: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + sdk = NeMoPlatform(base_url="http://test", http_client=http_client, workspace="default") return MockKubernetesCPUJobBackend(nmp_sdk=sdk, execution_profile_config=MagicMock(), profile_name="default") @@ -575,10 +576,10 @@ def test_disabled_when_timeout_is_zero(self): assert backend.check_step_is_stale(step) is False - def test_disabled_when_lifecycle_is_none(self): + def test_disabled_when_step_spec_is_none(self): backend = _make_backend() step = _make_step() - step.step_spec.lifecycle = None + step.step_spec = None assert backend.check_step_is_stale(step) is False diff --git a/services/core/models/src/nmp/core/models/api/dependencies.py b/services/core/models/src/nmp/core/models/api/dependencies.py index b50a2ade03..eacf2da598 100644 --- a/services/core/models/src/nmp/core/models/api/dependencies.py +++ b/services/core/models/src/nmp/core/models/api/dependencies.py @@ -17,16 +17,18 @@ def get_model_entity_service( entity_client: EntityClient = Depends(get_entity_client), + nmp_sdk: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> ModelEntityService: """Dependency to get ModelEntityService instance.""" - return ModelEntityService(entity_client) + return ModelEntityService(entity_client, sdk=nmp_sdk) def get_adapter_entity_service( entity_client: EntityClient = Depends(get_entity_client), + nmp_sdk: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> AdapterEntityService: """Dependency to get AdapterEntityService instance.""" - return AdapterEntityService(entity_client) + return AdapterEntityService(entity_client, sdk=nmp_sdk) def get_model_provider_service( diff --git a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py index 9278971d8f..d259d475ee 100644 --- a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py @@ -11,7 +11,6 @@ from nmp.common.api.parsed_filter import ParsedFilter from nmp.common.entities import ALL_WORKSPACES, ListResponse from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.core.models.api.service.model_entity_service import _adapter_to_adapter_schema, get_fileset_and_files_list from nmp.core.models.constants import parse_model_ref from nmp.core.models.entities import Adapter, Model @@ -24,9 +23,9 @@ class AdapterEntityService: """Service for adapter CRUD, scoped to a workspace, with model reference from path or body.""" - def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform | None = None) -> None: + def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform) -> None: self.entity_client = entity_client - self.sdk = sdk or get_async_platform_sdk() + self.sdk = sdk async def _fetch_all_entities( self, diff --git a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py index e1958ac756..f4c1a60888 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py @@ -17,7 +17,6 @@ from nmp.common.auth import AuthClient from nmp.common.entities import ALL_WORKSPACES, ListResponse from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.core.models.api.permissions import can_set_tool_call_plugin, check_fileset_access from nmp.core.models.config import config from nmp.core.models.entities import Adapter, Model, ModelDeploymentConfig @@ -235,9 +234,9 @@ async def validate_tool_call_plugin_allowed(auth_client: AuthClient, workspace: class ModelEntityService: """Service layer for Model Entity operations.""" - def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform | None = None): + def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform): self.entity_client = entity_client - self.sdk = sdk or get_async_platform_sdk() + self.sdk = sdk async def _fetch_all_entities( self, diff --git a/services/core/models/src/nmp/core/models/api/v2/models.py b/services/core/models/src/nmp/core/models/api/v2/models.py index 4354ddfa3c..8f14863023 100644 --- a/services/core/models/src/nmp/core/models/api/v2/models.py +++ b/services/core/models/src/nmp/core/models/api/v2/models.py @@ -26,7 +26,6 @@ from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, get_auth_client from nmp.common.entities.client import EntityConflictError, EntityNotFoundError, EntityValidationError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service.dependencies import get_sdk_client from nmp.core.models.api.dependencies import get_adapter_entity_service, get_model_entity_service from nmp.core.models.api.permissions import check_fileset_access @@ -156,7 +155,7 @@ async def create_model( # add sdk job creation here for checkpoint metadata if created_model.fileset: - await start_update_model_spec_job(created_model) + await start_update_model_spec_job(created_model, nmp_sdk) return created_model @@ -264,8 +263,7 @@ async def get_model( return model_entity -async def start_update_model_spec_job(model_entity: ModelEntity): - sdk = get_async_platform_sdk(as_service="models", internal=True) +async def start_update_model_spec_job(model_entity: ModelEntity, sdk: AsyncNeMoPlatform) -> None: model_spec_task_config = ModelSpecTaskConfig(workspace=model_entity.workspace, name=model_entity.name) task_spec = PlatformJobSpec( steps=[ @@ -421,7 +419,7 @@ async def update_model( raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to update model entity") if updated_model.fileset and (updated_model.fileset != original_fileset or not updated_model.spec): - await start_update_model_spec_job(updated_model) + await start_update_model_spec_job(updated_model, nmp_sdk) return updated_model diff --git a/services/core/models/src/nmp/core/models/config.py b/services/core/models/src/nmp/core/models/config.py index f4b3319280..57f3bd2cc2 100644 --- a/services/core/models/src/nmp/core/models/config.py +++ b/services/core/models/src/nmp/core/models/config.py @@ -454,7 +454,7 @@ class ModelsConfig(create_service_config_class("models")): # type: ignore # Module-level singleton instances config = get_service_config(ModelsConfig) -backends = merge_backends( +backends: dict[BackendName, BackendConfig] = merge_backends( config.controller.backends, get_default_backends_for_runtime(get_platform_config().runtime), ) diff --git a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py index cfb016280b..36d5dd93ca 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py @@ -3,6 +3,7 @@ """Models ServiceBackend backed by nemo-deployments plugin entities.""" +import asyncio import logging from typing import Any @@ -41,6 +42,7 @@ class DeploymentsPluginServiceBackend(ServiceBackend): def __init__(self, nmp_sdk: AsyncNeMoPlatform, config: dict[str, Any], huggingface_model_puller: str) -> None: self._backend_config: DeploymentsPluginConfig | None = None self._entities: NemoEntitiesClient | None = None + self._entities_sdk: AsyncNeMoPlatform | None = None self._huggingface_model_puller = huggingface_model_puller super().__init__(nmp_sdk, config) @@ -48,11 +50,16 @@ def init(self) -> None: self._backend_config = DeploymentsPluginConfig(**self._config) def shutdown(self) -> None: + entities_sdk = self._entities_sdk self._entities = None + self._entities_sdk = None + if entities_sdk is not None: + asyncio.run(entities_sdk.close()) def _entity_client(self) -> NemoEntitiesClient: if self._entities is None: sdk = get_async_platform_sdk(as_service="models", internal=True) + self._entities_sdk = sdk self._entities = NemoEntitiesClient(client_from_platform(sdk, AsyncEntitiesClient)) return self._entities diff --git a/services/core/models/src/nmp/core/models/controllers/backends/registry.py b/services/core/models/src/nmp/core/models/controllers/backends/registry.py index 06914ef188..bd1cb19a8c 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/registry.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/registry.py @@ -3,8 +3,9 @@ """Backend registry for Models Controller service.""" +from collections.abc import Mapping from logging import getLogger -from typing import Dict, Self +from typing import Any, Dict, Protocol, Self from nemo_platform import AsyncNeMoPlatform from nmp.core.models.controllers.backends.backends import ServiceBackend @@ -23,9 +24,19 @@ # Type alias for the backend name BackendName = str + +class BackendFactory(Protocol): + def __call__( + self, + nmp_sdk: AsyncNeMoPlatform, + config: dict[str, Any], + huggingface_model_puller: str, + ) -> ServiceBackend: ... + + # The deployments_plugin backend is resolved lazily because it imports the # optional `nemo_deployments_plugin` package. -backend_classes: Dict[BackendName, type[ServiceBackend]] = {} +backend_classes: Dict[BackendName, BackendFactory] = {} _LAZY_BACKEND_NAMES = frozenset({"deployments_plugin"}) @@ -37,8 +48,8 @@ def _resolve_backend_class( - name: BackendName, available_backends: Dict[BackendName, type[ServiceBackend]] -) -> type[ServiceBackend]: + name: BackendName, available_backends: Mapping[BackendName, BackendFactory] +) -> BackendFactory: """Return the backend class for ``name``, importing optional backends lazily.""" if name in available_backends: return available_backends[name] @@ -83,9 +94,9 @@ def __init__(self, registry: Dict[BackendName, ServiceBackend]) -> None: def from_config( cls, nmp_sdk: AsyncNeMoPlatform, - backend_configs: Dict[BackendName, BackendConfig], + backend_configs: Mapping[BackendName, BackendConfig], huggingface_model_puller: str, - available_backends: Dict[BackendName, type[ServiceBackend]] | None = None, + available_backends: Mapping[BackendName, BackendFactory] | None = None, ) -> Self: """Create a BackendRegistry from backend configurations. diff --git a/services/core/models/src/nmp/core/models/controllers/main.py b/services/core/models/src/nmp/core/models/controllers/main.py index 92824a9ee8..bae0cfd66f 100644 --- a/services/core/models/src/nmp/core/models/controllers/main.py +++ b/services/core/models/src/nmp/core/models/controllers/main.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import asyncio import logging import signal import threading @@ -11,7 +12,7 @@ from nmp.common.service.api.health import wait_for_service_ready from nmp.core.models.config import backends from nmp.core.models.config import config as models_config -from nmp.core.models.controllers.backends.registry import BackendRegistry +from nmp.core.models.controllers.backends.registry import BackendConfig, BackendRegistry from nmp.core.models.controllers.models_controller import ModelsController stop_signal = threading.Event() @@ -35,7 +36,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: """Run the Models Controller with its control loop.""" global models_controller_monitored @@ -56,48 +57,50 @@ def run(parent_stop_signal: threading.Event | None = None): # Initialize NeMo Platform SDK (used for all API interactions including secrets) nmp_sdk = get_async_platform_sdk(as_service="models", internal=True) - # Initialize backend registry from configuration - logger.info("Initializing backend registry...") - logger.debug(f"Models backend configs: {backends}") - backend_registry = BackendRegistry.from_config( - nmp_sdk=nmp_sdk, - backend_configs=backends, - huggingface_model_puller=models_config.huggingface_model_puller, - ) - logger.info(f"Backend registry initialized with: {', '.join(backend_registry.list_backends())}") - - # Wait for the models service to be ready before starting control loops (polls /status so we can start once models is ready) - platform_config = get_platform_config() - if not wait_for_service_ready(platform_config, "models", local_stop_signal): - if local_stop_signal.is_set(): - logger.info("Shutdown requested before server became ready") - return - logger.warning("Server did not become ready in time, starting loops anyway") - - # Create the Models Controller - models_controller = ModelsController(backend_registry=backend_registry, stop_signal=local_stop_signal) - models_controller_monitored = TrackLastExecutionTime(models_controller) - - # Create the control loop with configured interval. - # shutdown_func ensures controller resources (event loop, backends) are cleaned up - # inside the loop thread after it exits, avoiding race conditions with step(). - models_controller_loop = Loop( - TimedLoopWaiter(models_config.controller.interval_seconds, stop_signal=local_stop_signal), - models_controller_monitored, - shutdown_func=models_controller.shutdown, - stop_signal=local_stop_signal, - ) - - # Register loop with ControllerManager - controller_manager = ControllerManager.get_instance() - controller_manager.register("models_controller", models_controller_loop) - - # Start the control loop - logger.debug("Starting Models Controller control loop...") - models_controller_loop.start() - logger.debug("Models controller started successfully") - + models_controller_loop: Loop | None = None try: + # Initialize backend registry from configuration + logger.info("Initializing backend registry...") + logger.debug(f"Models backend configs: {backends}") + backend_configs: dict[str, BackendConfig] = {name: config for name, config in backends.items()} + backend_registry = BackendRegistry.from_config( + nmp_sdk=nmp_sdk, + backend_configs=backend_configs, + huggingface_model_puller=models_config.huggingface_model_puller, + ) + logger.info(f"Backend registry initialized with: {', '.join(backend_registry.list_backends())}") + + # Wait for the models service to be ready before starting control loops (polls /status so we can start once models is ready) + platform_config = get_platform_config() + if not wait_for_service_ready(platform_config, "models", local_stop_signal): + if local_stop_signal.is_set(): + logger.info("Shutdown requested before server became ready") + return + logger.warning("Server did not become ready in time, starting loops anyway") + + # Create the Models Controller + models_controller = ModelsController(backend_registry=backend_registry, stop_signal=local_stop_signal) + models_controller_monitored = TrackLastExecutionTime(models_controller) + + # Create the control loop with configured interval. + # shutdown_func ensures controller resources (event loop, backends) are cleaned up + # inside the loop thread after it exits, avoiding race conditions with step(). + models_controller_loop = Loop( + TimedLoopWaiter(models_config.controller.interval_seconds, stop_signal=local_stop_signal), + models_controller_monitored, + shutdown_func=models_controller.shutdown, + stop_signal=local_stop_signal, + ) + + # Register loop with ControllerManager + controller_manager = ControllerManager.get_instance() + controller_manager.register("models_controller", models_controller_loop) + + # Start the control loop + logger.debug("Starting Models Controller control loop...") + models_controller_loop.start() + logger.debug("Models controller started successfully") + # Wait for stop signal or control loop to finish while not local_stop_signal.is_set(): local_stop_signal.wait(timeout=1) @@ -105,15 +108,17 @@ def run(parent_stop_signal: threading.Event | None = None): logger.info("Received keyboard interrupt, stopping models controller") finally: # Tiered shutdown: graceful wait → force cancel → forced cleanup - models_controller_loop.stop() - models_controller_loop.join(timeout=10) - if models_controller_loop.is_alive(): - logger.warning("Models controller step did not finish in 10s, cancelling...") - models_controller.cancel_step() - models_controller_loop.join(timeout=5) - if models_controller_loop.is_alive(): - logger.warning("Models controller loop did not stop, forcing cleanup") - models_controller.shutdown() + if models_controller_loop is not None: + models_controller_loop.stop() + models_controller_loop.join(timeout=10) + if models_controller_loop.is_alive(): + logger.warning("Models controller step did not finish in 10s, cancelling...") + models_controller.cancel_step() + models_controller_loop.join(timeout=5) + if models_controller_loop.is_alive(): + logger.warning("Models controller loop did not stop, forcing cleanup") + models_controller.shutdown() + asyncio.run(nmp_sdk.close()) logger.info("Models controller stopped") 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 4c09e819f8..0f50cc866b 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 @@ -6,14 +6,11 @@ from logging import getLogger from typing import Optional -from nemo_platform import DefaultAsyncHttpxClient # type: ignore[deprecated] +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_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 ModelEntity +from nemo_platform.types.models.model_entity import 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 @@ -64,7 +61,6 @@ def __init__( self._models_sdk = get_async_platform_sdk( as_service="models", internal=True, - http_client=DefaultAsyncHttpxClient(), ) self._service_backends = backend_registry.list_backends() @@ -238,13 +234,10 @@ 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. - models = client_from_platform(self._models_sdk, AsyncModelsClient) - model_entity = ( - await models.get_model( - name=full_model_name, - workspace=workspace, - ) - ).data() + model_entity = await self._models_sdk.models.retrieve( + name=full_model_name, + workspace=workspace, + ) else: model_entity = self._entity_cache.get(workspace, model_name) if model_entity is None: @@ -316,7 +309,9 @@ async def retrieve_non_terminal_deployments(self) -> list[ModelContext]: if deployment.config and deployment.config_version: try: config = await self._retrieve_deployment_config( - deployment.config, deployment.config_version, deployment.workspace + deployment.config, + str(deployment.config_version), + deployment.workspace, ) except Exception as e: logger.warning( @@ -471,9 +466,10 @@ async def async_controller_step(self) -> None: logger.debug(f"Found {len(deployment_contexts)} total deployment(s) in non-terminal states") 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 - } + known_deployment_ids = set() + for ctx in deployment_contexts: + if ctx.model_deployment is not None: + known_deployment_ids.add(f"{ctx.model_deployment.workspace}/{ctx.model_deployment.name}") await self._deployment_reconciler.reconcile_orphans(known_deployment_ids) self.emit_heartbeat() 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 28e9d8e9fc..35fd1e24d0 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 @@ -14,9 +14,8 @@ import urllib.request from nemo_platform import NeMoPlatform, NotFoundError -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 nemo_platform.types.models import ModelEntity +from nemo_platform.types.models.adapter import Adapter from nmp.common.config import get_platform_config from nmp.common.controller import ( Controller, @@ -39,7 +38,11 @@ class AdaptersController(HeartbeatMixin, Controller): - def __init__(self, stop_signal: threading.Event | None = None): + def __init__( + self, + *, + stop_signal: threading.Event | None = None, + ) -> None: self.nim_peft_source = os.getenv("NIM_PEFT_SOURCE", "") if not self.nim_peft_source: msg = "NIM_PEFT_SOURCE is not set on the container" @@ -93,7 +96,6 @@ def __init__(self, stop_signal: threading.Event | None = None): as_service="models", internal=True, ) - self._models = client_from_platform(self._sdk, ModelsClient) def download_fileset(self, dest_dir: str, workspace: str, name: str) -> bool: try: @@ -158,12 +160,12 @@ 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] = list( - self._models.list_models( - workspace=self.workspace, - query_params={"filter": json.dumps({"base_model": self.model_name})}, - ).items() - ) + model_entities = self._sdk.models.list( + workspace=self.workspace, + filter={ + "base_model": self.model_name, + }, + ).data for model_entity in model_entities: if model_entity.prompt: dirs_to_keep.add(model_entity.name) @@ -418,7 +420,7 @@ 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._models.get_model(name=self.model_name, workspace=self.workspace).data() + model_entity: ModelEntity = self._sdk.models.retrieve(name=self.model_name, workspace=self.workspace) if not model_entity.adapters: return @@ -466,7 +468,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: """Run the Adapters Controller with its control loop.""" global adapters_controller_monitored diff --git a/services/core/models/tests/integration/conftest.py b/services/core/models/tests/integration/conftest.py index 0ac087d285..b46ff12a34 100644 --- a/services/core/models/tests/integration/conftest.py +++ b/services/core/models/tests/integration/conftest.py @@ -6,7 +6,7 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, Generator, Optional +from typing import Any, Generator, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -160,12 +160,20 @@ def shutdown(self) -> None: async def create_model_deployment(self, ctx: ModelContext) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.create_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) + deployment = ctx.model_deployment + config = ctx.model_deployment_config + assert deployment is not None + assert config is not None + self.create_calls.append((deployment, config, ctx.model_entity)) return self.create_response async def update_model_deployment(self, ctx: ModelContext) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.update_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) + deployment = ctx.model_deployment + config = ctx.model_deployment_config + assert deployment is not None + assert config is not None + self.update_calls.append((deployment, config, ctx.model_entity)) return self.create_response # Update returns same as create async def get_model_deployment_status(self, ctx: ModelContext) -> DeploymentStatusUpdate: @@ -179,6 +187,7 @@ async def get_model_deployment_status(self, ctx: ModelContext) -> DeploymentStat otherwise falls back to default_status_response. """ deployment = ctx.model_deployment + assert deployment is not None self.status_calls.append(deployment) return self.status_responses.get(deployment.name, self.default_status_response) @@ -243,7 +252,7 @@ def controller_with_mock_backend( Yields: Tuple of (controller, mock_backend, sync_sdk) for testing """ - mock_backend = mock_backend_registry.get_backend() + mock_backend = cast(MockServiceBackend, mock_backend_registry.get_backend()) # Create controller with mock backend registry # We need to patch the SDK factory and platform config (used in config and main modules) @@ -445,11 +454,15 @@ def reconcile_stack(models_controller: ModelsController) -> None: return_value=mock_platform_config, ), patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk") as mock_models_sdk, - patch("nemo_platform_plugin.sdk_provider.get_async_platform_sdk") as mock_sdk, + patch( + "nmp.core.models.controllers.backends.deployments_plugin.backend.get_async_platform_sdk" + ) as mock_backend_sdk, + patch("nemo_deployments_plugin.controller.get_async_platform_sdk") as mock_deployments_sdk, patch("nemo_deployments_plugin.config.DeploymentsConfig.get", return_value=deployments_config), ): mock_models_sdk.return_value = test_clients.async_sdk - mock_sdk.return_value = test_clients.async_sdk + mock_backend_sdk.return_value = test_clients.async_sdk + mock_deployments_sdk.return_value = test_clients.async_sdk models_controller = ModelsController( backend_registry=backend_registry, @@ -479,7 +492,7 @@ def reconcile_stack(models_controller: ModelsController) -> None: @pytest.hookimpl(tryfirst=True, hookwrapper=True) -def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, None, None]: +def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, Any, None]: """Store test results on the item for fixture access.""" outcome = yield rep = outcome.get_result() diff --git a/services/core/models/tests/unit/api/test_models_api.py b/services/core/models/tests/unit/api/test_models_api.py index 83997c3ced..077baed3b1 100644 --- a/services/core/models/tests/unit/api/test_models_api.py +++ b/services/core/models/tests/unit/api/test_models_api.py @@ -412,18 +412,16 @@ def test_create_model_entity_validation_error_returns_422(client, mock_model_ent @pytest.mark.asyncio -async def test_model_spec_job_transport_failure_does_not_fail_persisted_model(sample_model_entity): +async def test_start_update_model_spec_job_swallows_nemo_transport_error(sample_model_entity): request = httpx.Request("POST", "http://test/apis/jobs/v2/workspaces/nvidia/jobs") + sdk = MagicMock() jobs = MagicMock() jobs.create_job = AsyncMock( side_effect=NemoTransportError(httpx.ConnectError("Connection refused", request=request)) ) - with ( - patch("nmp.core.models.api.v2.models.get_async_platform_sdk"), - patch("nmp.core.models.api.v2.models.client_from_platform", return_value=jobs), - ): - await start_update_model_spec_job(sample_model_entity) + with patch("nmp.core.models.api.v2.models.client_from_platform", return_value=jobs): + await start_update_model_spec_job(sample_model_entity, sdk) jobs.create_job.assert_awaited_once() diff --git a/services/core/models/tests/unit/controllers/conftest.py b/services/core/models/tests/unit/controllers/conftest.py index 4b11091713..9023bbd9d5 100644 --- a/services/core/models/tests/unit/controllers/conftest.py +++ b/services/core/models/tests/unit/controllers/conftest.py @@ -155,12 +155,12 @@ def _assert_controller_healthy(controller, is_healthy=True): def _assert_sdk_initialized_correctly(mock_sdk_class_patch): """Assert that get_async_platform_sdk was called with correct args for Models API.""" - # Controller initializes ONE SDK for Models API (base_url is resolved from config inside the factory) + # Controller initializes one SDK and lets the factory own endpoint/client selection. assert mock_sdk_class_patch.call_count == 1 call_kwargs = mock_sdk_class_patch.call_args.kwargs assert call_kwargs["as_service"] == "models" assert call_kwargs["internal"] is True - assert "http_client" in call_kwargs + assert "http_client" not in call_kwargs def _assert_asyncio_run_called_once(mock_asyncio_run_patch): 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 1c7e5d1fc0..e60d398f42 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 @@ -12,7 +12,7 @@ from nmp.core.models.controllers.context import ModelContext from nmp.core.models.controllers.models_controller import NON_TERMINAL_STATES, ModelsController -from .conftest import _ModelResponse, make_async_models_client +from .conftest import make_async_models_client class MockAsyncPaginator: @@ -32,20 +32,12 @@ async def __anext__(self): @pytest.fixture(autouse=True) def _patch_typed_model_client(mock_models_sdk): - """Route ``client_from_platform(sdk, AsyncModelsClient)`` in the models controller - and entity cache back to a typed ``AsyncModelsClient`` mock on - ``mock_models_sdk.models_client``, so tests drive ``get_model``/``list_models`` - directly instead of the legacy SDK resource.""" + """Route entity-cache typed client calls while controller tests use generated SDK resources directly.""" mock_models_sdk.models_client = make_async_models_client() - with ( - patch( - "nmp.core.models.controllers.models_controller.client_from_platform", - side_effect=lambda sdk, cls: sdk.models_client, - ), - patch( - "nmp.core.models.controllers.entity_cache.client_from_platform", - side_effect=lambda sdk, cls: sdk.models_client, - ), + mock_models_sdk.models.retrieve = AsyncMock() + with patch( + "nmp.core.models.controllers.entity_cache.client_from_platform", + side_effect=lambda sdk, cls: sdk.models_client, ): yield @@ -467,7 +459,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_client.get_model = AsyncMock(return_value=_ModelResponse(mock_entity)) + mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) config = MagicMock() config.model_entity_id = "my-ws/my-model" @@ -481,7 +473,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_client.get_model.assert_awaited_once_with(name="my-model", workspace="my-ws") + mock_models_sdk.models.retrieve.assert_awaited_once_with(name="my-model", workspace="my-ws") @pytest.mark.asyncio @@ -490,7 +482,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_client.get_model = AsyncMock(return_value=_ModelResponse(mock_entity)) + mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) config = MagicMock() config.model_entity_id = "my-ws/my-model@v2" @@ -501,8 +493,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_client.get_model.assert_awaited_once_with(name="my-model@v2", workspace="my-ws") - call = mock_models_sdk.models_client.get_model.await_args + mock_models_sdk.models.retrieve.assert_awaited_once_with(name="my-model@v2", workspace="my-ws") + call = mock_models_sdk.models.retrieve.await_args assert call is not None call_kw = call.kwargs assert call_kw["name"] == "my-model@v2" @@ -515,7 +507,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_client.get_model = AsyncMock(return_value=_ModelResponse(mock_entity)) + mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) config = MagicMock() config.model_entity_id = None @@ -529,7 +521,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_client.get_model.assert_awaited_once_with(name="nim-model@v1", workspace="nim-ns") + mock_models_sdk.models.retrieve.assert_awaited_once_with(name="nim-model@v1", workspace="nim-ns") @pytest.mark.asyncio @@ -546,7 +538,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_client.get_model.assert_not_awaited() + mock_models_sdk.models.retrieve.assert_not_awaited() @pytest.mark.asyncio @@ -555,7 +547,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_client.get_model = AsyncMock(return_value=_ModelResponse(mock_entity)) + mock_models_sdk.models.retrieve = AsyncMock(return_value=mock_entity) config = MagicMock() config.model_entity_id = "bogus" @@ -569,7 +561,7 @@ 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_client.get_model.assert_awaited_once_with(name="fallback-model", workspace="fallback-ns") + mock_models_sdk.models.retrieve.assert_awaited_once_with(name="fallback-model", workspace="fallback-ns") # ============================================================================= 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 5ebe9016a5..3d4079181b 100644 --- a/services/core/models/tests/unit/sidecars/test_adapters_controller.py +++ b/services/core/models/tests/unit/sidecars/test_adapters_controller.py @@ -43,7 +43,6 @@ def controller(tmp_path): ctrl = AdaptersController() ctrl.nim_peft_source = str(tmp_path) ctrl._sdk = MagicMock() - ctrl._models = MagicMock() ctrl.workspace = "default" ctrl.model_name = "base-model" # Default to NIM behavior (no rewrite, no eager vLLM load); vLLM tests @@ -400,7 +399,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -432,7 +431,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity dirs_to_keep: set[str] = set() controller._update_lora_adapters(dirs_to_keep) @@ -449,7 +448,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -478,7 +477,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -499,7 +498,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [] @@ -528,7 +527,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [] @@ -557,7 +556,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -579,7 +578,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -614,7 +613,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] @@ -642,14 +641,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._models.get_model.return_value.data.return_value = mock_model_entity + 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 # No prompt-tuned models in this scenario. - controller._models.list_models.return_value.items.return_value = iter([]) + controller._sdk.models.list.return_value = MagicMock(data=[]) controller.step() @@ -671,7 +670,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._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity dirs_to_keep: set[str] = set() controller._update_lora_adapters(dirs_to_keep) @@ -693,7 +692,7 @@ def _model_entity(self, controller, adapter): me = MagicMock() me.workspace = "default" me.adapters = [adapter] - controller._models.get_model.return_value.data.return_value = me + controller._sdk.models.retrieve.return_value = me files_resp = MagicMock() files_resp.data = [MagicMock()] controller._sdk.files.list.return_value = files_resp @@ -824,7 +823,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._models.list_models.return_value.items.return_value = iter([]) # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models with patch.object(controller, "_vllm_api_call", return_value=(200, "")) as api: controller.step() @@ -844,7 +843,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._models.list_models.return_value.items.return_value = iter([]) # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models # vLLM unreachable: both the kept adapter's load and the stale one's unload # hit a transport error. @@ -865,7 +864,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._models.list_models.return_value.items.return_value = iter([]) # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": @@ -888,7 +887,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._models.list_models.return_value.items.return_value = iter([]) # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": @@ -986,7 +985,7 @@ def __init__(self, name: str, fileset: str): mock_model_entity = MagicMock() mock_model_entity.workspace = "base-ws" mock_model_entity.adapters = [adapter] - controller._models.get_model.return_value.data.return_value = mock_model_entity + controller._sdk.models.retrieve.return_value = mock_model_entity mock_files_response = MagicMock() mock_files_response.data = [MagicMock()] diff --git a/services/guardrails/src/nmp/guardrails/api/dependencies.py b/services/guardrails/src/nmp/guardrails/api/dependencies.py index cfdac1eaa2..9a13dfe3b0 100644 --- a/services/guardrails/src/nmp/guardrails/api/dependencies.py +++ b/services/guardrails/src/nmp/guardrails/api/dependencies.py @@ -10,7 +10,6 @@ from fastapi import Depends from nemo_platform import AsyncNeMoPlatform from nmp.common.entities.client import EntityClient -from nmp.common.http_clients import shared_async_http_client from nmp.common.service.dependencies import get_entity_client from nmp.guardrails.app.services.configs.registry import ConfigRegistry from nmp.guardrails.app.services.rails.registry import RailsRegistry @@ -47,17 +46,41 @@ def get_rails_service( return RailsService(config_registry=config_registry, rails_registry=rails_registry) -# Dependency for NeMo Platform -def get_nemo_platform() -> AsyncNeMoPlatform: +def _guardrails_inference_base_url() -> str: nim_endpoint_url = settings.nim_endpoint_settings.base_url - # Remove the /v1 from the end of the URL if it exists - # This is necessary because the NeMo Platform API SDK expects the base URL to not have the /v1 suffix unlike OpenAI SDK + # Remove the /v1 from the end of the URL if it exists. + # The platform SDK base URL does not include the OpenAI-compatible suffix. if nim_endpoint_url.endswith("/v1"): - nim_endpoint_url = nim_endpoint_url[: -len("/v1")] - return AsyncNeMoPlatform( - inference_base_url=nim_endpoint_url, - http_client=shared_async_http_client(), - ) + return nim_endpoint_url[: -len("/v1")] + return nim_endpoint_url + + +class _NeMoPlatformClientProvider: + def __init__(self) -> None: + self._sdk: AsyncNeMoPlatform | None = None + + def get(self) -> AsyncNeMoPlatform: + if self._sdk is None: + self._sdk = AsyncNeMoPlatform(inference_base_url=_guardrails_inference_base_url()) + return self._sdk + + async def aclose(self) -> None: + sdk = self._sdk + self._sdk = None + if sdk is not None: + await sdk.close() + + +_nemo_platform_client_provider = _NeMoPlatformClientProvider() + + +# Dependency for NeMo Platform +def get_nemo_platform() -> AsyncNeMoPlatform: + return _nemo_platform_client_provider.get() + + +async def close_nemo_platform() -> None: + await _nemo_platform_client_provider.aclose() RailsServiceDep = Annotated[RailsService, Depends(get_rails_service)] 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 4d1d980dc1..f8d86856ff 100644 --- a/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py +++ b/services/guardrails/src/nmp/guardrails/app/utils/model_routing.py @@ -61,10 +61,11 @@ 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() - models = client_from_platform(sdk, ModelsClient) - url = models.get_openai_route_base_url(workspace=workspace) - - return url + try: + models = client_from_platform(sdk, ModelsClient) + return models.get_openai_route_base_url(workspace=workspace) + finally: + sdk.close() def resolve_model_entity_references(rails_config: RailsConfig) -> RailsConfig: diff --git a/services/guardrails/src/nmp/guardrails/service.py b/services/guardrails/src/nmp/guardrails/service.py index ae7dff1383..a994c46f73 100644 --- a/services/guardrails/src/nmp/guardrails/service.py +++ b/services/guardrails/src/nmp/guardrails/service.py @@ -5,7 +5,7 @@ import logging import os -from typing import ClassVar, List +from typing import ClassVar, List, cast from fastapi import FastAPI from fastapi.exceptions import RequestValidationError @@ -28,6 +28,7 @@ from nmp.guardrails.app.patches import apply_langchain_patch from nmp.guardrails.config import GuardrailsServiceConfig from pydantic import ValidationError +from starlette.types import ExceptionHandler logger = logging.getLogger(__name__) @@ -70,6 +71,12 @@ async def on_startup(self) -> None: os.environ["NVIDIA_BASE_URL"] = inference_base_url logger.info(f"Set NVIDIA_BASE_URL to: {inference_base_url}") + async def on_shutdown(self) -> None: + from nmp.guardrails.api.dependencies import close_nemo_platform + + await close_nemo_platform() + await super().on_shutdown() + def configure_app(self, app: FastAPI) -> None: """Configure additional app settings after creation.""" # Add middlewares @@ -90,11 +97,15 @@ def _configure_middlewares(self, app: FastAPI) -> None: def _register_exception_handlers(self, app: FastAPI) -> None: """Register custom exception handlers.""" - app.add_exception_handler(GuardrailConfigurationNotFoundError, config_not_found_error_handler) - app.add_exception_handler(CustomHTTPException, custom_exception_handler) - app.add_exception_handler(LLMCallException, llm_call_exception_handler) - app.add_exception_handler(ModelInitializationError, model_initialization_error_handler) - app.add_exception_handler(InvalidRailsConfigurationError, invalid_rails_configuration_error_handler) - app.add_exception_handler(RequestValidationError, validation_error_handler) - app.add_exception_handler(ValidationError, validation_error_handler) - app.add_exception_handler(404, custom_404_handler) + app.add_exception_handler( + GuardrailConfigurationNotFoundError, cast(ExceptionHandler, config_not_found_error_handler) + ) + app.add_exception_handler(CustomHTTPException, cast(ExceptionHandler, custom_exception_handler)) + app.add_exception_handler(LLMCallException, cast(ExceptionHandler, llm_call_exception_handler)) + app.add_exception_handler(ModelInitializationError, cast(ExceptionHandler, model_initialization_error_handler)) + app.add_exception_handler( + InvalidRailsConfigurationError, cast(ExceptionHandler, invalid_rails_configuration_error_handler) + ) + app.add_exception_handler(RequestValidationError, cast(ExceptionHandler, validation_error_handler)) + app.add_exception_handler(ValidationError, cast(ExceptionHandler, validation_error_handler)) + app.add_exception_handler(404, cast(ExceptionHandler, custom_404_handler)) diff --git a/services/guardrails/tests/utils/test_model_routing.py b/services/guardrails/tests/utils/test_model_routing.py index f6920f4b17..1630ba8dd7 100644 --- a/services/guardrails/tests/utils/test_model_routing.py +++ b/services/guardrails/tests/utils/test_model_routing.py @@ -1,15 +1,28 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -from unittest.mock import MagicMock, patch +from unittest.mock import patch +import httpx import pytest +from nemo_platform import NeMoPlatform from nmp.guardrails.app.utils.model_routing import ( build_openai_gateway_url, parse_model_entity_reference, resolve_model_entity_references, ) -from nmp.guardrails.entities.values._private import Model, RailsConfig +from nmp.guardrails.entities.values._private import Model, ModelParameters, RailsConfig + + +def _model_routing_sdk(base_url: str = "http://localhost:8000") -> NeMoPlatform: + http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + return NeMoPlatform(base_url=base_url, workspace="default", http_client=http_client) + + +def _base_url(model: Model) -> str | None: + parameters = model.parameters + assert parameters is not None + return parameters.base_url class TestParseModelEntityReference: @@ -51,9 +64,7 @@ class TestBuildOpenAIGatewayUrl: @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") def test_url_construction(self, mock_get_sdk): """Test URL construction for Model Entity reference.""" - mock_sdk = MagicMock() - mock_sdk.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk url = build_openai_gateway_url("default/my-model") assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" @@ -65,9 +76,7 @@ def test_url_construction(self, mock_get_sdk): @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 typed client URL.""" - mock_sdk = MagicMock() - mock_sdk.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk url = build_openai_gateway_url("default/model") assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" @@ -76,9 +85,7 @@ def test_v1_suffix_preserved(self, mock_get_sdk): @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") def test_typed_client_adds_v1(self, mock_get_sdk): """Test the typed Models client helper adds the OpenAI /v1 suffix.""" - mock_sdk = MagicMock() - mock_sdk.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk url = build_openai_gateway_url("default/model") assert url == "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" @@ -95,9 +102,7 @@ class TestResolveModelEntityReferences: @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") def test_resolve_single_model(self, mock_get_sdk): """Test resolving a single model with Model Entity reference.""" - mock_sdk = MagicMock() - mock_sdk.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk rails_config = RailsConfig( models=[ @@ -107,16 +112,14 @@ def test_resolve_single_model(self, mock_get_sdk): resolved = resolve_model_entity_references(rails_config) - assert resolved.models[0].parameters["base_url"] == ( + assert _base_url(resolved.models[0]) == ( "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" ) @patch("nmp.guardrails.app.utils.model_routing.get_platform_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.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk rails_config = RailsConfig( models=[ @@ -129,9 +132,9 @@ def test_resolve_all_models(self, mock_get_sdk): resolved = resolve_model_entity_references(rails_config) expected_url = "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" - assert resolved.models[0].parameters["base_url"] == expected_url - assert resolved.models[1].parameters["base_url"] == expected_url - assert resolved.models[2].parameters["base_url"] == expected_url + assert _base_url(resolved.models[0]) == expected_url + assert _base_url(resolved.models[1]) == expected_url + assert _base_url(resolved.models[2]) == expected_url def test_explicit_base_url_preserved_for_non_entity(self): """Test that explicit base_url is preserved for non-entity models.""" @@ -141,14 +144,14 @@ def test_explicit_base_url_preserved_for_non_entity(self): type="main", engine="nim", model="gpt-4", - parameters={"base_url": "http://custom-endpoint/v1"}, + parameters=ModelParameters(base_url="http://custom-endpoint/v1"), ), ] ) resolved = resolve_model_entity_references(rails_config) - assert resolved.models[0].parameters["base_url"] == "http://custom-endpoint/v1" + assert _base_url(resolved.models[0]) == "http://custom-endpoint/v1" def test_explicit_base_url_preserved_for_entity_reference(self): """Test that explicit base_url is preserved even for Model Entity references.""" @@ -158,21 +161,19 @@ def test_explicit_base_url_preserved_for_entity_reference(self): type="main", engine="nim", model="default/my-model", - parameters={"base_url": "http://custom-override/v1"}, + parameters=ModelParameters(base_url="http://custom-override/v1"), ), ] ) resolved = resolve_model_entity_references(rails_config) - assert resolved.models[0].parameters["base_url"] == "http://custom-override/v1" + assert _base_url(resolved.models[0]) == "http://custom-override/v1" @patch("nmp.guardrails.app.utils.model_routing.get_platform_sdk") def test_mixed_config(self, mock_get_sdk): """Test config with one Model Entity ref, one explicit URLs.""" - mock_sdk = MagicMock() - mock_sdk.base_url = "http://localhost:8000" - mock_get_sdk.return_value = mock_sdk + mock_get_sdk.side_effect = _model_routing_sdk rails_config = RailsConfig( models=[ @@ -181,17 +182,17 @@ def test_mixed_config(self, mock_get_sdk): type="content_safety", engine="nim", model="default/nemoguard-content-safety", - parameters={"base_url": "https://direct-nim:8000/v1"}, + parameters=ModelParameters(base_url="https://direct-nim:8000/v1"), ), ] ) resolved = resolve_model_entity_references(rails_config) - assert resolved.models[0].parameters["base_url"] == ( + assert _base_url(resolved.models[0]) == ( "http://localhost:8000/apis/inference-gateway/v2/workspaces/default/openai/-/v1" ) - assert resolved.models[1].parameters["base_url"] == "https://direct-nim:8000/v1" + assert _base_url(resolved.models[1]) == "https://direct-nim:8000/v1" def test_empty_models_returns_unchanged(self): """Test that config with empty models list is returned unchanged.""" diff --git a/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py b/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py index 856db6f4a1..4d600e2268 100644 --- a/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py +++ b/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py @@ -10,6 +10,7 @@ import respx from nemo_platform import NeMoPlatform from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.oidc import NMPOIDCConfig from nmp.common.config import Configuration, PlatformConfig from nmp.common.jobs.constants import TASK_CONFIG_ENVVAR from nmp.hello_world.tasks.workload_workspace_get.run import run as task_run @@ -34,6 +35,17 @@ def platform_base_url(): Configuration.clear_override(PlatformConfig) +def _workload_oidc_config() -> NMPOIDCConfig: + return NMPOIDCConfig( + auth_enabled=True, + workload_token_exchange_enabled=True, + workload_client_id="workload-client", + workload_token_endpoint="https://idp.example.test/oauth2/token", + workload_audience="nemo-platform", + workload_scope="openid email groups", + ) + + class _StubWorkspaces: """Recording fake matching the typed WorkspacesClient shape (get_workspace + .data()).""" @@ -137,3 +149,59 @@ def workspace_response(request: httpx.Request) -> httpx.Response: assert exit_code == 0 assert workspace_route.called + + +@respx.mock +def test_workload_workspace_get_uses_task_sdk_with_workload_token(monkeypatch, tmp_path, platform_base_url): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "workload-access-token", "expires_in": 300} + + monkeypatch.setenv(TASK_CONFIG_ENVVAR, '{"workspace":"workload-read-target"}') + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", platform_base_url) + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.delenv("NEMO_WORKLOAD_TOKEN", raising=False) + monkeypatch.delenv("NEMO_WORKLOAD_TOKEN_FILE", raising=False) + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory._discover_oidc_client_settings", + lambda _base_url: _workload_oidc_config(), + ) + monkeypatch.setattr("nemo_platform_plugin.client.oidc.token_exchange_grant", token_exchange_grant) + + def workspace_response(request: httpx.Request) -> httpx.Response: + if request.headers.get("Authorization") != "Bearer workload-access-token": + return httpx.Response(401, json={"detail": "Unauthorized"}) + if request.headers.get("X-NMP-Internal") != "true": + return httpx.Response(403, json={"detail": "Forbidden"}) + if "X-NMP-Principal-Id" in request.headers: + return httpx.Response(403, json={"detail": "Trusted principal header was sent"}) + if "X-NMP-Principal-On-Behalf-Of" in request.headers: + return httpx.Response(403, json={"detail": "Trusted on-behalf-of header was sent"}) + return httpx.Response( + 200, + json={ + "id": "workspace-id", + "name": "workload-read-target", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-01T00:00:00Z", + }, + ) + + workspace_route = respx.get(f"{platform_base_url}/apis/entities/v2/workspaces/workload-read-target").mock( + side_effect=workspace_response + ) + + exit_code = task_run() + + assert exit_code == 0 + assert workspace_route.called + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" diff --git a/services/intake/tests/test_clickhouse_startup.py b/services/intake/tests/test_clickhouse_startup.py index 963141d1b2..9897cdf32e 100644 --- a/services/intake/tests/test_clickhouse_startup.py +++ b/services/intake/tests/test_clickhouse_startup.py @@ -32,6 +32,7 @@ def test_intake_ready_with_explicit_external_clickhouse_without_startup_warning( service = IntakeService().with_config(intake_config) async def check_readiness() -> bool: + service.dependency_provider.initialize() await service.on_startup() assert service.clickhouse_client is not None try: @@ -54,6 +55,7 @@ def test_intake_uses_reconciled_clickhouse_url(monkeypatch: pytest.MonkeyPatch) service = IntakeService().with_config(intake_config) async def start_and_stop() -> None: + service.dependency_provider.initialize() await service.on_startup() try: assert service.clickhouse_client is not None @@ -86,6 +88,7 @@ def test_intake_stays_ready_and_logs_docker_recovery_guidance( service = IntakeService().with_config(IntakeConfig(clickhouse_config=ClickHouseConfig())) async def check_readiness() -> bool: + service.dependency_provider.initialize() await service.on_startup() try: return await service.is_ready() @@ -107,6 +110,7 @@ def test_intake_stays_ready_for_non_docker_reconciliation_errors( service = IntakeService().with_config(IntakeConfig(clickhouse_config=ClickHouseConfig())) async def check_readiness() -> bool: + service.dependency_provider.initialize() await service.on_startup() try: assert service.clickhouse_client is not None diff --git a/services/studio/src/nmp/studio/service.py b/services/studio/src/nmp/studio/service.py index 9267770ad2..1812429e75 100644 --- a/services/studio/src/nmp/studio/service.py +++ b/services/studio/src/nmp/studio/service.py @@ -11,10 +11,10 @@ from pathlib import Path from typing import ClassVar +import httpx from fastapi import FastAPI, HTTPException, Request, status from fastapi.responses import FileResponse, HTMLResponse from nemo_platform_plugin.authz import Permission -from nmp.common.http_clients import shared_async_http_client from nmp.common.service import RouterConfig, Service from nmp.studio import assistant from nmp.studio.config import StudioConfig @@ -89,9 +89,17 @@ class StudioService(Service[StudioConfig]): dependencies: ClassVar[list[str]] = ["entities", "auth"] - def __init__(self): + def __init__(self, telemetry_http_client: httpx.AsyncClient | None = None): """Initialize the studio service.""" super().__init__(name="studio", module_name="nmp.studio") + self._owns_telemetry_http_client = telemetry_http_client is None + self._telemetry_http_client = telemetry_http_client or httpx.AsyncClient() + + async def on_shutdown(self) -> None: + """Close service-owned telemetry proxy resources.""" + if self._owns_telemetry_http_client: + await self._telemetry_http_client.aclose() + await super().on_shutdown() @property def title(self) -> str: @@ -189,7 +197,7 @@ async def _proxy_telemetry(self, request: Request, telemetry_path: str = "") -> target_url = self._build_telemetry_target_url(collector_url, telemetry_path, request.url.query) try: - upstream_response = await shared_async_http_client().request( + upstream_response = await self._telemetry_http_client.request( method=request.method, url=target_url, content=await request.body(), diff --git a/services/studio/tests/unit/test_service.py b/services/studio/tests/unit/test_service.py index 2e9d6237c8..0f1debbb29 100644 --- a/services/studio/tests/unit/test_service.py +++ b/services/studio/tests/unit/test_service.py @@ -7,8 +7,10 @@ from pathlib import Path from types import SimpleNamespace +from typing import cast from unittest.mock import patch +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient @@ -32,11 +34,21 @@ class FakeTelemetryClient: def __init__(self, response: FakeTelemetryResponse | None = None): self.response = response or FakeTelemetryResponse() self.calls: list[dict] = [] + self.close_calls = 0 + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc_value, traceback) -> None: + pass async def request(self, **kwargs): self.calls.append(kwargs) return self.response + async def aclose(self) -> None: + self.close_calls += 1 + class TestStudioService: """Tests for the StudioService class.""" @@ -75,24 +87,23 @@ class TestTelemetryProxy: def _client( self, config: StudioConfig, - monkeypatch: pytest.MonkeyPatch, fake_client: FakeTelemetryClient | None = None, ) -> tuple[TestClient, FakeTelemetryClient]: app = FastAPI() telemetry_client = fake_client or FakeTelemetryClient() - monkeypatch.setattr("nmp.studio.service.shared_async_http_client", lambda: telemetry_client) - StudioService().with_config(config).configure_app(app) + StudioService(telemetry_http_client=cast(httpx.AsyncClient, telemetry_client)).with_config( + config + ).configure_app(app) return TestClient(app), telemetry_client - def test_post_strips_studio_telemetry_prefix_and_proxies_request(self, monkeypatch: pytest.MonkeyPatch): + def test_post_strips_studio_telemetry_prefix_and_proxies_request(self): """Test that /studio/telemetry/* proxies to the collector without the route prefix.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post( @@ -113,15 +124,14 @@ def test_post_strips_studio_telemetry_prefix_and_proxies_request(self, monkeypat assert call["headers"]["X-Real-IP"] == "testclient" assert call["headers"]["X-Forwarded-For"] == "testclient" - def test_post_only_forwards_whitelisted_telemetry_headers(self, monkeypatch: pytest.MonkeyPatch): + def test_post_only_forwards_whitelisted_telemetry_headers(self): """Test that browser credentials and metadata are not forwarded to the collector.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post( @@ -151,15 +161,14 @@ def test_post_only_forwards_whitelisted_telemetry_headers(self, monkeypatch: pyt "X-Forwarded-For": "testclient", } - def test_post_strips_root_telemetry_prefix_and_proxies_request(self, monkeypatch: pytest.MonkeyPatch): + def test_post_strips_root_telemetry_prefix_and_proxies_request(self): """Test that /telemetry/* keeps parity with the old nginx route.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post("/telemetry/v1/logs", headers={"origin": origin}) @@ -167,15 +176,14 @@ def test_post_strips_root_telemetry_prefix_and_proxies_request(self, monkeypatch assert response.status_code == 200 assert telemetry_client.calls[0]["url"] == "http://collector:4318/v1/logs" - def test_options_returns_preflight_response_without_proxying(self, monkeypatch: pytest.MonkeyPatch): + def test_options_returns_preflight_response_without_proxying(self): """Test that CORS preflight requests are handled locally.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.options("/studio/telemetry/v1/traces", headers={"origin": origin}) @@ -188,11 +196,10 @@ def test_options_returns_preflight_response_without_proxying(self, monkeypatch: assert response.headers["access-control-max-age"] == "1728000" assert telemetry_client.calls == [] - def test_disabled_telemetry_returns_404(self, monkeypatch: pytest.MonkeyPatch): + def test_disabled_telemetry_returns_404(self): """Test that disabled telemetry preserves the old nginx 404 behavior.""" client, telemetry_client = self._client( - StudioConfig(telemetry_enabled=False, otel={"collector_url": "http://collector:4318"}), - monkeypatch, + StudioConfig(telemetry_enabled=False, otel={"collector_url": "http://collector:4318"}) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://testserver"}) @@ -200,14 +207,13 @@ def test_disabled_telemetry_returns_404(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 404 assert telemetry_client.calls == [] - def test_disallowed_origin_returns_403(self, monkeypatch: pytest.MonkeyPatch): + def test_disallowed_origin_returns_403(self): """Test that disallowed origins preserve the old nginx 403 behavior.""" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": ["http://studio.test"]}, - ), - monkeypatch, + ) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://not-allowed.test"}) @@ -215,14 +221,13 @@ def test_disallowed_origin_returns_403(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 403 assert telemetry_client.calls == [] - def test_same_origin_request_is_allowed(self, monkeypatch: pytest.MonkeyPatch): + def test_same_origin_request_is_allowed(self): """Test that same-origin Studio deployments work without hard-coded host config.""" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": []}, - ), - monkeypatch, + ) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://testserver"}) @@ -230,6 +235,59 @@ def test_same_origin_request_is_allowed(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 200 assert telemetry_client.calls[0]["url"] == "http://collector:4318/v1/traces" + def test_post_reuses_service_owned_telemetry_client(self, monkeypatch: pytest.MonkeyPatch): + """Test that proxied telemetry requests reuse one service-scoped client.""" + origin = "http://studio.test" + app = FastAPI() + created_clients: list[FakeTelemetryClient] = [] + + def create_client() -> FakeTelemetryClient: + client = FakeTelemetryClient() + created_clients.append(client) + return client + + monkeypatch.setattr("nmp.studio.service.httpx.AsyncClient", create_client) + StudioService().with_config( + StudioConfig( + telemetry_enabled=True, + otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, + ) + ).configure_app(app) + client = TestClient(app) + + first_response = client.post("/studio/telemetry/v1/traces", headers={"origin": origin}) + second_response = client.post("/studio/telemetry/v1/logs", headers={"origin": origin}) + + assert first_response.status_code == 200 + assert second_response.status_code == 200 + assert len(created_clients) == 1 + assert [call["url"] for call in created_clients[0].calls] == [ + "http://collector:4318/v1/traces", + "http://collector:4318/v1/logs", + ] + + @pytest.mark.asyncio + async def test_shutdown_does_not_close_injected_telemetry_client(self): + """Test that injected telemetry HTTP clients remain caller-owned.""" + telemetry_client = FakeTelemetryClient() + service = StudioService(telemetry_http_client=cast(httpx.AsyncClient, telemetry_client)) + + await service.on_shutdown() + + assert telemetry_client.close_calls == 0 + + @pytest.mark.asyncio + async def test_shutdown_closes_owned_telemetry_client(self): + """Test that the service closes telemetry clients it creates.""" + telemetry_client = FakeTelemetryClient() + service = StudioService() + service._telemetry_http_client = cast(httpx.AsyncClient, telemetry_client) + service._owns_telemetry_http_client = True + + await service.on_shutdown() + + assert telemetry_client.close_calls == 1 + class TestStaticFilesPath: """Tests for static_files_path configuration.""" diff --git a/tests/auth_idp/authentik_live.py b/tests/auth_idp/authentik_live.py index 2376918245..1aeae4e5a5 100644 --- a/tests/auth_idp/authentik_live.py +++ b/tests/auth_idp/authentik_live.py @@ -12,7 +12,7 @@ from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import rsa from cryptography.x509.oid import ExtendedKeyUsageOID, NameOID -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR REPO_ROOT = Path(__file__).resolve().parents[2] AUTHENTIK_ROOT = REPO_ROOT / "contrib/auth/authentik" diff --git a/tests/auth_idp/compose/test_authentik_cli_login.py b/tests/auth_idp/compose/test_authentik_cli_login.py index a323dfe884..e3063a6139 100644 --- a/tests/auth_idp/compose/test_authentik_cli_login.py +++ b/tests/auth_idp/compose/test_authentik_cli_login.py @@ -4,7 +4,7 @@ import httpx import pytest from nemo_platform_ext.auth.helpers import discover_nmp_config -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from tests.auth_idp.authentik_live import AUTHENTIK_DOCKER_E2E_CONFIG @@ -33,7 +33,7 @@ def test_authentik_discovery_exposes_gateway_reachable_device_flow(authentik_sta "scope": oidc.default_scopes, }, timeout=30.0, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) response.raise_for_status() body = response.json() @@ -55,7 +55,7 @@ def test_authentik_cli_provider_rejects_unseeded_human_app_password(authentik_st "scope": "openid email offline_access groups", }, timeout=30.0, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) assert token_response.status_code == 400 diff --git a/tests/auth_idp/conftest.py b/tests/auth_idp/conftest.py index 5193750f09..29af33712e 100644 --- a/tests/auth_idp/conftest.py +++ b/tests/auth_idp/conftest.py @@ -13,8 +13,8 @@ import httpx import pytest from nemo_platform import NeMoPlatform -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, client_verify_from_env from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, httpx_tls_config_from_env from nemo_platform_plugin.workspaces.client import WorkspacesClient from nemo_platform_plugin.workspaces.types import CreateWorkspaceQueryParams, CreateWorkspaceRequest @@ -152,14 +152,14 @@ def _exchange_token_with_retries( ) -> str: deadline = time.monotonic() + timeout last_error: Exception | None = None - request_verify = client_verify_from_env() if verify is None else verify + request_kwargs = httpx_tls_config_from_env() if verify is None else {"verify": verify} while time.monotonic() < deadline: try: response = httpx.post( token_endpoint, data=_token_request_body(grant), timeout=30.0, - verify=request_verify, + **request_kwargs, ) if response.status_code >= 500: last_error = httpx.HTTPStatusError( diff --git a/tests/auth_idp/contracts/test_access_keys.py b/tests/auth_idp/contracts/test_access_keys.py index 5651962e77..a3500ef64c 100644 --- a/tests/auth_idp/contracts/test_access_keys.py +++ b/tests/auth_idp/contracts/test_access_keys.py @@ -6,7 +6,7 @@ import httpx import pytest -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from nmp.testing import grant_workspace_role from tests.auth_idp.common import jwt_claims, require_capability @@ -24,7 +24,7 @@ def _runtime_verify(auth_idp_runtime) -> str | bool: - return getattr(auth_idp_runtime, "verify", client_verify_from_env()) + return getattr(auth_idp_runtime, "verify", httpx_tls_config_from_env().get("verify", True)) def _create_access_key_with_body(auth_idp_runtime, bearer_token: str, body: dict[str, object]) -> dict[str, object]: diff --git a/tests/auth_idp/contracts/test_cli_refresh.py b/tests/auth_idp/contracts/test_cli_refresh.py index 014badb944..4b026636f8 100644 --- a/tests/auth_idp/contracts/test_cli_refresh.py +++ b/tests/auth_idp/contracts/test_cli_refresh.py @@ -9,8 +9,8 @@ import yaml from nemo_platform_ext.auth.helpers import decode_jwt_claims, discover_nmp_config, generate_unsigned_jwt from nemo_platform_ext.cli.app import app -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, client_verify_from_env from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR, httpx_tls_config_from_env from typer.testing import CliRunner from tests.auth_idp.common import require_capability @@ -92,7 +92,7 @@ def test_cli_api_command_auto_refreshes_expired_device_flow_token( assert oidc.token_endpoint assert "offline_access" in oidc.default_scopes.split() - verify = getattr(auth_idp_runtime, "verify", client_verify_from_env()) + verify = getattr(auth_idp_runtime, "verify", httpx_tls_config_from_env().get("verify", True)) runtime_device_authorization_endpoint = with_url_origin( oidc.device_authorization_endpoint, auth_idp_runtime.gateway_base_url, diff --git a/tests/auth_idp/contracts/test_discovery.py b/tests/auth_idp/contracts/test_discovery.py index 1cbb1ceb5d..b5b5edd3eb 100644 --- a/tests/auth_idp/contracts/test_discovery.py +++ b/tests/auth_idp/contracts/test_discovery.py @@ -6,7 +6,7 @@ import httpx import pytest from nemo_platform_ext.auth.helpers import discover_nmp_config -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from tests.auth_idp.common import require_capability from tests.auth_idp.device_flow import ( @@ -23,10 +23,14 @@ ] +def _runtime_verify(auth_idp_runtime) -> str | bool: + return getattr(auth_idp_runtime, "verify", httpx_tls_config_from_env().get("verify", True)) + + def test_provider_gateway_serves_oidc_discovery(auth_idp_case, auth_idp_runtime): require_capability(auth_idp_case, "gateway_discovery") - verify = getattr(auth_idp_runtime, "verify", client_verify_from_env()) + verify = _runtime_verify(auth_idp_runtime) response = httpx.get(auth_idp_runtime.discovery_url, timeout=10.0, verify=verify) response.raise_for_status() @@ -51,7 +55,7 @@ def test_provider_device_authorization_endpoint_issues_user_code(auth_idp_case, require_capability(auth_idp_case, "device_flow") oidc = discover_nmp_config(auth_idp_runtime.gateway_base_url) - verify = getattr(auth_idp_runtime, "verify", client_verify_from_env()) + verify = _runtime_verify(auth_idp_runtime) device_authorization_endpoint = with_url_origin( oidc.device_authorization_endpoint, auth_idp_runtime.gateway_base_url, @@ -88,7 +92,7 @@ def test_provider_device_flow_returns_refresh_token(auth_idp_case, auth_idp_runt assert oidc.device_authorization_endpoint assert "offline_access" in oidc.default_scopes.split() - verify = getattr(auth_idp_runtime, "verify", client_verify_from_env()) + verify = _runtime_verify(auth_idp_runtime) device_authorization_endpoint = with_url_origin( oidc.device_authorization_endpoint, auth_idp_runtime.gateway_base_url, diff --git a/tests/auth_idp/contracts/test_gateway.py b/tests/auth_idp/contracts/test_gateway.py index cfbbfb10a7..17a4112a23 100644 --- a/tests/auth_idp/contracts/test_gateway.py +++ b/tests/auth_idp/contracts/test_gateway.py @@ -6,7 +6,7 @@ import httpx import pytest -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from nmp.testing import grant_workspace_role from tests.auth_idp.common import jwt_claims, require_capability @@ -25,7 +25,7 @@ def _runtime_verify(auth_idp_runtime) -> str | bool: - return getattr(auth_idp_runtime, "verify", client_verify_from_env()) + return getattr(auth_idp_runtime, "verify", httpx_tls_config_from_env().get("verify", True)) def _gateway_get_with_transient_retries( diff --git a/tests/auth_idp/runtime_compose.py b/tests/auth_idp/runtime_compose.py index a3f27fb524..04eef2b8b3 100644 --- a/tests/auth_idp/runtime_compose.py +++ b/tests/auth_idp/runtime_compose.py @@ -7,7 +7,7 @@ import httpx from nemo_platform import NeMoPlatform -from nemo_platform_ext.client.tls import client_verify_from_env +from nemo_platform_plugin.client.tls import httpx_tls_config_from_env from tests.auth_idp.common import jwt_claims from tests.auth_idp.runtime_contract import AuthIdpCase, TokenSet @@ -89,7 +89,7 @@ def exchange_workload_token(self, subject_token: str) -> TokenSet: "scope": workload_grant.get("scope", "openid email groups"), }, timeout=TOKEN_EXCHANGE_TIMEOUT_SECONDS, - verify=client_verify_from_env(), + **httpx_tls_config_from_env(), ) response.raise_for_status() token_response = response.json() diff --git a/tests/auth_idp/runtime_kubernetes.py b/tests/auth_idp/runtime_kubernetes.py index ccfdce3d77..edae04bb44 100644 --- a/tests/auth_idp/runtime_kubernetes.py +++ b/tests/auth_idp/runtime_kubernetes.py @@ -21,7 +21,7 @@ import httpx import pytest from nemo_platform import DefaultHttpxClient, NeMoPlatform -from nemo_platform_ext.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR +from nemo_platform_plugin.client.tls import NMP_CLIENT_SSL_CERT_FILE_ENVVAR from tests.auth_idp.common import jwt_claims from tests.auth_idp.runtime_contract import AuthIdpCase, TokenSet diff --git a/tests/test_e2e_services_pool.py b/tests/test_e2e_services_pool.py index 281337b967..794d227136 100644 --- a/tests/test_e2e_services_pool.py +++ b/tests/test_e2e_services_pool.py @@ -22,7 +22,7 @@ def test_render_e2e_config_for_docker_preserves_container_paths(tmp_path) -> Non "files": {"default_storage_config": {"type": "local", "path": "/data/files"}}, } - rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "docker"}) + rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "docker"}, "abc123") assert rendered["jobs"]["executors"][0]["config"]["working_directory"] == "/data/subprocess-jobs" assert rendered["files"]["default_storage_config"]["path"] == "/data/files" @@ -41,7 +41,7 @@ def test_render_e2e_config_for_subprocess_rewrites_instance_paths(tmp_path) -> N "files": {"default_storage_config": {"type": "local", "path": ".tmp/e2e/files"}}, } - rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "subprocess"}) + rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "subprocess"}, "abc123") assert rendered["jobs"]["executors"][0]["config"]["working_directory"] == str(tmp_path / "subprocess-jobs") assert rendered["files"]["default_storage_config"]["path"] == str(tmp_path / "files") @@ -88,7 +88,7 @@ def test_render_e2e_config_for_docker_compose_preserves_container_paths(tmp_path "files": {"default_storage_config": {"type": "local", "path": "/data/files"}}, } - rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "docker_compose"}) + rendered = services_pool._render_e2e_config_for_backend(config, tmp_path, {"backend": "docker_compose"}, "abc123") assert rendered["jobs"]["executors"][0]["config"]["working_directory"] == "/data/subprocess-jobs" assert rendered["files"]["default_storage_config"]["path"] == "/data/files" diff --git a/tools/nemo-platform-sdk-tools/src/nemo_platform_sdk_tools/sdk/vendor/vendor_package.py b/tools/nemo-platform-sdk-tools/src/nemo_platform_sdk_tools/sdk/vendor/vendor_package.py index b46a2c6d26..c010e41f65 100644 --- a/tools/nemo-platform-sdk-tools/src/nemo_platform_sdk_tools/sdk/vendor/vendor_package.py +++ b/tools/nemo-platform-sdk-tools/src/nemo_platform_sdk_tools/sdk/vendor/vendor_package.py @@ -78,17 +78,35 @@ class ResourceReplacement(BaseModel): _CLIENT_CLASS_NAMES = ("NeMoPlatform", "AsyncNeMoPlatform") -_CLIENT_METHOD_NAMES = ("__init__", "__getattr__", "copy") -_CLIENT_HELPER_FUNCTION_NAMES = ("_should_bootstrap_config", "_copy_requires_bootstrap") +_CLIENT_METHOD_NAMES = ( + "__init__", + "__getattr__", + "copy", + "custom_auth", + "http_client", + "token_provider", + "typed_client_options", +) +_CLIENT_HELPER_FUNCTION_NAMES = ( + "_should_bootstrap_config", + "_copy_requires_bootstrap", + "_typed_client_default_headers", + "_typed_client_retry", + "_typed_client_timeout", +) _CLIENT_INIT_REQUIRED_IMPORTS: dict[str, tuple[str, ...]] = { "nemo_platform._base_client": ("DefaultAsyncHttpxClient", "DefaultHttpxClient"), + "nemo_platform_plugin.client.auth": ("AsyncTokenProvider", "TokenProvider", "TokenProviderAuth"), "nemo_platform_plugin.client.constants": ("WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR",), - "nemo_platform_plugin.client.tls": ("client_verify_from_env",), + "nemo_platform_plugin.client.platform_options": ("AsyncPlatformClientOptions", "SyncPlatformClientOptions"), + "nemo_platform_plugin.client.tls": ("httpx_tls_config_from_env",), + "nemo_platform_plugin.client.types": ("RetryPolicy",), "pathlib": ("Path",), } _STALE_CLIENT_INIT_IMPORTS: dict[str, tuple[str, ...]] = { "nemo_platform.client.tls": ("client_verify_from_env",), "nemo_platform_ext.client.tls": ("client_verify_from_env",), + "nemo_platform_plugin.client.tls": ("client_verify_from_env",), } diff --git a/tools/nemo-platform-sdk-tools/tests/sdk/vendor/test_vendor_package.py b/tools/nemo-platform-sdk-tools/tests/sdk/vendor/test_vendor_package.py index 6664b4813a..90e1b34e48 100644 --- a/tools/nemo-platform-sdk-tools/tests/sdk/vendor/test_vendor_package.py +++ b/tools/nemo-platform-sdk-tools/tests/sdk/vendor/test_vendor_package.py @@ -1102,7 +1102,7 @@ class DefaultHttpxClient: encoding="utf-8", ) (plugin_client_path / "tls.py").write_text( - "def client_verify_from_env() -> bool:\n return True\n", + "def httpx_tls_config_from_env() -> dict[str, str]:\n return {}\n", encoding="utf-8", ) @@ -1169,8 +1169,9 @@ def __getattr__(self, name: str) -> Any: assert "from pathlib import Path" in updated assert "from nemo_platform._base_client import DefaultAsyncHttpxClient, DefaultHttpxClient" in updated assert "from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR" in updated - assert "from nemo_platform_plugin.client.tls import client_verify_from_env" in updated + assert "from nemo_platform_plugin.client.tls import httpx_tls_config_from_env" in updated assert "from nemo_platform_ext.client.tls import client_verify_from_env" not in updated + assert "from nemo_platform_plugin.client.tls import client_verify_from_env" not in updated assert "def _should_bootstrap_config(config_path: Path | None = None) -> bool:" in updated assert "return config_path is not None" in updated assert "return False" not in updated