diff --git a/packages/nemo_platform/pyproject.toml b/packages/nemo_platform/pyproject.toml index 7bf5caa741..e23e635342 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..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 @@ -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,142 @@ 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 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(sorted(eps)) + raise RuntimeError( + f"Multiple NemoClient providers registered under 'nemo.client_provider': {names}. " + "Only the platform (nmp-common) should register a provider." + ) + + if eps: + ep = next(iter(eps.values())) + try: + obj = ep.load() + if isinstance(obj, type): + obj = obj() + 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 + + +# --------------------------------------------------------------------------- +# 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 +292,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/dependencies.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py index f24531bdb6..a7a655fde4 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/dependencies.py @@ -11,6 +11,8 @@ from typing import TYPE_CHECKING, Any +from nemo_platform_plugin.client.client import AsyncNemoClient + if TYPE_CHECKING: from nemo_platform import AsyncNeMoPlatform from nemo_platform_plugin.config import PlatformConfig @@ -51,6 +53,18 @@ def get_sdk_client() -> "AsyncNeMoPlatform": ) +def get_nemo_client() -> AsyncNemoClient: + """FastAPI dependency for getting the async NemoClient. + + This is a placeholder. The actual client is injected via + app.dependency_overrides in Service.create_app(). + """ + raise RuntimeError( + "get_nemo_client() was called without being overridden. " + "Ensure your Service subclass calls super().create_app()." + ) + + def get_entity_client() -> "EntityClient": """FastAPI dependency for getting the EntityClient. 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..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,27 +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 - 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 new file mode 100644 index 0000000000..76ca4636fc --- /dev/null +++ b/packages/nemo_platform_plugin/tests/test_client_provider.py @@ -0,0 +1,344 @@ +# 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 +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, + *, + 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 + + +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_raises(self): + class _NotAProvider: + pass + + eps = [_FakeEntryPoint("platform", _NotAProvider())] + with patch("nemo_platform_plugin.client_provider.entry_points", return_value=eps): + 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, 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_name_same_target_is_deduplicated(self, monkeypatch): + monkeypatch.delenv("NMP_BASE_URL", raising=False) + 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) + 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_dependencies.py b/packages/nemo_platform_plugin/tests/test_dependencies.py new file mode 100644 index 0000000000..979c17f137 --- /dev/null +++ b/packages/nemo_platform_plugin/tests/test_dependencies.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for plugin-owned FastAPI dependency placeholders.""" + +import pytest +from nemo_platform_plugin.dependencies import get_nemo_client + + +def test_get_nemo_client_requires_platform_override() -> None: + with pytest.raises(RuntimeError, match=r"get_nemo_client\(\) was called without being overridden"): + get_nemo_client() diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index 5dc461454f..3fc96eb3e9 100644 --- a/packages/nemo_platform_plugin/tests/test_sdk_provider.py +++ b/packages/nemo_platform_plugin/tests/test_sdk_provider.py @@ -211,6 +211,26 @@ def get_platform_sdk(self, **kwargs) -> NeMoPlatform: return NeMoPlatform(base_url="http://custom:1234") +class _FakeEntryPoint: + 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 + + class TestProviderResolution: def setup_method(self): # Reset global state before each test. @@ -249,6 +269,49 @@ 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_raises(self): + class _NotAProvider: + pass + + eps = [_FakeEntryPoint("platform", _NotAProvider())] + with patch("nemo_platform_plugin.sdk_provider.entry_points", return_value=eps): + 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) + # --------------------------------------------------------------------------- # 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..3417d81c88 --- /dev/null +++ b/packages/nmp_common/src/nmp/common/client_factory.py @@ -0,0 +1,193 @@ +# 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) + 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``). + """ + 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/src/nmp/common/service/__init__.py b/packages/nmp_common/src/nmp/common/service/__init__.py index 084a7a65e8..1d8874e47a 100644 --- a/packages/nmp_common/src/nmp/common/service/__init__.py +++ b/packages/nmp_common/src/nmp/common/service/__init__.py @@ -6,6 +6,7 @@ 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, @@ -20,6 +21,7 @@ "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 f1d59d573c..4661821abd 100644 --- a/packages/nmp_common/src/nmp/common/service/base.py +++ b/packages/nmp_common/src/nmp/common/service/base.py @@ -10,12 +10,14 @@ 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 import httpx from fastapi import APIRouter, FastAPI from fastapi.openapi.utils import get_openapi from nemo_platform import AsyncNeMoPlatform, DefaultAsyncHttpxClient +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.controller import Controller @@ -57,7 +59,7 @@ def _get_config_class_from_generic(cls: type) -> Type[ServiceConfig] | None: class DependencyProvider: """ - Manages SDK, entity client, HTTP client, and config lifecycle for NeMo Platform services. + Manages SDK, NemoClient, entity client, HTTP client, and config lifecycle for NeMo Platform services. Provides lazy initialization, FastAPI dependency wiring, and cleanup. @@ -67,6 +69,7 @@ class DependencyProvider: """ def __init__(self) -> None: + self._client_lock = RLock() self._http_client: Optional[httpx.AsyncClient] = None self._sdk_client: Optional[AsyncNeMoPlatform] = None self._platform_config: Optional[PlatformConfig] = None @@ -79,9 +82,10 @@ def get_http_client(self) -> httpx.AsyncClient: If you need to share a client across providers (e.g., for connection pooling), you can inject the same client via _http_client. """ - if self._http_client is None: - self._http_client = DefaultAsyncHttpxClient() - return self._http_client + with self._client_lock: + if self._http_client is None: + self._http_client = DefaultAsyncHttpxClient() + return self._http_client def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: """Return the async platform SDK client. @@ -101,12 +105,13 @@ def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: # 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._http_client) + return get_async_platform_sdk(as_service=as_service, internal=True, http_client=self.get_http_client()) # For request handling, use cached SDK. EntityClient adds auth headers per-request. - if self._sdk_client is None: - self._sdk_client = get_async_platform_sdk(http_client=self._http_client) - return self._sdk_client + 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 def get_entity_client(self, as_service: str | None = None) -> Optional[EntityClient]: """Return the EntityClient. @@ -170,34 +175,39 @@ 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 setup_dependencies(self, app: FastAPI, service: "Service") -> None: """Configure FastAPI dependency overrides.""" from nmp.common.service.dependencies import ( 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_platform_config] = self.get_platform_config if service._service_config is not None: app.dependency_overrides[get_service_config] = lambda: service._service_config async def close(self) -> None: - """Close managed clients. - - Each DependencyProvider owns its HTTP client and SDK, so closing them - here is safe. Called by Service.on_shutdown() during lifespan cleanup. - """ - if self._http_client is not None: - await self._http_client.aclose() + """Close the provider-owned HTTP transport and clear cached wrappers.""" + with self._client_lock: + http_client = self._http_client self._http_client = None - if self._sdk_client is not None: - await self._sdk_client.close() self._sdk_client = None + if http_client is not None: + await http_client.aclose() + class Service(ABC, Generic[TConfig]): """ diff --git a/packages/nmp_common/src/nmp/common/service/dependencies.py b/packages/nmp_common/src/nmp/common/service/dependencies.py index 8eade0d015..e8daab07a3 100644 --- a/packages/nmp_common/src/nmp/common/service/dependencies.py +++ b/packages/nmp_common/src/nmp/common/service/dependencies.py @@ -13,6 +13,7 @@ from fastapi import Request 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 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..44410016ef --- /dev/null +++ b/packages/nmp_common/tests/client_factory/test_client_factory.py @@ -0,0 +1,252 @@ +# 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" 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 7d29d156e3..23dbee480b 100644 --- a/packages/nmp_common/tests/nmp_common/test_common_service.py +++ b/packages/nmp_common/tests/nmp_common/test_common_service.py @@ -3,11 +3,23 @@ """Tests for nmp.common.service module.""" +import time +from concurrent.futures import ThreadPoolExecutor +from threading import Barrier, Event, Lock from typing import List +from unittest.mock import AsyncMock, patch +import httpx import pytest -from fastapi import APIRouter, FastAPI +from fastapi import APIRouter, Depends, FastAPI +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.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 def _route_paths(app: FastAPI) -> set[str]: @@ -159,6 +171,18 @@ async def test_service_is_ready_default(self): assert await service.is_ready() is True +class CloseCountingAsyncClient(httpx.AsyncClient): + """Async transport that records lifecycle closure while retaining real HTTPX behavior.""" + + def __init__(self) -> None: + super().__init__(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) + self.close_count = 0 + + async def aclose(self) -> None: + self.close_count += 1 + await super().aclose() + + class TestDependencyProvider: """Tests for DependencyProvider class.""" @@ -168,6 +192,177 @@ def test_init(self): assert provider._sdk_client is None assert provider._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_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): + provider = DependencyProvider() + transport = AsyncMock(spec=httpx.AsyncClient) + provider._http_client = transport + + with patch( + "nmp.common.sdk_factory.get_principal_auth_headers", + return_value={ + "X-NMP-Principal-Id": "user-one@example.com", + "X-NMP-Principal-On-Behalf-Of": "delegate-one@example.com", + }, + ): + with scoped_otel_headers({"traceparent": "00-trace-one-span-one-01"}): + first = provider.get_request_scoped_nemo_client() + 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() + + 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" + + @pytest.mark.asyncio + async def test_close_closes_shared_sdk_and_nemo_transport_exactly_once(self): + provider = DependencyProvider() + transport = CloseCountingAsyncClient() + provider._http_client = transport + 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 provider._sdk_client is None + + @pytest.mark.asyncio + async def test_concurrent_first_dependency_resolution_creates_one_transport_and_sdk( + self, monkeypatch: pytest.MonkeyPatch + ): + from nmp.common import sdk_factory + from nmp.common.service import base as service_base + + 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: + resolution_ready.wait(timeout=5) + factory = provider.get_request_scoped_sdk if index % 2 == 0 else provider.get_request_scoped_nemo_client + return factory() + + def create_transport() -> CloseCountingAsyncClient: + transport = CloseCountingAsyncClient() + with created_lock: + created.append(transport) + factory_started.set() + assert release_factory.wait(timeout=5) + return transport + + monkeypatch.setattr(service_base, "DefaultAsyncHttpxClient", create_transport) + + with patch.object( + sdk_factory, "get_async_platform_sdk", wraps=sdk_factory.get_async_platform_sdk + ) as sdk_factory_call: + with ThreadPoolExecutor(max_workers=12) as executor: + futures = [executor.submit(resolve_dependency, index) 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) + + 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): + provider = DependencyProvider() + app = FastAPI() + provider.setup_dependencies(app, MockService()) + resolved: list[tuple[AsyncNemoClient, AsyncNemoClient]] = [] + + @app.get("/clients") + async def clients( + first: AsyncNemoClient = Depends(get_nemo_client), + second: AsyncNemoClient = Depends(get_nemo_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") + + assert first_response.json() == {"same": True} + assert second_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 + + await provider.close() + + @pytest.mark.asyncio + async def test_service_principal_sdk_shares_provider_transport(self): + provider = DependencyProvider() + cached_sdk = provider.get_sdk_client() + service_sdk = provider.get_sdk_client(as_service="entities") + + assert service_sdk is not cached_sdk + assert service_sdk._client is provider.get_http_client() + assert cached_sdk._client is provider.get_http_client() + + await provider.close() + @pytest.mark.asyncio async def test_close_without_clients(self): """Test close when no clients were created."""