From 9ffb820610ca43f27f042c975a122181b9fdcbb4 Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Mon, 20 Jul 2026 15:26:13 -0400 Subject: [PATCH 1/2] feat(client): add NemoClient provider foundation Signed-off-by: Max Dubrinsky --- packages/nemo_platform/pyproject.toml | 4 + .../nemo_platform_plugin/client_provider.py | 210 +++++++++++- .../src/nemo_platform_plugin/sdk_provider.py | 1 + .../tests/test_client_provider.py | 298 ++++++++++++++++++ .../tests/test_sdk_provider.py | 27 ++ packages/nmp_common/pyproject.toml | 3 + .../src/nmp/common/client_factory.py | 189 +++++++++++ .../client_factory/test_client_factory.py | 236 ++++++++++++++ 8 files changed, 952 insertions(+), 16 deletions(-) create mode 100644 packages/nemo_platform_plugin/tests/test_client_provider.py create mode 100644 packages/nmp_common/src/nmp/common/client_factory.py create mode 100644 packages/nmp_common/tests/client_factory/test_client_factory.py diff --git a/packages/nemo_platform/pyproject.toml b/packages/nemo_platform/pyproject.toml index 34d0d595e3..01b9f57ed8 100644 --- a/packages/nemo_platform/pyproject.toml +++ b/packages/nemo_platform/pyproject.toml @@ -490,6 +490,10 @@ auditor = "nemo_auditor.cli:AuditorPluginCLI" data-designer = "nemo_data_designer_plugin.cli.main:DataDesignerCLI" evaluator = "nemo_evaluator.cli:EvaluatorPluginCLI" +# 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_plugin/src/nemo_platform_plugin/client_provider.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py index a719dc9ed3..0ffb239c22 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,14 +1,33 @@ # 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. +"""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. -Builds :class:`~nemo_platform_plugin.client.client.NemoClient` / -:class:`~nemo_platform_plugin.client.client.AsyncNemoClient` from -environment variables (``NMP_BASE_URL``, ``NMP_PRINCIPAL``). +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. -For user-facing / CLI usage, prefer ``NemoClient.from_config()`` which -reads ``~/.config/nmp/config.yaml`` and wires up OIDC token refresh. +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. +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. """ from __future__ import annotations @@ -16,7 +35,8 @@ import json import logging import os -from typing import Any +from importlib.metadata import entry_points +from typing import Any, Protocol, runtime_checkable from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient @@ -26,7 +46,52 @@ _NMP_PRINCIPAL_ENVVAR = "NMP_PRINCIPAL" +# --------------------------------------------------------------------------- +# Protocol +# --------------------------------------------------------------------------- + + +@runtime_checkable +class NemoClientProvider(Protocol): + """Contract for building authenticated NemoClient handles. + + Implementations live outside this module — the default is below; + ``nmp-common`` ships a richer one registered via entry-point. + """ + + def get_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + workspace: str | None = None, + ) -> NemoClient: + """Build a sync NemoClient for the current service context.""" + + def get_async_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + workspace: str | None = None, + ) -> AsyncNemoClient: + """Build an async NemoClient for the current service context.""" + + +# --------------------------------------------------------------------------- +# Default provider (env-var based, zero nmp-common dependency) +# --------------------------------------------------------------------------- + + def _read_principal_from_env() -> dict[str, Any] | None: + """Read and parse ``NMP_PRINCIPAL`` from the environment. + + Returns ``None`` when the variable is absent or empty. Raises + :class:`ValueError` on malformed JSON so task containers surface the same + error as ``nmp.common``. + """ raw = os.environ.get(_NMP_PRINCIPAL_ENVVAR) if not raw: return None @@ -69,6 +134,12 @@ def _build_headers( headers["X-NMP-Principal-On-Behalf-Of-Groups"] = ",".join(principal["on_behalf_of_groups"]) if on_behalf_of is not None: + # An explicit override wins over any on-behalf-of delegation carried by + # the env principal. Drop the principal's stale sub-headers so we don't + # ship a mismatched delegated identity (correct id but wrong + # email/groups) -- mirrors nmp.common.sdk_factory._get_default_headers. + headers.pop("X-NMP-Principal-On-Behalf-Of-Email", None) + headers.pop("X-NMP-Principal-On-Behalf-Of-Groups", None) headers["X-NMP-Principal-On-Behalf-Of"] = on_behalf_of return headers @@ -78,19 +149,121 @@ def _base_url() -> str: return os.environ.get("NMP_BASE_URL", "http://localhost:8080") +class DefaultNemoClientProvider: + """Env-var-based provider that ships with the plugin package. + + Reads ``NMP_BASE_URL`` (default ``http://localhost:8080``) and + ``NMP_PRINCIPAL`` — both are set by the jobs backend before launching task + containers. No ``nmp-common`` imports, so it works standalone. + """ + + def get_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + workspace: str | None = None, + ) -> NemoClient: + 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) + + def get_async_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | None = None, + workspace: str | None = None, + ) -> AsyncNemoClient: + 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) + + +# --------------------------------------------------------------------------- +# Provider resolution +# --------------------------------------------------------------------------- + +_cached_provider: NemoClientProvider | None = None + + +def set_nemo_client_provider(provider: NemoClientProvider | None) -> None: + """Override the provider (primarily for tests). + + Pass ``None`` to clear the override and fall back to entry-point discovery + on the next call. + """ + global _cached_provider + _cached_provider = provider + + +def _resolve_provider() -> NemoClientProvider: + """Resolve the provider once: explicit override → entry-point → default.""" + global _cached_provider + if _cached_provider is not None: + return _cached_provider + + # Scan entry-points. nmp-common registers a provider; the nemo-platform + # bundle inherits the same entry-point, so deduplicate by name. + eps = {ep.name: ep for ep in entry_points(group="nemo.client_provider")} + if len(eps) > 1: + names = ", ".join(eps) + raise RuntimeError( + f"Multiple NemoClient providers registered under 'nemo.client_provider': {names}. " + "Only the platform (nmp-common) should register a provider." + ) + for ep in eps.values(): + try: + obj = ep.load() + if isinstance(obj, type): + obj = obj() + if isinstance(obj, NemoClientProvider): + logger.debug("Using NemoClient provider from entry-point %r", ep.name) + _cached_provider = obj + return obj + logger.warning("Entry-point %r loaded but does not satisfy NemoClientProvider; skipping", ep.name) + except Exception: + logger.warning("Failed to load NemoClient provider %r; skipping", ep.name, exc_info=True) + + # Fall back to the built-in default. + logger.debug("No entry-point NemoClient provider found; using DefaultNemoClientProvider") + _cached_provider = DefaultNemoClientProvider() + return _cached_provider + + +# --------------------------------------------------------------------------- +# Public API +# --------------------------------------------------------------------------- + + def get_nemo_client( *, as_service: str | None = None, internal: bool = False, on_behalf_of: str | None = None, + workspace: str | None = None, ) -> NemoClient: """Build a sync NemoClient for the current service context. - Reads ``NMP_BASE_URL`` (default ``http://localhost:8080``) and - ``NMP_PRINCIPAL`` from the environment. + 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 + OTEL headers. + + Args: + as_service: If provided, authenticate as ``service:{as_service}``. + If ``None``, propagate the principal read from ``NMP_PRINCIPAL``. + internal: Mark requests as internal (service-to-service). + on_behalf_of: Principal ID to act on behalf of. + workspace: Default workspace used to fill ``{workspace}`` path params. """ - headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return NemoClient(base_url=_base_url(), default_headers=headers or None) + return _resolve_provider().get_nemo_client( + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + workspace=workspace, + ) def get_async_nemo_client( @@ -98,11 +271,16 @@ def get_async_nemo_client( as_service: str | None = None, internal: bool = False, on_behalf_of: str | None = None, + workspace: str | None = None, ) -> AsyncNemoClient: - """Build an async NemoClient for the current service context. + """Async counterpart of :func:`get_nemo_client`. - Reads ``NMP_BASE_URL`` (default ``http://localhost:8080``) and - ``NMP_PRINCIPAL`` from the environment. + Used by middleware and controllers that run inside the platform service + process and need an async client. """ - headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return AsyncNemoClient(base_url=_base_url(), default_headers=headers or None) + return _resolve_provider().get_async_nemo_client( + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + workspace=workspace, + ) 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 b93bb909b4..9606a7cfc1 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 @@ -312,6 +312,7 @@ def _resolve_provider() -> SDKProvider: logger.debug("Using SDK provider from entry-point %r", ep.name) _cached_provider = obj return obj + logger.warning("Entry-point %r loaded but does not satisfy SDKProvider; skipping", ep.name) except Exception: logger.warning("Failed to load SDK provider %r; skipping", ep.name, exc_info=True) diff --git a/packages/nemo_platform_plugin/tests/test_client_provider.py b/packages/nemo_platform_plugin/tests/test_client_provider.py new file mode 100644 index 0000000000..a22b807a6d --- /dev/null +++ b/packages/nemo_platform_plugin/tests/test_client_provider.py @@ -0,0 +1,298 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""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``. +""" + +from __future__ import annotations + +import json +import logging +from unittest.mock import patch + +import pytest +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client_provider import ( + DefaultNemoClientProvider, + NemoClientProvider, + _build_headers, + _read_principal_from_env, + get_async_nemo_client, + get_nemo_client, + set_nemo_client_provider, +) + +# --------------------------------------------------------------------------- +# _read_principal_from_env +# --------------------------------------------------------------------------- + + +class TestReadPrincipalFromEnv: + def test_returns_none_when_unset(self, monkeypatch): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + assert _read_principal_from_env() is None + + def test_returns_none_when_empty(self, monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", "") + assert _read_principal_from_env() is None + + def test_returns_none_when_id_missing(self, monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"email": "a@b.com"})) + assert _read_principal_from_env() is None + + def test_returns_none_when_id_empty(self, monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": ""})) + assert _read_principal_from_env() is None + + def test_parses_valid_principal(self, monkeypatch): + principal = {"id": "user@example.com", "email": "user@example.com", "groups": ["team-a"]} + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(principal)) + assert _read_principal_from_env() == principal + + def test_raises_on_malformed_json(self, monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", "not-json") + with pytest.raises(ValueError, match="Invalid JSON"): + _read_principal_from_env() + + +# --------------------------------------------------------------------------- +# _build_headers +# --------------------------------------------------------------------------- + + +class TestBuildHeaders: + def test_internal_flag(self): + assert _build_headers(internal=True)["X-NMP-Internal"] == "true" + + def test_service_principal(self): + assert _build_headers(as_service="svc")["X-NMP-Principal-Id"] == "service:svc" + + def test_explicit_on_behalf_of(self): + headers = _build_headers(as_service="svc", on_behalf_of="user@ex.com") + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user@ex.com" + + def test_principal_from_env_when_no_service(self, monkeypatch): + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps( + { + "id": "user@ex.com", + "email": "user@ex.com", + "groups": ["g1", "g2"], + "on_behalf_of": "boss@ex.com", + "on_behalf_of_email": "boss@ex.com", + "on_behalf_of_groups": ["admin"], + } + ), + ) + headers = _build_headers() + assert headers["X-NMP-Principal-Id"] == "user@ex.com" + assert headers["X-NMP-Principal-Email"] == "user@ex.com" + assert headers["X-NMP-Principal-Groups"] == "g1,g2" + assert headers["X-NMP-Principal-On-Behalf-Of"] == "boss@ex.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Email"] == "boss@ex.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "admin" + + def test_service_principal_ignores_env_principal(self, monkeypatch): + # The as_service branch does not read NMP_PRINCIPAL (matches legacy behavior). + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "user@ex.com", "on_behalf_of": "boss@ex.com"})) + headers = _build_headers(as_service="svc") + assert headers["X-NMP-Principal-Id"] == "service:svc" + assert "X-NMP-Principal-On-Behalf-Of" not in headers + + def test_explicit_on_behalf_of_overrides_env_delegation(self, monkeypatch): + # An explicit on_behalf_of must not leave behind the env principal's + # delegated email/groups sub-headers: those describe a different + # identity. Only the overridden -On-Behalf-Of id should survive, matching + # nmp.common.sdk_factory._get_default_headers. + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps( + { + "id": "owner@ex.com", + "on_behalf_of": "boss@ex.com", + "on_behalf_of_email": "boss@ex.com", + "on_behalf_of_groups": ["admin"], + } + ), + ) + headers = _build_headers(on_behalf_of="override@ex.com") + assert headers["X-NMP-Principal-On-Behalf-Of"] == "override@ex.com" + assert "X-NMP-Principal-On-Behalf-Of-Email" not in headers + assert "X-NMP-Principal-On-Behalf-Of-Groups" not in headers + + +# --------------------------------------------------------------------------- +# DefaultNemoClientProvider +# --------------------------------------------------------------------------- + + +class TestDefaultNemoClientProvider: + def test_sync_default_base_url(self, monkeypatch): + monkeypatch.delenv("NMP_BASE_URL", raising=False) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + client = DefaultNemoClientProvider().get_nemo_client() + assert isinstance(client, NemoClient) + assert client.base_url == "http://localhost:8080" + + def test_sync_env_base_url_and_service_internal(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + client = DefaultNemoClientProvider().get_nemo_client(as_service="evaluator", internal=True) + assert client.base_url == "http://test:9090" + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert client._default_headers["X-NMP-Internal"] == "true" + + def test_sync_workspace_passthrough(self, monkeypatch): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + client = DefaultNemoClientProvider().get_nemo_client(workspace="team-a") + assert client.workspace == "team-a" + + def test_sync_propagates_env_principal_on_behalf_of(self, monkeypatch): + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps({"id": "creator@ex.com", "on_behalf_of": "real@ex.com"}), + ) + client = DefaultNemoClientProvider().get_nemo_client() + assert client._default_headers["X-NMP-Principal-Id"] == "creator@ex.com" + assert client._default_headers["X-NMP-Principal-On-Behalf-Of"] == "real@ex.com" + + def test_async_service_internal(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + client = DefaultNemoClientProvider().get_async_nemo_client(as_service="evaluator", internal=True) + assert isinstance(client, AsyncNemoClient) + assert client.base_url == "http://test:9090" + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert client._default_headers["X-NMP-Internal"] == "true" + + def test_async_workspace_passthrough(self, monkeypatch): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + client = DefaultNemoClientProvider().get_async_nemo_client(workspace="team-a") + assert client.workspace == "team-a" + + +# --------------------------------------------------------------------------- +# Provider resolution +# --------------------------------------------------------------------------- + + +class _FakeEntryPoint: + def __init__(self, name: str, obj: object) -> None: + self.name = name + self._obj = obj + + def load(self) -> object: + return self._obj + + +class _CustomProvider: + def get_nemo_client(self, **kwargs) -> NemoClient: + return NemoClient(base_url="http://custom:1234") + + def get_async_nemo_client(self, **kwargs) -> AsyncNemoClient: + return AsyncNemoClient(base_url="http://custom:1234") + + +class TestProviderResolution: + def setup_method(self): + set_nemo_client_provider(None) + + def teardown_method(self): + set_nemo_client_provider(None) + + def test_explicit_provider_takes_precedence(self, monkeypatch): + monkeypatch.delenv("NMP_BASE_URL", raising=False) + set_nemo_client_provider(_CustomProvider()) + assert get_nemo_client().base_url == "http://custom:1234" + assert get_async_nemo_client().base_url == "http://custom:1234" + + def test_falls_back_to_default_when_no_entry_points(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://fallback:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=[]): + client = get_nemo_client() + assert client.base_url == "http://fallback:8080" + + def test_entry_point_provider_is_discovered(self, monkeypatch): + monkeypatch.delenv("NMP_BASE_URL", raising=False) + eps = [_FakeEntryPoint("platform", _CustomProvider)] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + client = get_nemo_client() + assert client.base_url == "http://custom:1234" + + def test_entry_point_instance_is_discovered(self, monkeypatch): + # An entry-point that loads an instance (not a class) is used as-is. + monkeypatch.delenv("NMP_BASE_URL", raising=False) + eps = [_FakeEntryPoint("platform", _CustomProvider())] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + assert get_nemo_client().base_url == "http://custom:1234" + + def test_entry_point_not_satisfying_protocol_is_skipped(self, monkeypatch, caplog): + # An EP that loads an object which is not a NemoClientProvider is skipped + # with a WARNING, and resolution falls back to the default provider + # rather than degrading silently. + monkeypatch.setenv("NMP_BASE_URL", "http://fallback:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + class _NotAProvider: + pass + + eps = [_FakeEntryPoint("platform", _NotAProvider())] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + with caplog.at_level(logging.WARNING): + client = get_nemo_client() + assert client.base_url == "http://fallback:8080" + assert any("does not satisfy NemoClientProvider" in r.message for r in caplog.records) + + def test_multiple_named_providers_raise(self): + eps = [_FakeEntryPoint("platform", _CustomProvider), _FakeEntryPoint("other", _CustomProvider)] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + with pytest.raises(RuntimeError, match="Multiple NemoClient providers"): + get_nemo_client() + + def test_duplicate_named_providers_dedup(self, monkeypatch): + # The bundle re-registers the same named entry-point; dedup by name avoids a spurious error. + monkeypatch.delenv("NMP_BASE_URL", raising=False) + eps = [_FakeEntryPoint("platform", _CustomProvider), _FakeEntryPoint("platform", _CustomProvider)] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + assert get_nemo_client().base_url == "http://custom:1234" + + def test_set_none_clears_and_re_resolves(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://re-resolved:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + set_nemo_client_provider(_CustomProvider()) + assert get_nemo_client().base_url == "http://custom:1234" + set_nemo_client_provider(None) + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=[]): + assert get_nemo_client().base_url == "http://re-resolved:8080" + + def test_public_functions_pass_workspace_through(self): + captured: dict[str, object] = {} + + class _CapturingProvider: + def get_nemo_client(self, **kwargs): + captured.update(kwargs) + return NemoClient(base_url="http://x") + + def get_async_nemo_client(self, **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"} + + +# --------------------------------------------------------------------------- +# Protocol conformance +# --------------------------------------------------------------------------- + + +class TestProtocolConformance: + def test_default_provider_is_protocol_instance(self): + assert isinstance(DefaultNemoClientProvider(), NemoClientProvider) diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index 5dc461454f..69a39821d4 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 @@ -211,6 +212,15 @@ def get_platform_sdk(self, **kwargs) -> NeMoPlatform: return NeMoPlatform(base_url="http://custom:1234") +class _FakeEntryPoint: + def __init__(self, name: str, obj: object) -> None: + self.name = name + self._obj = obj + + def load(self) -> object: + return self._obj + + class TestProviderResolution: def setup_method(self): # Reset global state before each test. @@ -249,6 +259,23 @@ def test_set_none_clears_and_re_resolves(self, monkeypatch): sdk = get_task_sdk("x") assert sdk.base_url == "http://re-resolved:8080" + def test_entry_point_not_satisfying_protocol_is_skipped(self, monkeypatch, caplog): + # An EP that loads an object which is not an SDKProvider is skipped with a + # WARNING, and resolution falls back to the default provider rather than + # degrading silently. + monkeypatch.setenv("NMP_BASE_URL", "http://fallback:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + class _NotAProvider: + pass + + eps = [_FakeEntryPoint("platform", _NotAProvider())] + with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=eps): + with caplog.at_level(logging.WARNING): + sdk = get_task_sdk("test") + assert sdk.base_url == "http://fallback:8080" + assert any("does not satisfy SDKProvider" in r.message for r in caplog.records) + # --------------------------------------------------------------------------- # Protocol conformance diff --git a/packages/nmp_common/pyproject.toml b/packages/nmp_common/pyproject.toml index 93c745b022..95b2f05d44 100644 --- a/packages/nmp_common/pyproject.toml +++ b/packages/nmp_common/pyproject.toml @@ -69,5 +69,8 @@ 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/client_factory.py b/packages/nmp_common/src/nmp/common/client_factory.py new file mode 100644 index 0000000000..0ba6c8d486 --- /dev/null +++ b/packages/nmp_common/src/nmp/common/client_factory.py @@ -0,0 +1,189 @@ +# 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 +from collections.abc import Callable + +import httpx +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nmp.common.auth import Principal +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.otel import get_otel_headers +from nmp.common.sdk_factory import PlatformRequestRouter, _get_default_headers + +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 _base_url() -> str: + return get_platform_config().base_url + + +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) + headers.update(get_otel_headers()) + 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``). + """ + return NemoClient( + base_url=_base_url(), + workspace=workspace, + default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, + http_client=http_client or shared_sync_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``. + """ + effective_client = http_client or _test_http_client or shared_async_http_client() + return AsyncNemoClient( + base_url=_base_url(), + workspace=workspace, + default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, + http_client=effective_client, + url_resolver=_platform_url_resolver(), + ) + + +# --------------------------------------------------------------------------- +# 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, + ) diff --git a/packages/nmp_common/tests/client_factory/test_client_factory.py b/packages/nmp_common/tests/client_factory/test_client_factory.py new file mode 100644 index 0000000000..cfd2f630eb --- /dev/null +++ b/packages/nmp_common/tests/client_factory/test_client_factory.py @@ -0,0 +1,236 @@ +# 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(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" + + 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" From 248640398ce063bc537c7a26b6774527a2a1b23f Mon Sep 17 00:00:00 2001 From: Max Dubrinsky Date: Mon, 20 Jul 2026 18:41:49 -0400 Subject: [PATCH 2/2] fix(client): harden NemoClient provider discovery Signed-off-by: Max Dubrinsky --- .../nemo_platform_plugin/client_provider.py | 47 ++++++++--- .../src/nemo_platform_plugin/sdk_provider.py | 44 +++++++--- .../tests/test_client_provider.py | 80 +++++++++++++++---- .../tests/test_sdk_provider.py | 62 +++++++++++--- .../src/nmp/common/client_factory.py | 6 +- .../client_factory/test_client_factory.py | 18 ++++- 6 files changed, 199 insertions(+), 58 deletions(-) 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 0ffb239c22..5e9b3590a5 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 @@ -204,28 +204,49 @@ def _resolve_provider() -> NemoClientProvider: return _cached_provider # Scan entry-points. nmp-common registers a provider; the nemo-platform - # bundle inherits the same entry-point, so deduplicate by name. - eps = {ep.name: ep for ep in entry_points(group="nemo.client_provider")} + # bundle inherits the same entry-point, so identical registrations are + # legitimate duplicates. A duplicate name pointing elsewhere is a + # conflicting registration and must not depend on metadata ordering. + eps = {} + for ep in sorted( + entry_points(group="nemo.client_provider"), key=lambda candidate: (candidate.name, candidate.value) + ): + existing = eps.get(ep.name) + if existing is not None and existing.value != ep.value: + targets = ", ".join(sorted((existing.value, ep.value))) + raise RuntimeError( + f"Conflicting NemoClient providers registered under 'nemo.client_provider' with name {ep.name!r}: " + f"{targets}. Provider names must resolve to a single target." + ) + eps[ep.name] = ep + if len(eps) > 1: - names = ", ".join(eps) + names = ", ".join(sorted(eps)) raise RuntimeError( f"Multiple NemoClient providers registered under 'nemo.client_provider': {names}. " "Only the platform (nmp-common) should register a provider." ) - for ep in eps.values(): + + if eps: + ep = next(iter(eps.values())) try: obj = ep.load() if isinstance(obj, type): obj = obj() - if isinstance(obj, NemoClientProvider): - logger.debug("Using NemoClient provider from entry-point %r", ep.name) - _cached_provider = obj - return obj - logger.warning("Entry-point %r loaded but does not satisfy NemoClientProvider; skipping", ep.name) - except Exception: - logger.warning("Failed to load NemoClient provider %r; skipping", ep.name, exc_info=True) - - # Fall back to the built-in default. + except Exception as exc: + raise RuntimeError( + f"Failed to load or construct NemoClient provider {ep.name!r} from entry-point target {ep.value!r}." + ) from exc + if not isinstance(obj, NemoClientProvider): + raise RuntimeError( + f"NemoClient provider {ep.name!r} from entry-point target {ep.value!r} " + "does not satisfy NemoClientProvider." + ) + logger.debug("Using NemoClient provider from entry-point %r", ep.name) + _cached_provider = obj + return obj + + # Fall back to the built-in default only when no provider is registered. logger.debug("No entry-point NemoClient provider found; using DefaultNemoClientProvider") _cached_provider = DefaultNemoClientProvider() return _cached_provider 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 9606a7cfc1..7c2a8bcf44 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 @@ -295,28 +295,46 @@ def _resolve_provider() -> SDKProvider: return _cached_provider # Scan entry-points. nmp-common registers a provider; the nemo-platform - # bundle inherits the same entry-point, so deduplicate by name. - eps = {ep.name: ep for ep in entry_points(group="nemo.sdk_provider")} + # bundle inherits the same entry-point, so identical registrations are + # legitimate duplicates. A duplicate name pointing elsewhere is a + # conflicting registration and must not depend on metadata ordering. + eps = {} + for ep in sorted(entry_points(group="nemo.sdk_provider"), key=lambda candidate: (candidate.name, candidate.value)): + existing = eps.get(ep.name) + if existing is not None and existing.value != ep.value: + targets = ", ".join(sorted((existing.value, ep.value))) + raise RuntimeError( + f"Conflicting SDK providers registered under 'nemo.sdk_provider' with name {ep.name!r}: " + f"{targets}. Provider names must resolve to a single target." + ) + eps[ep.name] = ep + if len(eps) > 1: - names = ", ".join(eps) + names = ", ".join(sorted(eps)) raise RuntimeError( f"Multiple SDK providers registered under 'nemo.sdk_provider': {names}. " "Only the platform (nmp-common) should register a provider." ) - for ep in eps.values(): + + if eps: + ep = next(iter(eps.values())) try: obj = ep.load() if isinstance(obj, type): obj = obj() - if isinstance(obj, SDKProvider): - logger.debug("Using SDK provider from entry-point %r", ep.name) - _cached_provider = obj - return obj - logger.warning("Entry-point %r loaded but does not satisfy SDKProvider; skipping", ep.name) - except Exception: - logger.warning("Failed to load SDK provider %r; skipping", ep.name, exc_info=True) - - # Fall back to the built-in default. + except Exception as exc: + raise RuntimeError( + f"Failed to load or construct SDK provider {ep.name!r} from entry-point target {ep.value!r}." + ) from exc + if not isinstance(obj, SDKProvider): + raise RuntimeError( + f"SDK provider {ep.name!r} from entry-point target {ep.value!r} does not satisfy SDKProvider." + ) + logger.debug("Using SDK provider from entry-point %r", ep.name) + _cached_provider = obj + return obj + + # Fall back to the built-in default only when no provider is registered. logger.debug("No entry-point SDK provider found; using DefaultSDKProvider") _cached_provider = DefaultSDKProvider() return _cached_provider diff --git a/packages/nemo_platform_plugin/tests/test_client_provider.py b/packages/nemo_platform_plugin/tests/test_client_provider.py index a22b807a6d..76ca4636fc 100644 --- a/packages/nemo_platform_plugin/tests/test_client_provider.py +++ b/packages/nemo_platform_plugin/tests/test_client_provider.py @@ -11,7 +11,6 @@ from __future__ import annotations import json -import logging from unittest.mock import patch import pytest @@ -182,11 +181,22 @@ def test_async_workspace_passthrough(self, monkeypatch): class _FakeEntryPoint: - def __init__(self, name: str, obj: object) -> None: + def __init__( + self, + name: str, + obj: object, + *, + value: str = "tests:_CustomProvider", + load_error: Exception | None = None, + ) -> None: self.name = name + self.value = value self._obj = obj + self._load_error = load_error def load(self) -> object: + if self._load_error is not None: + raise self._load_error return self._obj @@ -232,36 +242,72 @@ def test_entry_point_instance_is_discovered(self, monkeypatch): with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): assert get_nemo_client().base_url == "http://custom:1234" - def test_entry_point_not_satisfying_protocol_is_skipped(self, monkeypatch, caplog): - # An EP that loads an object which is not a NemoClientProvider is skipped - # with a WARNING, and resolution falls back to the default provider - # rather than degrading silently. - monkeypatch.setenv("NMP_BASE_URL", "http://fallback:8080") - monkeypatch.delenv("NMP_PRINCIPAL", raising=False) - + def test_entry_point_not_satisfying_protocol_raises(self): class _NotAProvider: pass eps = [_FakeEntryPoint("platform", _NotAProvider())] with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): - with caplog.at_level(logging.WARNING): - client = get_nemo_client() - assert client.base_url == "http://fallback:8080" - assert any("does not satisfy NemoClientProvider" in r.message for r in caplog.records) + with pytest.raises(RuntimeError, match="does not satisfy NemoClientProvider"): + get_nemo_client() + + def test_entry_point_load_exception_raises(self): + eps = [_FakeEntryPoint("platform", _CustomProvider, load_error=ImportError("missing provider"))] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + with pytest.raises(RuntimeError, match="Failed to load or construct NemoClient provider") as exc_info: + get_nemo_client() + assert isinstance(exc_info.value.__cause__, ImportError) + + def test_entry_point_constructor_exception_raises(self): + class _BrokenProvider: + def __init__(self) -> None: + raise ValueError("invalid configuration") + + eps = [_FakeEntryPoint("platform", _BrokenProvider, value="tests:_BrokenProvider")] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + with pytest.raises(RuntimeError, match="Failed to load or construct NemoClient provider") as exc_info: + get_nemo_client() + assert isinstance(exc_info.value.__cause__, ValueError) + + def test_resolution_retries_after_entry_point_failure(self): + ep = _FakeEntryPoint("platform", _CustomProvider, load_error=ImportError("temporarily unavailable")) + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=[ep]): + with pytest.raises(RuntimeError, match="Failed to load or construct NemoClient provider"): + get_nemo_client() + ep._load_error = None + assert get_nemo_client().base_url == "http://custom:1234" def test_multiple_named_providers_raise(self): - eps = [_FakeEntryPoint("platform", _CustomProvider), _FakeEntryPoint("other", _CustomProvider)] + eps = [ + _FakeEntryPoint("platform", _CustomProvider, value="tests:_CustomProvider"), + _FakeEntryPoint("other", _CustomProvider, value="tests:_OtherProvider"), + ] with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): with pytest.raises(RuntimeError, match="Multiple NemoClient providers"): get_nemo_client() - def test_duplicate_named_providers_dedup(self, monkeypatch): - # The bundle re-registers the same named entry-point; dedup by name avoids a spurious error. + def test_duplicate_name_same_target_is_deduplicated(self, monkeypatch): monkeypatch.delenv("NMP_BASE_URL", raising=False) - eps = [_FakeEntryPoint("platform", _CustomProvider), _FakeEntryPoint("platform", _CustomProvider)] + eps = [ + _FakeEntryPoint("platform", _CustomProvider, value="tests:_CustomProvider"), + _FakeEntryPoint("platform", _CustomProvider, value="tests:_CustomProvider"), + ] with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): assert get_nemo_client().base_url == "http://custom:1234" + @pytest.mark.parametrize("reverse", [False, True]) + def test_duplicate_name_different_targets_raises_deterministically(self, reverse): + eps = [ + _FakeEntryPoint("platform", _CustomProvider, value="z_package:Provider"), + _FakeEntryPoint("platform", _CustomProvider, value="a_package:Provider"), + ] + if reverse: + eps.reverse() + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + with pytest.raises(RuntimeError, match="Conflicting NemoClient providers") as exc_info: + get_nemo_client() + assert "a_package:Provider, z_package:Provider" in str(exc_info.value) + def test_set_none_clears_and_re_resolves(self, monkeypatch): monkeypatch.setenv("NMP_BASE_URL", "http://re-resolved:8080") monkeypatch.delenv("NMP_PRINCIPAL", raising=False) diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index 69a39821d4..3fc96eb3e9 100644 --- a/packages/nemo_platform_plugin/tests/test_sdk_provider.py +++ b/packages/nemo_platform_plugin/tests/test_sdk_provider.py @@ -6,7 +6,6 @@ from __future__ import annotations import json -import logging from unittest.mock import patch import pytest @@ -213,11 +212,22 @@ def get_platform_sdk(self, **kwargs) -> NeMoPlatform: class _FakeEntryPoint: - def __init__(self, name: str, obj: object) -> None: + def __init__( + self, + name: str, + obj: object, + *, + value: str = "tests:DefaultSDKProvider", + load_error: Exception | None = None, + ) -> None: self.name = name + self.value = value self._obj = obj + self._load_error = load_error def load(self) -> object: + if self._load_error is not None: + raise self._load_error return self._obj @@ -259,22 +269,48 @@ def test_set_none_clears_and_re_resolves(self, monkeypatch): sdk = get_task_sdk("x") assert sdk.base_url == "http://re-resolved:8080" - def test_entry_point_not_satisfying_protocol_is_skipped(self, monkeypatch, caplog): - # An EP that loads an object which is not an SDKProvider is skipped with a - # WARNING, and resolution falls back to the default provider rather than - # degrading silently. - monkeypatch.setenv("NMP_BASE_URL", "http://fallback:8080") - monkeypatch.delenv("NMP_PRINCIPAL", raising=False) - + def test_entry_point_not_satisfying_protocol_raises(self): class _NotAProvider: pass eps = [_FakeEntryPoint("platform", _NotAProvider())] with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=eps): - with caplog.at_level(logging.WARNING): - sdk = get_task_sdk("test") - assert sdk.base_url == "http://fallback:8080" - assert any("does not satisfy SDKProvider" in r.message for r in caplog.records) + with pytest.raises(RuntimeError, match="does not satisfy SDKProvider"): + get_task_sdk("test") + + def test_entry_point_load_exception_raises_and_retries(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://retried:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + ep = _FakeEntryPoint("platform", DefaultSDKProvider, load_error=ImportError("temporarily unavailable")) + with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=[ep]): + with pytest.raises(RuntimeError, match="Failed to load or construct SDK provider") as exc_info: + get_task_sdk("test") + assert isinstance(exc_info.value.__cause__, ImportError) + ep._load_error = None + assert get_task_sdk("test").base_url == "http://retried:8080" + + def test_duplicate_name_same_target_is_deduplicated(self, monkeypatch): + monkeypatch.setenv("NMP_BASE_URL", "http://deduplicated:8080") + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + eps = [ + _FakeEntryPoint("platform", DefaultSDKProvider, value="tests:DefaultSDKProvider"), + _FakeEntryPoint("platform", DefaultSDKProvider, value="tests:DefaultSDKProvider"), + ] + with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=eps): + assert get_task_sdk("test").base_url == "http://deduplicated:8080" + + @pytest.mark.parametrize("reverse", [False, True]) + def test_duplicate_name_different_targets_raises_deterministically(self, reverse): + eps = [ + _FakeEntryPoint("platform", DefaultSDKProvider, value="z_package:Provider"), + _FakeEntryPoint("platform", DefaultSDKProvider, value="a_package:Provider"), + ] + if reverse: + eps.reverse() + with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=eps): + with pytest.raises(RuntimeError, match="Conflicting SDK providers") as exc_info: + get_task_sdk("test") + assert "a_package:Provider, z_package:Provider" in str(exc_info.value) # --------------------------------------------------------------------------- diff --git a/packages/nmp_common/src/nmp/common/client_factory.py b/packages/nmp_common/src/nmp/common/client_factory.py index 0ba6c8d486..3417d81c88 100644 --- a/packages/nmp_common/src/nmp/common/client_factory.py +++ b/packages/nmp_common/src/nmp/common/client_factory.py @@ -75,7 +75,11 @@ def _platform_headers( request scope). """ headers = _get_default_headers(as_service, internal, on_behalf_of) - headers.update(get_otel_headers()) + 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 diff --git a/packages/nmp_common/tests/client_factory/test_client_factory.py b/packages/nmp_common/tests/client_factory/test_client_factory.py index cfd2f630eb..44410016ef 100644 --- a/packages/nmp_common/tests/client_factory/test_client_factory.py +++ b/packages/nmp_common/tests/client_factory/test_client_factory.py @@ -170,11 +170,27 @@ def test_propagates_request_principal_when_no_service(self): 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(self): + 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.