diff --git a/packages/nemo_platform/pyproject.toml b/packages/nemo_platform/pyproject.toml index e2acf75f19..f3d1f78308 100644 --- a/packages/nemo_platform/pyproject.toml +++ b/packages/nemo_platform/pyproject.toml @@ -505,6 +505,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..e59ba4dc46 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,9 +35,16 @@ import json import logging import os -from typing import Any +from importlib.metadata import entry_points +from pathlib import Path +from typing import Any, Protocol, runtime_checkable +from nemo_platform_plugin.client.auth import TokenProvider from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.constants import ( + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + is_workload_identity_token_file_set, +) logger = logging.getLogger(__name__) @@ -26,7 +52,74 @@ _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.""" + + def get_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + ) -> NemoClient: + """Build a sync NemoClient for use inside a task container. + + Mirrors ``nmp.common.sdk_factory.get_task_sdk``: authenticate as + ``service:{service_name}`` while acting on behalf of the job creator + (read from ``NMP_PRINCIPAL``), or bootstrap workload-identity bearer-token + exchange when ``NMP_WORKLOAD_IDENTITY_TOKEN_FILE`` is set. + """ + + def get_async_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + ) -> AsyncNemoClient: + """Async counterpart of :meth:`get_task_nemo_client`.""" + + +# --------------------------------------------------------------------------- +# 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,28 +162,251 @@ 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 +def _effective_on_behalf_of(principal: dict[str, Any]) -> tuple[str, list[str], str | None]: + """Collapse an env principal to its acting identity (id, groups, email). + + Mirrors :pyattr:`nmp.common.auth.Principal.effective_principal`: if the job + creator's principal is itself delegated, the ``on_behalf_of`` identity wins; + otherwise the principal's own identity is used. + """ + if principal.get("on_behalf_of"): + return ( + principal["on_behalf_of"], + list(principal.get("on_behalf_of_groups") or []), + principal.get("on_behalf_of_email"), + ) + return principal["id"], list(principal.get("groups") or []), principal.get("email") + + +def _build_task_headers(service_name: str) -> dict[str, str]: + """Headers for a task container: service principal + creator delegation. + + Wire-equivalent to ``get_task_sdk(as_service=service_name)`` in the + non-workload-identity path -- ``service:{service_name}`` plus the full + ``X-NMP-Principal-On-Behalf-Of*`` set derived from ``NMP_PRINCIPAL``. + """ + headers: dict[str, str] = { + _INTERNAL_REQUEST_HEADER: "true", + "X-NMP-Principal-Id": f"service:{service_name}", + } + principal = _read_principal_from_env() + if principal is None: + logger.warning( + "NMP_PRINCIPAL not set; task NemoClient will authenticate as service:%s without on-behalf-of delegation", + service_name, + ) + return headers + obo_id, obo_groups, obo_email = _effective_on_behalf_of(principal) + headers["X-NMP-Principal-On-Behalf-Of"] = obo_id + if obo_groups: + headers["X-NMP-Principal-On-Behalf-Of-Groups"] = ",".join(obo_groups) + if obo_email: + headers["X-NMP-Principal-On-Behalf-Of-Email"] = obo_email + return headers + + +def _workload_identity_auth(base_url: str) -> TokenProvider: + """Build a workload-identity token-exchange auth provider. + + Only call when :func:`is_workload_identity_token_file_set` is true. + """ + from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider + + token_file = os.environ[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] + return resolve_workload_exchange_provider(base_url=base_url, subject_token_file=Path(token_file)) + + def _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) + + def get_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + ) -> NemoClient: + base_url = _base_url() + if is_workload_identity_token_file_set(): + return NemoClient( + base_url=base_url, + workspace=workspace, + auth=_workload_identity_auth(base_url), + default_headers={_INTERNAL_REQUEST_HEADER: "true"}, + ) + return NemoClient( + base_url=base_url, + workspace=workspace, + default_headers=_build_task_headers(service_name), + ) + + def get_async_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + ) -> AsyncNemoClient: + base_url = _base_url() + if is_workload_identity_token_file_set(): + return AsyncNemoClient( + base_url=base_url, + workspace=workspace, + auth=_workload_identity_auth(base_url), + default_headers={_INTERNAL_REQUEST_HEADER: "true"}, + ) + return AsyncNemoClient( + base_url=base_url, + workspace=workspace, + default_headers=_build_task_headers(service_name), + ) + + +# --------------------------------------------------------------------------- +# 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 +414,37 @@ 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. + """ + return _resolve_provider().get_async_nemo_client( + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + workspace=workspace, + ) + + +def get_task_nemo_client(service_name: str, *, workspace: str | None = None) -> NemoClient: + """Build a sync NemoClient for use inside a task container. + + NemoClient counterpart of ``nmp.common.sdk_factory.get_task_sdk``. Reads the + job creator's principal from ``NMP_PRINCIPAL`` and authenticates as + ``service:{service_name}`` while acting on behalf of that creator, or -- + when ``NMP_WORKLOAD_IDENTITY_TOKEN_FILE`` is set -- bootstraps + workload-identity bearer-token exchange instead of trusted ``X-NMP-*`` + principal headers. + + Use this from task containers rather than ``get_nemo_client(as_service=...)``, + which authenticates as an *undelegated* service principal. """ - 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_task_nemo_client(service_name, workspace=workspace) + + +def get_async_task_nemo_client(service_name: str, *, workspace: str | None = None) -> AsyncNemoClient: + """Async counterpart of :func:`get_task_nemo_client`.""" + return _resolve_provider().get_async_task_nemo_client(service_name, 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 4057c4f748..49379fa2a2 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 @@ -316,27 +316,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..0c043f136c --- /dev/null +++ b/packages/nemo_platform_plugin/tests/test_client_provider.py @@ -0,0 +1,350 @@ +# 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") + + def get_task_nemo_client(self, service_name, **kwargs) -> NemoClient: + return NemoClient(base_url="http://custom:1234") + + def get_async_task_nemo_client(self, service_name, **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_nemo_client_task_auth.py b/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py new file mode 100644 index 0000000000..f31021c72f --- /dev/null +++ b/packages/nemo_platform_plugin/tests/test_nemo_client_task_auth.py @@ -0,0 +1,161 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Regression tests for the plugin-side (zero-nmp-common) NemoClient task path. + +Guards the three PR-800 review findings against the DefaultNemoClientProvider +and the public ``get_task_nemo_client`` API: + +1. task clients must carry job-creator on-behalf-of delegation; +2. workload-identity must bootstrap bearer-token auth (no trusted X-NMP-* headers); +3. ``get_nemo_client(as_service=...)`` must stay *undelegated* (background + controllers rely on that), so the task path is a distinct entry point. +""" + +from __future__ import annotations + +import json + +import pytest +from nemo_platform_plugin import client_provider as cp +from nemo_platform_plugin.client_provider import ( + DefaultNemoClientProvider, + get_async_task_nemo_client, + get_nemo_client, + get_task_nemo_client, + set_nemo_client_provider, +) + +CREATOR = { + "id": "user:alice@acme.com", + "email": "alice@acme.com", + "groups": ["team-a", "team-b"], +} + + +@pytest.fixture(autouse=True) +def _force_default_provider(monkeypatch): + # Pin the built-in default provider so the public helpers don't pick up an + # entry-point-registered platform provider from the environment. + set_nemo_client_provider(DefaultNemoClientProvider()) + monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") + monkeypatch.delenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", raising=False) + yield + set_nemo_client_provider(None) + + +# --------------------------------------------------------------------------- +# Claim 1 — task clients preserve creator delegation +# --------------------------------------------------------------------------- +def test_task_client_delegates_to_job_creator(monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) + headers = get_task_nemo_client("evaluator")._default_headers + + assert headers["X-NMP-Internal"] == "true" + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Email"] == "alice@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "team-a,team-b" + + +def test_task_client_collapses_already_delegated_creator(monkeypatch): + # If the creator principal was itself delegated, the acting identity wins. + delegated = { + "id": "service:jobs", + "on_behalf_of": "user:bob@acme.com", + "on_behalf_of_email": "bob@acme.com", + "on_behalf_of_groups": ["team-z"], + } + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(delegated)) + headers = get_task_nemo_client("evaluator")._default_headers + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:bob@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Email"] == "bob@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "team-z" + + +def test_task_client_without_principal_warns_and_stays_service(monkeypatch, caplog): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + with caplog.at_level("WARNING"): + headers = get_task_nemo_client("evaluator")._default_headers + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert "X-NMP-Principal-On-Behalf-Of" not in headers + assert "without on-behalf-of delegation" in caplog.text + + +async def test_async_task_client_delegates(monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) + headers = get_async_task_nemo_client("evaluator")._default_headers + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" + + +def test_general_as_service_stays_undelegated(monkeypatch): + # Contrast: the generic entry point must NOT silently delegate, so + # background controllers can act as an unscoped service principal. + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) + headers = get_nemo_client(as_service="evaluator")._default_headers + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert "X-NMP-Principal-On-Behalf-Of" not in headers + + +# --------------------------------------------------------------------------- +# Claim 2 — workload identity bootstraps bearer auth, drops trusted headers +# --------------------------------------------------------------------------- +class _FakeExchangeProvider: + def get_access_token(self) -> str: + return "exchanged-token" + + async def get_access_token_async(self) -> str: + return "exchanged-token" + + +@pytest.fixture +def _stub_workload_exchange(monkeypatch): + captured = {} + + def _fake(*, base_url, subject_token_file): + captured["base_url"] = base_url + captured["subject_token_file"] = str(subject_token_file) + return _FakeExchangeProvider() + + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + _fake, + ) + return captured + + +def test_task_client_uses_workload_identity(monkeypatch, tmp_path, _stub_workload_exchange): + token_file = tmp_path / "token" + token_file.write_text("subject-token") + monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) # must be ignored in WI mode + + client = get_task_nemo_client("evaluator") + + # Bearer exchange wired up... + assert isinstance(client._auth, _FakeExchangeProvider) + assert _stub_workload_exchange["subject_token_file"] == str(token_file) + assert _stub_workload_exchange["base_url"] == "http://platform:8080" + # ...and NO trusted principal headers are sent (they'd be stripped/rejected + # at a workload-identity trust boundary). + assert "X-NMP-Principal-Id" not in client._default_headers + assert client._default_headers.get("X-NMP-Internal") == "true" + + +async def test_async_task_client_uses_workload_identity(monkeypatch, tmp_path, _stub_workload_exchange): + token_file = tmp_path / "token" + token_file.write_text("subject-token") + monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + + client = get_async_task_nemo_client("evaluator") + assert isinstance(client._auth, _FakeExchangeProvider) + assert "X-NMP-Principal-Id" not in client._default_headers + + +def test_default_provider_directly(monkeypatch): + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps(CREATOR)) + provider = DefaultNemoClientProvider() + assert isinstance(provider, cp.NemoClientProvider) + headers = provider.get_task_nemo_client("evaluator")._default_headers + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index cc1a1ec77e..c4ff61c53b 100644 --- a/packages/nemo_platform_plugin/tests/test_sdk_provider.py +++ b/packages/nemo_platform_plugin/tests/test_sdk_provider.py @@ -250,6 +250,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. @@ -288,6 +308,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..26d883a306 --- /dev/null +++ b/packages/nmp_common/src/nmp/common/client_factory.py @@ -0,0 +1,351 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Rich NemoClient factory backed by platform internals. + +This is the :class:`~nemo_platform_plugin.client.client.NemoClient` sibling of +:mod:`nmp.common.sdk_factory`. It builds typed clients that reuse the same +platform machinery the SDK factory uses: + +- base URL from :class:`~nmp.common.config.Configuration`; +- per-service URL routing via :class:`~nmp.common.sdk_factory.PlatformRequestRouter`; +- the shared sync/async HTTP clients (connection-pool + SSL-context reuse); +- principal / auth + internal-request headers via ``_get_default_headers``; +- OTEL trace-propagation headers captured on the current request. + +:class:`PlatformNemoClientProvider` is registered under the ``nemo.client_provider`` +entry-point group so :func:`nemo_platform_plugin.client_provider.get_nemo_client` +discovers it automatically whenever ``nmp-common`` is installed. +""" + +from __future__ import annotations + +import logging +import os +from collections.abc import Callable +from pathlib import Path + +import httpx +from nemo_platform_plugin.client.auth import TokenProvider +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.constants import ( + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + is_workload_identity_token_file_set, +) +from nmp.common.auth import Principal, principal_from_env +from nmp.common.config import get_platform_config +from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client +from nmp.common.observability import MARK_INTERNAL_REQUEST_HEADERS +from nmp.common.observability.otel import get_otel_headers +from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint +from nmp.common.sdk_factory import PlatformRequestRouter, _get_default_headers, _should_bootstrap_workload_identity + +logger = logging.getLogger(__name__) + +# Test-only: async HTTP client to use for NemoClient requests in test context. +# Set by test fixtures to route requests through the in-process test transport, +# mirroring ``nmp.common.sdk_factory._test_http_client``. +_test_http_client: httpx.AsyncClient | None = None + + +def _sync_http_client_for_endpoint( + endpoint: PlatformEndpoint, + http_client: httpx.Client | None, +) -> httpx.Client: + """Endpoint-aware sync client: honour an explicit client, else a UDS + transport for ``unix://`` endpoints, else the shared TCP client.""" + if http_client is not None: + return http_client + if endpoint.transport == "uds": + return endpoint.sync_http_client() + return shared_sync_http_client() + + +def _async_http_client_for_endpoint( + endpoint: PlatformEndpoint, + http_client: httpx.AsyncClient | None, +) -> httpx.AsyncClient: + """Async counterpart of :func:`_sync_http_client_for_endpoint`. + + Preserves the module-level ``_test_http_client`` fixture hook ahead of the + UDS / shared-client selection (mirrors ``sdk_factory``). + """ + if http_client is not None: + return http_client + if _test_http_client is not None: + return _test_http_client + if endpoint.transport == "uds": + return endpoint.async_http_client() + return shared_async_http_client() + + +def _workload_identity_auth(base_url: str) -> TokenProvider: + """Build a workload-identity token-exchange auth provider. + + Only call when :func:`is_workload_identity_token_file_set` is true. + """ + from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider + + token_file = os.environ[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] + return resolve_workload_exchange_provider(base_url=base_url, subject_token_file=Path(token_file)) + + +def _workload_identity_headers(internal: bool) -> dict[str, str]: + return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} + + +def _absolute_url(url: str) -> httpx.URL: + """Default resolver for the request router. + + :class:`NemoClient` hands its ``url_resolver`` the fully-qualified request + URL (``base_url`` + path), so — unlike the generated SDK's ``_prepare_url``, + which resolves a relative path — the router just needs to parse it. + """ + return httpx.URL(url) + + +def _platform_url_resolver() -> Callable[[str], httpx.URL]: + """Build a per-service URL router bound to the current platform config.""" + router = PlatformRequestRouter( + platform_config=get_platform_config(), + default_resolver=_absolute_url, + ) + return router.resolve + + +def _platform_headers( + as_service: str | None, + internal: bool, + on_behalf_of: str | Principal | None, +) -> dict[str, str]: + """Auth / internal headers plus OTEL trace-propagation headers. + + ``_get_default_headers`` supplies the principal + internal-request markers + (wire-identical to the SDK factory); ``get_otel_headers`` layers on the + trace-propagation context captured on the current request (empty outside a + request scope). + """ + headers = _get_default_headers(as_service, internal, on_behalf_of) + for name, value in get_otel_headers().items(): + normalized_name = name.lower() + if normalized_name == "x-nmp-internal" or normalized_name.startswith("x-nmp-principal-"): + continue + headers[name] = value + return headers + + +def get_nemo_client( + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + workspace: str | None = None, + http_client: httpx.Client | None = None, +) -> NemoClient: + """Build a sync :class:`NemoClient` configured with platform internals. + + Args: + as_service: If provided, authenticate as ``service:{as_service}``. + If ``None``, propagate the current request's / env principal. + internal: Mark requests as internal (service-to-service). + on_behalf_of: Principal (or id) to act on behalf of. Passing a + :class:`~nmp.common.auth.Principal` (rather than a bare id string) + is only reachable through this direct entry point; the plugin-facing + :class:`~nemo_platform_plugin.client_provider.NemoClientProvider` + protocol narrows ``on_behalf_of`` to ``str | None``. + workspace: Default workspace used to fill ``{workspace}`` path params. + http_client: Optional sync HTTP client; defaults to the shared client. + + Note: + OTEL trace-propagation headers are captured once, at construction, from + the current request context. Build a fresh client per request scope + rather than caching one across requests, or its ``traceparent`` will be + stale (mirrors ``get_platform_sdk``). + """ + endpoint = resolve_platform_endpoint() + if _should_bootstrap_workload_identity( + as_service=as_service, + on_behalf_of=on_behalf_of, + http_client=http_client, + endpoint=endpoint, + ): + return NemoClient( + base_url=endpoint.connect_base_url, + workspace=workspace, + auth=_workload_identity_auth(endpoint.connect_base_url), + default_headers=_workload_identity_headers(internal) or None, + http_client=_sync_http_client_for_endpoint(endpoint, http_client), + url_resolver=_platform_url_resolver(), + ) + return NemoClient( + base_url=endpoint.connect_base_url, + workspace=workspace, + default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, + http_client=_sync_http_client_for_endpoint(endpoint, http_client), + url_resolver=_platform_url_resolver(), + ) + + +def get_async_nemo_client( + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + workspace: str | None = None, + http_client: httpx.AsyncClient | None = None, +) -> AsyncNemoClient: + """Async counterpart of :func:`get_nemo_client`. + + Uses the explicitly provided ``http_client`` (e.g. from a test fixture), then + the module-level ``_test_http_client`` fallback, then the shared async + client — mirroring ``nmp.common.sdk_factory.get_async_platform_sdk``. + """ + endpoint = resolve_platform_endpoint() + if _should_bootstrap_workload_identity( + as_service=as_service, + on_behalf_of=on_behalf_of, + http_client=http_client, + endpoint=endpoint, + ): + return AsyncNemoClient( + base_url=endpoint.connect_base_url, + workspace=workspace, + auth=_workload_identity_auth(endpoint.connect_base_url), + default_headers=_workload_identity_headers(internal) or None, + http_client=_async_http_client_for_endpoint(endpoint, http_client), + url_resolver=_platform_url_resolver(), + ) + return AsyncNemoClient( + base_url=endpoint.connect_base_url, + workspace=workspace, + default_headers=_platform_headers(as_service, internal, on_behalf_of) or None, + http_client=_async_http_client_for_endpoint(endpoint, http_client), + url_resolver=_platform_url_resolver(), + ) + + +def get_task_nemo_client( + service_name: str, + *, + workspace: str | None = None, + http_client: httpx.Client | None = None, +) -> NemoClient: + """Build a sync :class:`NemoClient` for use inside a task container. + + NemoClient counterpart of :func:`nmp.common.sdk_factory.get_task_sdk`: + reads the job creator's principal from ``NMP_PRINCIPAL`` and authenticates + as ``service:{service_name}`` while acting on behalf of that creator, or -- + when ``NMP_WORKLOAD_IDENTITY_TOKEN_FILE`` is set -- bootstraps + workload-identity bearer-token exchange (via :func:`get_nemo_client` with + ``internal=True``) instead of trusted ``X-NMP-*`` principal headers. + """ + if http_client is None and is_workload_identity_token_file_set(): + return get_nemo_client(internal=True, workspace=workspace) + if http_client is None: + http_client = resolve_platform_endpoint().sync_sdk_http_client() + principal = principal_from_env() + if principal is None: + logger.warning( + "NMP_PRINCIPAL not set; task NemoClient will authenticate as service:%s without on-behalf-of delegation", + service_name, + ) + return get_nemo_client( + as_service=service_name, + internal=True, + on_behalf_of=principal.effective_principal if principal else None, + workspace=workspace, + http_client=http_client, + ) + + +def get_async_task_nemo_client( + service_name: str, + *, + workspace: str | None = None, + http_client: httpx.AsyncClient | None = None, +) -> AsyncNemoClient: + """Async counterpart of :func:`get_task_nemo_client`. Wire-identical.""" + if http_client is None and is_workload_identity_token_file_set(): + return get_async_nemo_client(internal=True, workspace=workspace) + if http_client is None: + http_client = resolve_platform_endpoint().async_sdk_http_client() + principal = principal_from_env() + if principal is None: + logger.warning( + "NMP_PRINCIPAL not set; async task NemoClient will authenticate as service:%s without on-behalf-of delegation", + service_name, + ) + return get_async_nemo_client( + as_service=service_name, + internal=True, + on_behalf_of=principal.effective_principal if principal else None, + workspace=workspace, + http_client=http_client, + ) + + +# --------------------------------------------------------------------------- +# Entry-point provider for nemo_platform_plugin.client_provider +# --------------------------------------------------------------------------- + + +class PlatformNemoClientProvider: + """Rich :class:`~nemo_platform_plugin.client_provider.NemoClientProvider` + that uses platform internals (shared HTTP clients, URL routing, OTEL + headers, auth context). + + Registered as a ``nemo.client_provider`` entry-point so it is discovered + automatically when ``nmp-common`` is installed. + """ + + def get_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + workspace: str | None = None, + http_client: httpx.Client | None = None, + ) -> NemoClient: + return get_nemo_client( + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + workspace=workspace, + http_client=http_client, + ) + + def get_async_nemo_client( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + workspace: str | None = None, + http_client: httpx.AsyncClient | None = None, + ) -> AsyncNemoClient: + return get_async_nemo_client( + as_service=as_service, + internal=internal, + on_behalf_of=on_behalf_of, + workspace=workspace, + http_client=http_client, + ) + + def get_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + http_client: httpx.Client | None = None, + ) -> NemoClient: + return get_task_nemo_client(service_name, workspace=workspace, http_client=http_client) + + def get_async_task_nemo_client( + self, + service_name: str, + *, + workspace: str | None = None, + http_client: httpx.AsyncClient | None = None, + ) -> AsyncNemoClient: + return get_async_task_nemo_client(service_name, workspace=workspace, http_client=http_client) diff --git a/packages/nmp_common/src/nmp/common/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 34a51502c3..815ba54e9c 100644 --- a/packages/nmp_common/src/nmp/common/service/base.py +++ b/packages/nmp_common/src/nmp/common/service/base.py @@ -10,17 +10,19 @@ 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 import AsyncNeMoPlatform +from nemo_platform_plugin.client.client import AsyncNemoClient from nmp.common.api.utils import register_query_param_schemas from nmp.common.config import Configuration, PlatformConfig, ServiceConfig from nmp.common.controller import Controller from nmp.common.entities.client import EntityClient -from nmp.common.platform_endpoint import resolve_service_endpoint +from nmp.common.platform_endpoint import resolve_platform_endpoint, resolve_service_endpoint logger = logging.getLogger(__name__) @@ -58,7 +60,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. @@ -68,6 +70,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 @@ -76,13 +79,21 @@ def __init__(self) -> None: def get_http_client(self) -> httpx.AsyncClient: """Return the httpx.AsyncClient for this provider, creating it lazily. + The client is transport-aware: for a ``unix://`` platform endpoint it is + bound to the Unix domain socket, otherwise it is the SDK's default TCP + client. Because this cached client is injected into the SDK and + NemoClient factories (which skip their own transport selection when a + client is supplied), building it endpoint-aware here is what makes + service-to-service requests work over UDS. + Each DependencyProvider manages its own HTTP client by default. If you need to share a client across providers (e.g., for connection pooling), you can inject the same client via _http_client. """ - 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 = resolve_platform_endpoint().async_sdk_http_client() + return self._http_client def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: """Return the async platform SDK client. @@ -102,12 +113,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 +182,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..7a905cbb38 --- /dev/null +++ b/packages/nmp_common/tests/client_factory/test_client_factory.py @@ -0,0 +1,394 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for :mod:`nmp.common.client_factory` — the rich NemoClient provider. + +Covers what the platform provider adds over the plugin's env-var default: +per-service URL routing, shared HTTP clients, principal/auth + internal + +OTEL headers, workspace defaults, and test-client injection. +""" + +from unittest.mock import patch + +import httpx +import pytest +from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient +from nemo_platform_plugin.client.types import PreparedRequest +from nemo_platform_plugin.client_provider import NemoClientProvider +from nmp.common import client_factory as cf +from nmp.common.config import Configuration +from nmp.common.observability.otel import scoped_otel_headers + + +@pytest.fixture(autouse=True) +def _reset_client_factory_state(): + """Keep tests order-independent: clear the injected test client and config cache.""" + old = cf._test_http_client + cf._test_http_client = None + Configuration.clear_cache() + try: + yield + finally: + cf._test_http_client = old + Configuration.clear_cache() + + +def _get(path_template: str, **path_params: str) -> PreparedRequest: + return PreparedRequest( + method="GET", + path_template=path_template, + path_params=path_params, + content=None, + content_type=None, + response_type=None, + ) + + +def _mock_client(sink: list[httpx.Request]) -> httpx.Client: + def handler(request: httpx.Request) -> httpx.Response: + sink.append(request) + return httpx.Response(200, json={"ok": True}) + + return httpx.Client(transport=httpx.MockTransport(handler)) + + +# --------------------------------------------------------------------------- +# Sync construction +# --------------------------------------------------------------------------- + + +class TestSyncConstruction: + def test_base_url_from_config(self): + client = cf.get_nemo_client() + assert isinstance(client, NemoClient) + assert client.base_url == str(Configuration.get_platform_config().base_url).rstrip("/") + + def test_service_principal_and_internal_headers(self): + client = cf.get_nemo_client(as_service="evaluator", internal=True) + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert client._default_headers["X-NMP-Internal"] == "true" + + def test_on_behalf_of(self): + client = cf.get_nemo_client(as_service="svc", on_behalf_of="user@example.com") + assert client._default_headers["X-NMP-Principal-On-Behalf-Of"] == "user@example.com" + + def test_workspace_passthrough(self): + client = cf.get_nemo_client(workspace="team-a") + assert client.workspace == "team-a" + + def test_reuses_shared_sync_http_client(self): + client = cf.get_nemo_client() + assert client._http is cf.shared_sync_http_client() + + def test_explicit_http_client_wins(self): + with httpx.Client() as explicit: + client = cf.get_nemo_client(http_client=explicit) + assert client._http is explicit + + +# --------------------------------------------------------------------------- +# Async construction +# --------------------------------------------------------------------------- + + +class TestAsyncConstruction: + def test_base_url_from_config(self): + client = cf.get_async_nemo_client() + assert isinstance(client, AsyncNemoClient) + assert client.base_url == str(Configuration.get_platform_config().base_url).rstrip("/") + + def test_service_principal_and_internal_headers(self): + client = cf.get_async_nemo_client(as_service="evaluator", internal=True) + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert client._default_headers["X-NMP-Internal"] == "true" + + def test_falls_back_to_shared_async_client(self): + client = cf.get_async_nemo_client() + assert isinstance(client._http, httpx.AsyncClient) + + +# --------------------------------------------------------------------------- +# URL routing +# --------------------------------------------------------------------------- + + +class TestUrlRouting: + def test_routes_service_path_to_discovered_origin(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") + monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-svc:9999") + Configuration.clear_cache() + + captured: list[httpx.Request] = [] + client = cf.get_nemo_client(as_service="entities", internal=True, http_client=_mock_client(captured)) + client.send(_get("/apis/entities/v2/foo")) + + assert str(captured[0].url) == "http://entities-svc:9999/apis/entities/v2/foo" + assert captured[0].headers["X-NMP-Principal-Id"] == "service:entities" + assert captured[0].headers["X-NMP-Internal"] == "true" + + def test_preserves_query_string_when_routing(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") + monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-svc:9999") + Configuration.clear_cache() + + captured: list[httpx.Request] = [] + client = cf.get_nemo_client(http_client=_mock_client(captured)) + client.send(_get("/apis/entities/v2/models?limit=5")) + + assert str(captured[0].url) == "http://entities-svc:9999/apis/entities/v2/models?limit=5" + + def test_non_discovered_path_stays_on_platform_origin(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") + monkeypatch.delenv("NMP_MODELS_URL", raising=False) + Configuration.clear_cache() + + captured: list[httpx.Request] = [] + client = cf.get_nemo_client(http_client=_mock_client(captured)) + client.send(_get("/apis/models/v1/bar")) + + assert str(captured[0].url) == "https://nemo-gateway:8080/apis/models/v1/bar" + + def test_workspace_default_fills_path_param(self): + captured: list[httpx.Request] = [] + client = cf.get_nemo_client(workspace="team-a", http_client=_mock_client(captured)) + client.send(_get("/apis/entities/v2/workspaces/{workspace}/models")) + + assert "/workspaces/team-a/models" in str(captured[0].url) + + +# --------------------------------------------------------------------------- +# Headers / auth +# --------------------------------------------------------------------------- + + +class TestHeadersAuth: + def test_propagates_request_principal_when_no_service(self): + auth_headers = {"X-NMP-Principal-Id": "user@example.com", "X-NMP-Principal-Groups": "g1,g2"} + # _get_default_headers reads the request principal via sdk_factory's binding. + with patch("nmp.common.sdk_factory.get_principal_auth_headers", return_value=auth_headers): + client = cf.get_nemo_client() + assert client._default_headers["X-NMP-Principal-Id"] == "user@example.com" + assert client._default_headers["X-NMP-Principal-Groups"] == "g1,g2" + + def test_merges_otel_propagation_headers_without_adding_internal_auth(self): + with scoped_otel_headers({"traceparent": "00-trace-span-01", "X-NMP-Internal": "true"}): + client = cf.get_nemo_client(as_service="svc") + assert client._default_headers["traceparent"] == "00-trace-span-01" + assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" + assert "X-NMP-Internal" not in client._default_headers + + def test_explicit_auth_headers_win_over_conflicting_otel_context(self): + with scoped_otel_headers( + { + "traceparent": "00-trace-span-01", + "x-nmp-principal-id": "attacker@example.com", + "X-NMP-Principal-Groups": "admins", + "x-NMP-Internal": "false", + } + ): + client = cf.get_async_nemo_client(as_service="evaluator", internal=True) + assert client._default_headers["traceparent"] == "00-trace-span-01" + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert client._default_headers["X-NMP-Internal"] == "true" + assert all(name.lower() != "x-nmp-principal-groups" for name in client._default_headers) + + def test_no_headers_leaves_default_headers_none(self): + # No service, no principal context, no OTEL, no internal → no default headers. + with patch("nmp.common.sdk_factory.get_principal_auth_headers", return_value={}): + with patch("nmp.common.sdk_factory.principal_from_env", return_value=None): + client = cf.get_nemo_client() + assert client._default_headers == {} + + +# --------------------------------------------------------------------------- +# Test-client injection +# --------------------------------------------------------------------------- + + +class TestTestClientInjection: + def test_async_uses_module_level_test_client(self): + test_client = httpx.AsyncClient(base_url="http://testserver") + cf._test_http_client = test_client + try: + client = cf.get_async_nemo_client(as_service="evaluator") + assert client._http is test_client + finally: + cf._test_http_client = None + + def test_async_explicit_http_client_beats_module_level(self): + module_client = httpx.AsyncClient(base_url="http://module") + explicit = httpx.AsyncClient(base_url="http://explicit") + cf._test_http_client = module_client + try: + client = cf.get_async_nemo_client(http_client=explicit) + assert client._http is explicit + finally: + cf._test_http_client = None + + +# --------------------------------------------------------------------------- +# Provider class +# --------------------------------------------------------------------------- + + +class TestPlatformNemoClientProvider: + def test_satisfies_protocol(self): + assert isinstance(cf.PlatformNemoClientProvider(), NemoClientProvider) + + def test_get_nemo_client_returns_routed_sync_client(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") + Configuration.clear_cache() + provider = cf.PlatformNemoClientProvider() + client = provider.get_nemo_client(as_service="svc", internal=True, workspace="ws1") + assert isinstance(client, NemoClient) + assert client.base_url == "https://nemo-gateway:8080" + assert client.workspace == "ws1" + assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" + + def test_get_async_nemo_client_returns_async_client(self): + provider = cf.PlatformNemoClientProvider() + client = provider.get_async_nemo_client(as_service="svc") + assert isinstance(client, AsyncNemoClient) + assert client._default_headers["X-NMP-Principal-Id"] == "service:svc" + + +# --------------------------------------------------------------------------- +# Task client: creator delegation (PR-800 claim 1) +# --------------------------------------------------------------------------- + + +class TestTaskClientDelegation: + def test_task_client_delegates_to_job_creator(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv( + "NMP_PRINCIPAL", + '{"id": "user:alice@acme.com", "email": "alice@acme.com", "groups": ["team-a"]}', + ) + client = cf.get_task_nemo_client("evaluator") + headers = client._default_headers + assert headers["X-NMP-Internal"] == "true" + assert headers["X-NMP-Principal-Id"] == "service:evaluator" + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Email"] == "alice@acme.com" + assert headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "team-a" + + def test_task_client_without_principal_warns(self, monkeypatch: pytest.MonkeyPatch, caplog): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + with caplog.at_level("WARNING"): + client = cf.get_task_nemo_client("evaluator") + assert client._default_headers["X-NMP-Principal-Id"] == "service:evaluator" + assert "X-NMP-Principal-On-Behalf-Of" not in client._default_headers + assert "without on-behalf-of delegation" in caplog.text + + async def test_async_task_client_delegates(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv( + "NMP_PRINCIPAL", + '{"id": "user:alice@acme.com", "email": "alice@acme.com", "groups": ["team-a"]}', + ) + client = cf.get_async_task_nemo_client("evaluator") + assert client._default_headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" + + def test_provider_exposes_task_methods(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') + provider = cf.PlatformNemoClientProvider() + headers = provider.get_task_nemo_client("evaluator")._default_headers + assert headers["X-NMP-Principal-On-Behalf-Of"] == "user:alice@acme.com" + + +# --------------------------------------------------------------------------- +# Task client: workload identity (PR-800 claim 2) +# --------------------------------------------------------------------------- + + +class _FakeExchangeProvider: + def get_access_token(self) -> str: + return "exchanged-token" + + async def get_access_token_async(self) -> str: + return "exchanged-token" + + +class TestTaskClientWorkloadIdentity: + @pytest.fixture + def _stub_exchange(self, monkeypatch: pytest.MonkeyPatch): + captured: dict[str, str] = {} + + def _fake(*, base_url, subject_token_file): + captured["base_url"] = base_url + captured["subject_token_file"] = str(subject_token_file) + return _FakeExchangeProvider() + + monkeypatch.setattr( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + _fake, + ) + return captured + + def test_task_client_bootstraps_workload_identity(self, monkeypatch, tmp_path, _stub_exchange): + token_file = tmp_path / "token" + token_file.write_text("subject-token") + monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") + monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') # ignored in WI mode + Configuration.clear_cache() + + client = cf.get_task_nemo_client("evaluator") + assert isinstance(client._auth, _FakeExchangeProvider) + assert _stub_exchange["base_url"] == "http://platform:8080" + # No trusted principal headers in workload-identity mode. + assert "X-NMP-Principal-Id" not in client._default_headers + assert client._default_headers.get("X-NMP-Internal") == "true" + + def test_uds_does_not_bootstrap_workload_identity(self, monkeypatch, tmp_path, _stub_exchange): + # Matches get_task_sdk exactly: with the WI token file set the task path + # delegates to get_nemo_client(internal=True); on UDS transport that skips + # bearer exchange and propagates the env principal as its own identity + # (no service principal, no bearer auth). + token_file = tmp_path / "token" + token_file.write_text("subject-token") + monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_FILE", str(token_file)) + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + monkeypatch.setenv("NMP_PRINCIPAL", '{"id": "user:alice@acme.com"}') + Configuration.clear_cache() + + client = cf.get_task_nemo_client("evaluator") + assert client._auth is None + assert client._default_headers["X-NMP-Principal-Id"] == "user:alice@acme.com" + assert "X-NMP-Principal-On-Behalf-Of" not in client._default_headers + + +# --------------------------------------------------------------------------- +# UDS endpoint routing + transport (PR-800 claim 3) +# --------------------------------------------------------------------------- + + +class TestUdsTransport: + def test_uds_base_url_is_normalized_not_pathed(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + Configuration.clear_cache() + client = cf.get_nemo_client() + # base_url is the routable host, not the raw unix:// socket path. + assert client.base_url == "http://nemo-platform.local" + # concatenating an API path yields a valid URL, not a broken one. + assert client.base_url + "/apis/entities/v2/foo" == "http://nemo-platform.local/apis/entities/v2/foo" + + def test_uds_sync_client_binds_socket_transport(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + Configuration.clear_cache() + client = cf.get_nemo_client() + transport = client._http._transport + assert isinstance(transport, httpx.HTTPTransport) + assert transport._pool._uds == "/tmp/nemo-platform.sock" + + async def test_uds_async_client_binds_socket_transport(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + Configuration.clear_cache() + client = cf.get_async_nemo_client() + transport = client._http._transport + assert isinstance(transport, httpx.AsyncHTTPTransport) + assert transport._pool._uds == "/tmp/nemo-platform.sock" + + def test_tcp_client_uses_shared_client(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") + Configuration.clear_cache() + client = cf.get_nemo_client() + assert client.base_url == "http://platform:8080" diff --git a/packages/nmp_common/tests/nmp_common/test_common_service.py b/packages/nmp_common/tests/nmp_common/test_common_service.py index 45d1bcce6f..4339f7eb3f 100644 --- a/packages/nmp_common/tests/nmp_common/test_common_service.py +++ b/packages/nmp_common/tests/nmp_common/test_common_service.py @@ -5,15 +5,26 @@ import asyncio import threading +import time +from concurrent.futures import ThreadPoolExecutor +from threading import Barrier, Event, Lock +from types import SimpleNamespace 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 fastapi.testclient import TestClient +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.client.client import AsyncNemoClient +from nemo_platform_plugin.dependencies import get_nemo_client as plugin_get_nemo_client from nmp.common.config import PlatformConfig +from nmp.common.observability.otel import scoped_otel_headers from nmp.common.service import DependencyProvider, RouterConfig, Service +from nmp.common.service import __all__ as service_exports +from nmp.common.service import get_nemo_client as facade_get_nemo_client +from nmp.common.service.dependencies import get_nemo_client def _route_paths(app: FastAPI) -> set[str]: @@ -210,6 +221,18 @@ def handler(request: httpx.Request) -> httpx.Response: assert len(requests) == 1 +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.""" @@ -219,6 +242,178 @@ 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 + + endpoint = SimpleNamespace(async_sdk_http_client=lambda: create_transport()) + monkeypatch.setattr(service_base, "resolve_platform_endpoint", lambda: endpoint) + + 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.""" diff --git a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py index bfee97783d..1875d0ff4a 100644 --- a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py +++ b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py @@ -5,24 +5,73 @@ from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import FastAPI from nemo_platform_plugin.entities.client import AsyncEntitiesClient +from nmp.common.config import Configuration from nmp.common.service import DependencyProvider from nmp.common.service.dependencies import get_entity_client, get_platform_config, get_sdk_client -def test_get_http_client_caches_default_client() -> None: +def test_get_http_client_caches_endpoint_client() -> None: provider = DependencyProvider() client = MagicMock() - with patch("nmp.common.service.base.DefaultAsyncHttpxClient", return_value=client) as factory: + with patch( + "nmp.common.service.base.resolve_platform_endpoint", + ) as resolve: + resolve.return_value.async_sdk_http_client.return_value = client first = provider.get_http_client() second = provider.get_http_client() assert first is client assert second is client - factory.assert_called_once_with() + resolve.assert_called_once_with() + resolve.return_value.async_sdk_http_client.assert_called_once_with() + + +def _uds_of(client: httpx.AsyncClient) -> str | None: + """Socket path bound to the client's transport pool, or None for TCP.""" + return getattr(client._transport._pool, "_uds", None) + + +def test_get_http_client_binds_uds_transport(monkeypatch: pytest.MonkeyPatch) -> None: + """Under a unix:// endpoint the provider-owned client must be socket-bound. + + Regression: the provider used to build a plain TCP DefaultAsyncHttpxClient + and inject it into the SDK/NemoClient factories, which then skipped their + own UDS selection, so service-to-service calls went to http://nemo-platform + .local over TCP (_uds=None) and failed. + """ + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + Configuration.clear_cache() + try: + provider = DependencyProvider() + assert _uds_of(provider.get_http_client()) == "/tmp/nemo-platform.sock" + finally: + Configuration.clear_cache() + + +def test_request_scoped_nemo_client_binds_uds_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NMP_BASE_URL", "unix:///tmp/nemo-platform.sock") + Configuration.clear_cache() + try: + provider = DependencyProvider() + client = provider.get_request_scoped_nemo_client() + assert _uds_of(client._http) == "/tmp/nemo-platform.sock" + finally: + Configuration.clear_cache() + + +def test_tcp_endpoint_client_is_not_socket_bound(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NMP_BASE_URL", "http://platform:8080") + Configuration.clear_cache() + try: + provider = DependencyProvider() + assert _uds_of(provider.get_http_client()) is None + finally: + Configuration.clear_cache() def test_get_sdk_client_caches_request_sdk_and_creates_fresh_service_sdk() -> None: @@ -35,11 +84,13 @@ def test_get_sdk_client_caches_request_sdk_and_creates_fresh_service_sdk() -> No assert provider.get_sdk_client() is request_sdk assert provider.get_sdk_client(as_service="jobs") is service_sdk - assert factory.call_args_list[0].kwargs == {"http_client": None} + # The provider now shares its pooled HTTP client with every SDK it builds. + http_client = provider.get_http_client() + assert factory.call_args_list[0].kwargs == {"http_client": http_client} assert factory.call_args_list[1].kwargs == { "as_service": "jobs", "internal": True, - "http_client": None, + "http_client": http_client, } @@ -68,8 +119,10 @@ async def test_close_closes_managed_clients_and_clears_references() -> None: await provider.close() + # close() owns and closes the shared HTTP transport; the SDK borrows that + # transport, so it is dropped without a separate sdk.close(). http_client.aclose.assert_awaited_once_with() - sdk.close.assert_awaited_once_with() + sdk.close.assert_not_awaited() assert provider._http_client is None assert provider._sdk_client is None