diff --git a/packages/nemo_platform_ext/tests/local/test_health_child.py b/packages/nemo_platform_ext/tests/local/test_health_child.py index 04a975d2f0..0830896592 100644 --- a/packages/nemo_platform_ext/tests/local/test_health_child.py +++ b/packages/nemo_platform_ext/tests/local/test_health_child.py @@ -53,7 +53,11 @@ def controller_run(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -87,7 +91,11 @@ def stubborn_controller(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -107,25 +115,21 @@ def stubborn_controller(stop_signal: threading.Event) -> None: @pytest.mark.integration -def test_lifespan_cleanup_runs_on_app_shutdown() -> None: - """``close_shared_http_clients`` should be called during lifespan teardown.""" - cleanup_called = threading.Event() - +def test_lifespan_shutdown_marks_controller_stop_signal() -> None: + """Lifespan teardown should mark the controller stop signal.""" with ( patch("nmp.platform_runner.server.get_platform_config") as mock_pc, patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, - patch("nmp.platform_runner.server.close_shared_http_clients") as mock_close, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") - async def fake_close(): - cleanup_called.set() - - mock_close.side_effect = fake_close - from nmp.platform_runner.server import create_app app = create_app(services=[]) @@ -135,7 +139,7 @@ async def fake_close(): with TestClient(app): pass - assert cleanup_called.is_set(), "close_shared_http_clients was not called during shutdown" + assert app.state.controller_stop_signal.is_set() # --------------------------------------------------------------------------- diff --git a/packages/nemo_platform_ext/tests/local/test_services_contract.py b/packages/nemo_platform_ext/tests/local/test_services_contract.py index acfdc227ca..28d8c34b07 100644 --- a/packages/nemo_platform_ext/tests/local/test_services_contract.py +++ b/packages/nemo_platform_ext/tests/local/test_services_contract.py @@ -274,7 +274,7 @@ def get_routers(self): return [] dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": sidecar_run_func} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": sidecar_run_func} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -298,6 +298,7 @@ def get_routers(self): monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) diff --git a/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py b/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py index 33cadabaa7..85d4d7a642 100644 --- a/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py +++ b/packages/nemo_platform_ext/tests/local/test_sidecar_integration.py @@ -42,7 +42,7 @@ def get_routers(self): def _sidecar_with_events(started: threading.Event, stopped: threading.Event) -> Callable[[threading.Event], None]: - """Return a sidecar ``run(stop_signal)`` that signals start/stop via events.""" + """Return a sidecar run function that signals start/stop via events.""" def run(stop_signal: threading.Event) -> None: started.set() @@ -71,7 +71,7 @@ def patched_registry( test sidecar, plus minimal auth/platform config stubs.""" started, stopped = sidecar_events dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": _sidecar_with_events(started, stopped)} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": _sidecar_with_events(started, stopped)} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -95,6 +95,7 @@ def patched_registry( monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py index d67180020d..afd266570a 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/adapter.py @@ -17,10 +17,8 @@ def make_sync_resource(platform: NeMoPlatform) -> NemoClient: from __future__ import annotations -from collections.abc import Callable from typing import TypeVar, overload -import httpx from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient from nemo_platform_plugin.client.types import RetryPolicy @@ -29,14 +27,6 @@ def make_sync_resource(platform: NeMoPlatform) -> NemoClient: AsyncT = TypeVar("AsyncT", bound=AsyncNemoClient) -def _url_resolver_from_platform(platform: NeMoPlatform | AsyncNeMoPlatform) -> Callable[[str], str | httpx.URL]: - router = getattr(platform, "_nmp_request_router", None) - resolver = getattr(router, "resolve", None) - if resolver is not None: - return resolver - return platform._prepare_url - - @overload def client_from_platform(platform: NeMoPlatform, client_cls: type[SyncT]) -> SyncT: ... @overload @@ -62,7 +52,7 @@ def client_from_platform( headers = {k: v for k, v in platform._client.headers.items() if k.lower() not in _skip} # type: ignore[union-attr] retry = RetryPolicy(max_retries=platform.max_retries) - url_resolver = _url_resolver_from_platform(platform) + url_resolver = platform._prepare_url if isinstance(platform, AsyncNeMoPlatform): if not issubclass(client_cls, AsyncNemoClient): raise TypeError("AsyncNeMoPlatform requires an AsyncNemoClient class") 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..8f37d1a8f2 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/sdk_provider.py @@ -7,8 +7,8 @@ Plugin authors use :func:`get_task_sdk` in their ``__main__.py`` entrypoints instead of importing from ``nmp.common.sdk_factory``. This keeps the ``nemo-platform-plugin`` package free of ``nmp-common`` dependencies while -still allowing the platform to register a richer provider (with URL routing, -shared HTTP clients, OTEL headers, etc.) when ``nmp-common`` is installed. +still allowing the platform to register a richer provider (with platform auth +context, OTEL headers, etc.) when ``nmp-common`` is installed. Lookup order for the provider ----------------------------- @@ -162,6 +162,21 @@ def _workload_identity_headers(*, internal: bool) -> dict[str, str]: return {_INTERNAL_REQUEST_HEADER: "true"} if internal else {} +def _task_on_behalf_of_headers() -> dict[str, str] | None: + principal = _read_principal_from_env() + return _on_behalf_of_headers(principal) if principal is not None else None + + +def _warn_missing_task_principal(*, service_name: str, async_sdk: bool) -> None: + qualifier = "async task SDK" if async_sdk else "task SDK" + logger.warning( + "%s not set; %s will authenticate as service:%s without on-behalf-of delegation", + _NMP_PRINCIPAL_ENVVAR, + qualifier, + service_name, + ) + + class DefaultSDKProvider: """Env-var-based provider that ships with the plugin package. @@ -171,59 +186,27 @@ class DefaultSDKProvider: """ def get_task_sdk(self, service_name: str) -> NeMoPlatform: - if is_workload_identity_token_file_set(): - return NeMoPlatform( - base_url=self._base_url(), - default_headers=_workload_identity_headers(internal=True), - ) - - headers: dict[str, str] = { - "X-NMP-Principal-Id": f"service:{service_name}", - _INTERNAL_REQUEST_HEADER: "true", - } - - principal = _read_principal_from_env() - if principal is not None: - headers.update(_on_behalf_of_headers(principal)) - else: - logger.warning( - "%s not set; task SDK will authenticate as service:%s without on-behalf-of delegation", - _NMP_PRINCIPAL_ENVVAR, - service_name, - ) - - return NeMoPlatform( - base_url=self._base_url(), - default_headers=headers, + on_behalf_of_headers = _task_on_behalf_of_headers() + if on_behalf_of_headers is None and not is_workload_identity_token_file_set(): + _warn_missing_task_principal(service_name=service_name, async_sdk=False) + return self._make_sdk( + NeMoPlatform, + as_service=service_name, + internal=True, + on_behalf_of_headers=on_behalf_of_headers, ) def get_async_task_sdk(self, service_name: str) -> AsyncNeMoPlatform: # Async mirror of get_task_sdk: identical headers (service principal, # internal marker, and full on-behalf-of id/email/groups), async client. - if is_workload_identity_token_file_set(): - return AsyncNeMoPlatform( - base_url=self._base_url(), - default_headers=_workload_identity_headers(internal=True), - ) - - headers: dict[str, str] = { - "X-NMP-Principal-Id": f"service:{service_name}", - _INTERNAL_REQUEST_HEADER: "true", - } - - principal = _read_principal_from_env() - if principal is not None: - headers.update(_on_behalf_of_headers(principal)) - else: - logger.warning( - "%s not set; async task SDK will authenticate as service:%s without on-behalf-of delegation", - _NMP_PRINCIPAL_ENVVAR, - service_name, - ) - - return AsyncNeMoPlatform( - base_url=self._base_url(), - default_headers=headers, + on_behalf_of_headers = _task_on_behalf_of_headers() + if on_behalf_of_headers is None and not is_workload_identity_token_file_set(): + _warn_missing_task_principal(service_name=service_name, async_sdk=True) + return self._make_sdk( + AsyncNeMoPlatform, + as_service=service_name, + internal=True, + on_behalf_of_headers=on_behalf_of_headers, ) def _make_sdk( @@ -233,12 +216,15 @@ def _make_sdk( as_service: str | None = None, internal: bool = False, on_behalf_of: str | None = None, + on_behalf_of_headers: dict[str, str] | None = None, ) -> _SDKT: - if as_service is None and on_behalf_of is None and is_workload_identity_token_file_set(): + if is_workload_identity_token_file_set(): headers = _workload_identity_headers(internal=internal) return cls(base_url=self._base_url(), default_headers=headers or None) headers = self._build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) + if on_behalf_of_headers: + headers.update(on_behalf_of_headers) return cls(base_url=self._base_url(), default_headers=headers or None) def get_platform_sdk( @@ -372,9 +358,8 @@ def get_async_task_sdk(service_name: str) -> AsyncNeMoPlatform: ``NMP_PRINCIPAL`` is set, on behalf of the job creator with the full delegated identity (on-behalf-of id, email, and groups) — wire-identical to :func:`get_task_sdk`. - A dedicated provider method (not a wrapper over :func:`get_async_platform_sdk`) so each provider - mirrors its own sync :meth:`SDKProvider.get_task_sdk` exactly; the platform provider routes URLs - and reuses its shared async client, the default provider uses env-var headers. + Delegates to the active provider so each provider can preserve its own SDK + construction lifecycle while keeping task and platform SDK auth policy aligned. """ return _resolve_provider().get_async_task_sdk(service_name) diff --git a/packages/nemo_platform_plugin/tests/client/test_adapter.py b/packages/nemo_platform_plugin/tests/client/test_adapter.py index e51c4fd19d..ad995ce15b 100644 --- a/packages/nemo_platform_plugin/tests/client/test_adapter.py +++ b/packages/nemo_platform_plugin/tests/client/test_adapter.py @@ -26,55 +26,18 @@ def test_client_from_platform_preserves_retry_count_with_nemoclient_defaults() - assert client.retry.retryable_status_codes == (502, 503, 504, 429) -def test_client_from_platform_prefers_platform_request_router() -> None: - http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) - platform = NeMoPlatform( - base_url="http://gateway", - workspace="default", - http_client=http_client, - ) - - class RequestRouter: - def resolve(self, url: str) -> str: - return url.replace("http://gateway/apis/jobs", "http://127.0.0.1:8080/apis/jobs") - - platform._nmp_request_router = RequestRouter() # type: ignore[attr-defined] - - client = client_from_platform(platform, JobsClient) - - request = endpoints.list_steps(workspace="default", name="job-1") - assert client._resolve_path(request) == ("http://127.0.0.1:8080/apis/jobs/v2/workspaces/default/jobs/job-1/steps") - - -def test_client_from_platform_falls_back_to_sdk_prepare_url() -> None: - http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) - platform = NeMoPlatform( - base_url="http://gateway", - workspace="default", - http_client=http_client, - ) - - def prepare_url(url: str) -> str: - return url.replace("http://gateway/apis/jobs", "http://127.0.0.1:8080/apis/jobs") - - platform._prepare_url = prepare_url # type: ignore[method-assign] - - client = client_from_platform(platform, JobsClient) - - request = endpoints.list_steps(workspace="default", name="job-1") - assert client._resolve_path(request) == ("http://127.0.0.1:8080/apis/jobs/v2/workspaces/default/jobs/job-1/steps") - - def test_from_client_preserves_url_resolver() -> None: http_client = httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request))) client = JobsClient( base_url="http://gateway", workspace="default", http_client=http_client, - url_resolver=lambda url: url.replace("http://gateway/apis/jobs", "http://127.0.0.1:8080/apis/jobs"), + url_resolver=lambda url: str(httpx.URL(url).copy_with(host="resolved.example.test")), ) clone = JobsClient.from_client(client) request = endpoints.list_steps(workspace="default", name="job-1") - assert clone._resolve_path(request) == ("http://127.0.0.1:8080/apis/jobs/v2/workspaces/default/jobs/job-1/steps") + assert ( + clone._resolve_path(request) == "http://resolved.example.test/apis/jobs/v2/workspaces/default/jobs/job-1/steps" + ) diff --git a/packages/nemo_platform_plugin/tests/test_sdk_provider.py b/packages/nemo_platform_plugin/tests/test_sdk_provider.py index cc1a1ec77e..4edbec891e 100644 --- a/packages/nemo_platform_plugin/tests/test_sdk_provider.py +++ b/packages/nemo_platform_plugin/tests/test_sdk_provider.py @@ -6,6 +6,7 @@ from __future__ import annotations import json +import logging from unittest.mock import patch import pytest @@ -128,15 +129,13 @@ def test_get_task_sdk_default_base_url(self, monkeypatch): sdk = provider.get_task_sdk("test") assert sdk.base_url == "http://localhost:8080" - def test_get_task_sdk_uses_workload_identity_when_token_file_configured(self, monkeypatch, tmp_path): + def test_get_task_sdk_uses_workload_identity_when_token_file_configured(self, monkeypatch, tmp_path, caplog): subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv( - "NMP_PRINCIPAL", - json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), - ) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + caplog.set_level(logging.WARNING, logger="nemo_platform_plugin.sdk_provider") provider = DefaultSDKProvider() sdk = provider.get_task_sdk("evaluator") @@ -146,6 +145,7 @@ def test_get_task_sdk_uses_workload_identity_when_token_file_configured(self, mo assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers finally: sdk.close() + assert "will authenticate as service:evaluator without on-behalf-of delegation" not in caplog.text def test_get_platform_sdk_as_service(self, monkeypatch): monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") @@ -157,6 +157,27 @@ def test_get_platform_sdk_as_service(self, monkeypatch): assert sdk.default_headers["X-NMP-Principal-Id"] == "service:my-svc" assert sdk.default_headers["X-NMP-Internal"] == "true" + def test_get_platform_sdk_uses_workload_identity_with_service_when_token_file_configured( + self, monkeypatch, tmp_path + ): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv( + "NMP_PRINCIPAL", + json.dumps({"id": "creator@ex.com", "email": "creator@ex.com", "groups": ["team"]}), + ) + + provider = DefaultSDKProvider() + sdk = provider.get_platform_sdk(as_service="my-svc", internal=True) + try: + assert sdk.default_headers["X-NMP-Internal"] == "true" + assert "X-NMP-Principal-Id" not in sdk.default_headers + assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers + finally: + sdk.close() + def test_get_platform_sdk_on_behalf_of(self, monkeypatch): monkeypatch.setenv("NMP_BASE_URL", "http://test:9090") monkeypatch.delenv("NMP_PRINCIPAL", raising=False) diff --git a/packages/nmp_common/src/nmp/common/entities/client.py b/packages/nmp_common/src/nmp/common/entities/client.py index f2214cb3d3..5f901b7b43 100644 --- a/packages/nmp_common/src/nmp/common/entities/client.py +++ b/packages/nmp_common/src/nmp/common/entities/client.py @@ -47,11 +47,11 @@ def as_service(self, service_name: str, *, internal: bool = False) -> "EntityCli """ from nemo_platform.resources.entities import AsyncEntitiesResource from nmp.common.observability import MARK_INTERNAL_REQUEST_HEADERS - from nmp.common.sdk_factory import with_options_preserving_request_router + from nmp.common.sdk_factory import with_options_reusing_http_client underlying_sdk = self.entities_api._client headers: dict[str, str] = {"X-NMP-Principal-Id": f"service:{service_name}"} if internal: headers.update(MARK_INTERNAL_REQUEST_HEADERS) - service_sdk = with_options_preserving_request_router(underlying_sdk, set_default_headers=headers) + service_sdk = with_options_reusing_http_client(underlying_sdk, set_default_headers=headers) return EntityClient(AsyncEntitiesResource(service_sdk)) diff --git a/packages/nmp_common/src/nmp/common/http_clients.py b/packages/nmp_common/src/nmp/common/http_clients.py deleted file mode 100644 index 0a8934bb01..0000000000 --- a/packages/nmp_common/src/nmp/common/http_clients.py +++ /dev/null @@ -1,107 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Shared HTTP clients for SDK and internal requests. - -Ideally, most consumers get their HTTP client from DependencyProvider, which -manages per-service client lifecycle. However, code often constructs SDKs in -isolation (tasks, controllers, background jobs) without access to a -DependencyProvider. These shared clients prevent each SDK instantiation from -creating a new HTTP client, avoiding connection pool and SSL context overhead. - -Used by: -- get_platform_sdk() / get_async_platform_sdk() for tasks, controllers, and - other code outside DependencyProvider context -- Platform shutdown cleanup via close_cached_http_clients() - -The wrapper classes make close()/aclose() a no-op so that SDK code can safely -call sdk.close() without affecting other users of the shared client. Actual -cleanup happens at shutdown via close_cached_http_clients(). -""" - -import asyncio -from functools import cache -from typing import cast - -import httpx -from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient - - -class _SharedAsyncHttpClient(DefaultAsyncHttpxClient): - """Shared async HTTP client that ignores aclose() calls.""" - - async def aclose(self) -> None: - pass - - async def _real_close(self) -> None: - await super().aclose() - - -class _SharedSyncHttpClient(DefaultHttpxClient): - """Shared sync HTTP client that ignores close() calls.""" - - def close(self) -> None: - pass - - def _real_close(self) -> None: - super().close() - - -_shared_async_http_clients: dict[asyncio.AbstractEventLoop, _SharedAsyncHttpClient] = {} - - -def shared_async_http_client() -> httpx.AsyncClient: - """Get the shared async HTTP client for SDK requests. - - Returns a cached async HTTP client scoped to the current event loop. - Use this when creating SDK instances that need to share a connection - pool and SSL context within a single loop. - - If called when no event loop is currently running, return a fresh async - client instead of caching one globally. This keeps sync setup code safe - without binding a shared client to an arbitrary thread-local loop. - - The returned client ignores aclose() calls - cleanup happens at shutdown - via close_cached_http_clients(). This allows SDKs to safely call close() - without breaking other users of the shared client. - """ - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return DefaultAsyncHttpxClient() - client = _shared_async_http_clients.get(loop) - if client is None or client.is_closed: - client = _SharedAsyncHttpClient() - _shared_async_http_clients[loop] = client - return client - - -@cache -def shared_sync_http_client() -> httpx.Client: - """Get the shared sync HTTP client for SDK requests. - - Returns a cached sync HTTP client with SDK-compatible defaults. - Use this when creating SDK instances that need to share a global - connection pool and SSL context. - - The returned client ignores close() calls - cleanup happens at shutdown - via close_cached_http_clients(). This allows SDKs to safely call close() - without breaking other users of the shared client. - """ - return _SharedSyncHttpClient() - - -async def close_shared_http_clients() -> None: - """Close shared HTTP clients during graceful shutdown. - - Called from the platform lifespan shutdown to clean up module-level - shared clients used by non-service code (tasks, jobs, legacy callers). - """ - loop = asyncio.get_running_loop() - async_client = _shared_async_http_clients.pop(loop, None) - if async_client is not None and not async_client.is_closed: - await async_client._real_close() - - sync_client = cast(_SharedSyncHttpClient, shared_sync_http_client()) - sync_client._real_close() - shared_sync_http_client.cache_clear() diff --git a/packages/nmp_common/src/nmp/common/immutable_http_client.py b/packages/nmp_common/src/nmp/common/immutable_http_client.py new file mode 100644 index 0000000000..ed29d7623f --- /dev/null +++ b/packages/nmp_common/src/nmp/common/immutable_http_client.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Immutability helpers for SDK-owned HTTP clients. + +NeMo Platform SDK instances reuse their underlying httpx client when callers +derive scoped SDKs via ``with_options()`` or pass the SDK into typed plugin +clients. Those clients must be created with their required transport-level +configuration and then left alone. + +Caller-specific request configuration, including auth headers, belongs on the +SDK instance or on a separate explicit client. Mutating an SDK-owned client +after it has been handed out is a bug because that state can leak into derived +SDKs or requests. These wrappers make those bugs fail immediately. +""" + +from types import MappingProxyType +from typing import NoReturn + +import httpx +from httpx._types import CookieTypes, HeaderTypes +from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient + +_IMMUTABLE_CLIENT_ATTRS = { + "_base_url", + "_cookies", + "_event_hooks", + "_headers", + "_params", + "_timeout", + "_trust_env", + "base_url", + "cookies", + "event_hooks", + "follow_redirects", + "headers", + "max_redirects", + "params", + "timeout", + "trust_env", +} + + +def _raise_immutable_client_mutation_error() -> NoReturn: + raise TypeError( + "SDK HTTP clients are immutable. Pass per-SDK options as SDK constructor " + "arguments, or pass a separate httpx client configured for that use case." + ) + + +class _ImmutableHeaders(httpx.Headers): + def __setitem__(self, key: str, value: str) -> None: + _raise_immutable_client_mutation_error() + + def __delitem__(self, key: str) -> None: + _raise_immutable_client_mutation_error() + + def clear(self) -> None: + _raise_immutable_client_mutation_error() + + def pop(self, key: str, default: object = None) -> str: + _raise_immutable_client_mutation_error() + + def popitem(self) -> tuple[str, str]: + _raise_immutable_client_mutation_error() + + def setdefault(self, key: str, default: str = "") -> str: + _raise_immutable_client_mutation_error() + + def update(self, headers: HeaderTypes | None = None) -> None: + _raise_immutable_client_mutation_error() + + +class _ImmutableCookies(httpx.Cookies): + def extract_cookies(self, response: httpx.Response) -> None: + # httpx normally persists response cookies on the client. SDK clients + # should not carry request/session state between derived SDK handles. + pass + + def set(self, name: str, value: str, domain: str = "", path: str = "/") -> None: + _raise_immutable_client_mutation_error() + + def delete( + self, + name: str, + domain: str | None = None, + path: str | None = None, + ) -> None: + _raise_immutable_client_mutation_error() + + def clear(self, domain: str | None = None, path: str | None = None) -> None: + _raise_immutable_client_mutation_error() + + def update(self, cookies: CookieTypes | None = None) -> None: + _raise_immutable_client_mutation_error() + + def __setitem__(self, name: str, value: str) -> None: + _raise_immutable_client_mutation_error() + + def __delitem__(self, name: str) -> None: + _raise_immutable_client_mutation_error() + + +class ImmutableHttpClientMixin: + """Mixin for httpx client subclasses that are immutable after construction.""" + + _immutable_http_client_frozen = False + + def __setattr__(self, name: str, value: object) -> None: + if self._immutable_http_client_frozen and name in _IMMUTABLE_CLIENT_ATTRS: + raise AttributeError( + "SDK HTTP clients are immutable. Pass a separate httpx client " + "when client-level configuration needs to differ." + ) + super().__setattr__(name, value) + + def _freeze_http_client(self) -> None: + # Assignment blocking is not enough because these attributes are mutable + # containers. Replace them with immutable versions before sharing the + # client through SDK clones or plugin adapters. + self._headers = _ImmutableHeaders(self._headers) + self._cookies = _ImmutableCookies(self._cookies) + self._event_hooks = MappingProxyType( + {hook_name: tuple(hooks) for hook_name, hooks in self._event_hooks.items()} + ) + self._immutable_http_client_frozen = True + + +class ImmutableHttpxClient(ImmutableHttpClientMixin, httpx.Client): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableAsyncHttpxClient(ImmutableHttpClientMixin, httpx.AsyncClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableDefaultHttpxClient(ImmutableHttpClientMixin, DefaultHttpxClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() + + +class ImmutableDefaultAsyncHttpxClient(ImmutableHttpClientMixin, DefaultAsyncHttpxClient): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self._freeze_http_client() diff --git a/packages/nmp_common/src/nmp/common/platform_endpoint.py b/packages/nmp_common/src/nmp/common/platform_endpoint.py index b8392881c8..6dc069f7ca 100644 --- a/packages/nmp_common/src/nmp/common/platform_endpoint.py +++ b/packages/nmp_common/src/nmp/common/platform_endpoint.py @@ -5,17 +5,34 @@ from __future__ import annotations +import logging import os -from dataclasses import dataclass +import re +import threading +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field from pathlib import Path -from typing import Literal +from types import MappingProxyType +from typing import Generic, Literal, TypeVar import httpx from httpx._types import TimeoutTypes -from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient -from nmp.common.config import PlatformConfig +from nmp.common.config import PlatformConfig, get_platform_config +from nmp.common.immutable_http_client import ( + ImmutableAsyncHttpxClient, + ImmutableDefaultAsyncHttpxClient, + ImmutableDefaultHttpxClient, + ImmutableHttpxClient, +) UDS_BASE_URL = "http://nemo-platform.local" +logger = logging.getLogger(__name__) + +_PLATFORM_SERVICE_URL_ENV_PATTERN = re.compile(r"^NMP_([A-Z0-9_]+)_URL$") + + +def _empty_service_endpoints() -> Mapping[str, "PlatformEndpoint"]: + return MappingProxyType({}) @dataclass(frozen=True) @@ -23,6 +40,12 @@ class PlatformEndpoint: connect_base_url: str socket_path: Path | None transport: Literal["tcp", "uds"] + service_pattern: re.Pattern[str] | None = field(default=None, repr=False, compare=False) + service_endpoints: Mapping[str, "PlatformEndpoint"] = field( + default_factory=_empty_service_endpoints, + repr=False, + compare=False, + ) def sync_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Client: if self.transport == "uds": @@ -48,38 +71,106 @@ def async_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Asy return httpx.AsyncClient(follow_redirects=True) return httpx.AsyncClient(follow_redirects=True, timeout=timeout) - def sync_sdk_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Client: + def sync_sdk_http_client( + self, + *, + timeout: TimeoutTypes | None = None, + ) -> httpx.Client: + if self.service_endpoints: + transport = _SyncPlatformEndpointRoutingTransport(endpoint=self) + if timeout is None: + return ImmutableDefaultHttpxClient(transport=transport) + return ImmutableDefaultHttpxClient(transport=transport, timeout=timeout) if self.transport == "uds": - return self.sync_http_client(timeout=timeout) + if self.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + transport = httpx.HTTPTransport(uds=str(self.socket_path)) + if timeout is None: + return ImmutableHttpxClient(transport=transport, follow_redirects=True) + return ImmutableHttpxClient(transport=transport, follow_redirects=True, timeout=timeout) if timeout is None: - return DefaultHttpxClient() - return DefaultHttpxClient(timeout=timeout) + return ImmutableDefaultHttpxClient() + return ImmutableDefaultHttpxClient(timeout=timeout) - def async_sdk_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.AsyncClient: + def async_sdk_http_client( + self, + *, + timeout: TimeoutTypes | None = None, + ) -> httpx.AsyncClient: + if self.service_endpoints: + transport = _AsyncPlatformEndpointRoutingTransport(endpoint=self) + if timeout is None: + return ImmutableDefaultAsyncHttpxClient(transport=transport) + return ImmutableDefaultAsyncHttpxClient(transport=transport, timeout=timeout) if self.transport == "uds": - return self.async_http_client(timeout=timeout) + if self.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + transport = httpx.AsyncHTTPTransport(uds=str(self.socket_path)) + if timeout is None: + return ImmutableAsyncHttpxClient(transport=transport, follow_redirects=True) + return ImmutableAsyncHttpxClient(transport=transport, follow_redirects=True, timeout=timeout) if timeout is None: - return DefaultAsyncHttpxClient() - return DefaultAsyncHttpxClient(timeout=timeout) + return ImmutableDefaultAsyncHttpxClient() + return ImmutableDefaultAsyncHttpxClient(timeout=timeout) + + def route_request_url(self, url: str | httpx.URL) -> "RoutedPlatformEndpointRequest": + """Resolve one outgoing SDK URL using this endpoint's fixed routing table.""" + + request_url = httpx.URL(url) + endpoint = self + service_name = "unknown" + + match = self.service_pattern.search(request_url.path) if self.service_pattern is not None else None + if match is not None: + service_name = match.group(1) + endpoint = self.service_endpoints.get(service_name, self) + + routed_url = _url_for_endpoint(request_url, endpoint) + logger.debug( + "Routing SDK URL to service endpoint" + if service_name != "unknown" + else "Routing SDK URL to default endpoint", + extra={ + "service": service_name, + "path": request_url.path, + "host": routed_url.host, + "port": routed_url.port, + "transport": endpoint.transport, + }, + ) + return RoutedPlatformEndpointRequest(url=routed_url, endpoint=endpoint) + + +@dataclass(frozen=True) +class RoutedPlatformEndpointRequest: + url: httpx.URL + endpoint: PlatformEndpoint def resolve_platform_endpoint(platform_config: PlatformConfig | None = None) -> PlatformEndpoint: """Resolve the default platform endpoint from ``NMP_BASE_URL`` / config.""" if platform_config is None: - from nmp.common.config import Configuration - - platform_config = Configuration.get_platform_config() - return parse_platform_endpoint(platform_config.base_url) + platform_config = get_platform_config() + default_endpoint = parse_platform_endpoint(platform_config.base_url) + service_endpoints = { + service_name: resolve_service_endpoint(service_name, platform_config) + for service_name in sorted(_service_route_names(platform_config)) + } + return PlatformEndpoint( + connect_base_url=default_endpoint.connect_base_url, + socket_path=default_endpoint.socket_path, + transport=default_endpoint.transport, + service_pattern=platform_config.create_service_pattern(), + service_endpoints=MappingProxyType(service_endpoints), + ) def resolve_service_endpoint(service_name: str, platform_config: PlatformConfig | None = None) -> PlatformEndpoint: """Resolve a service endpoint using ``NMP__URL`` before ``NMP_BASE_URL``.""" if platform_config is None: - from nmp.common.config import Configuration - - platform_config = Configuration.get_platform_config() + platform_config = get_platform_config() env_name = f"NMP_{service_name.upper().replace('-', '_')}_URL" endpoint = os.environ.get(env_name) or platform_config.get_service_url(service_name) return parse_platform_endpoint(endpoint) @@ -109,3 +200,121 @@ def _parse_unix_socket_path(endpoint: str) -> Path: if not raw_path.startswith("/"): raise ValueError(f"UDS endpoint must use an absolute socket path, got {endpoint!r}") return Path(raw_path) + + +def _service_route_names(platform_config: PlatformConfig) -> set[str]: + names = {_normalize_service_name(service_name) for service_name in platform_config.service_discovery} + names.update(_normalize_service_name(service_name) for service_name in platform_config.get_services()) + names.discard("") + names.update(_service_url_env_names(names)) + return names + + +def _service_url_env_names(known_service_names: set[str]) -> set[str]: + names: set[str] = set() + for key, value in os.environ.items(): + if not value: + continue + match = _PLATFORM_SERVICE_URL_ENV_PATTERN.match(key) + if match is None: + continue + service_name = match.group(1) + if service_name == "BASE": + continue + normalized_service_name = _normalize_service_name(service_name) + if normalized_service_name not in known_service_names: + continue + names.add(normalized_service_name) + return names + + +def _normalize_service_name(service_name: str) -> str: + return service_name.strip().lower().replace("_", "-") + + +def _url_for_endpoint(url: httpx.URL, endpoint: PlatformEndpoint) -> httpx.URL: + if endpoint.transport == "uds": + return url.copy_with(scheme="http", host="nemo-platform.local", port=None) + endpoint_url = httpx.URL(endpoint.connect_base_url) + path = url.path + if endpoint_url.path not in ("", "/"): + path = f"{endpoint_url.path.rstrip('/')}/{url.path.lstrip('/')}" + return url.copy_with(scheme=endpoint_url.scheme, host=endpoint_url.host, port=endpoint_url.port, path=path) + + +def _set_request_url(request: httpx.Request, url: httpx.URL) -> None: + if request.url == url: + return + request.url = url + if url.host: + request.headers["Host"] = url.netloc.decode("ascii") + + +_TransportT = TypeVar("_TransportT") + + +class _ThreadSafeUDSTransportCache(Generic[_TransportT]): + def __init__(self, factory: Callable[..., _TransportT]) -> None: + self._factory = factory + self._transports: dict[Path, _TransportT] = {} + self._lock = threading.Lock() + + def get(self, socket_path: Path) -> _TransportT: + with self._lock: + transport = self._transports.get(socket_path) + if transport is None: + transport = self._factory(uds=str(socket_path)) + self._transports[socket_path] = transport + return transport + + def values(self) -> tuple[_TransportT, ...]: + with self._lock: + return tuple(self._transports.values()) + + +class _SyncPlatformEndpointRoutingTransport(httpx.BaseTransport): + def __init__(self, *, endpoint: PlatformEndpoint) -> None: + self._endpoint = endpoint + self._tcp_transport = httpx.HTTPTransport() + self._uds_transports = _ThreadSafeUDSTransportCache(httpx.HTTPTransport) + + def handle_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return self._transport_for_endpoint(routed.endpoint).handle_request(request) + + def close(self) -> None: + self._tcp_transport.close() + for transport in self._uds_transports.values(): + transport.close() + + def _transport_for_endpoint(self, endpoint: PlatformEndpoint) -> httpx.HTTPTransport: + if endpoint.transport == "tcp": + return self._tcp_transport + if endpoint.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + return self._uds_transports.get(endpoint.socket_path) + + +class _AsyncPlatformEndpointRoutingTransport(httpx.AsyncBaseTransport): + def __init__(self, *, endpoint: PlatformEndpoint) -> None: + self._endpoint = endpoint + self._tcp_transport = httpx.AsyncHTTPTransport() + self._uds_transports = _ThreadSafeUDSTransportCache(httpx.AsyncHTTPTransport) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + routed = self._endpoint.route_request_url(request.url) + _set_request_url(request, routed.url) + return await self._transport_for_endpoint(routed.endpoint).handle_async_request(request) + + async def aclose(self) -> None: + await self._tcp_transport.aclose() + for transport in self._uds_transports.values(): + await transport.aclose() + + def _transport_for_endpoint(self, endpoint: PlatformEndpoint) -> httpx.AsyncHTTPTransport: + if endpoint.transport == "tcp": + return self._tcp_transport + if endpoint.socket_path is None: + raise ValueError("UDS endpoint is missing a socket path") + return self._uds_transports.get(endpoint.socket_path) diff --git a/packages/nmp_common/src/nmp/common/sdk_factory.py b/packages/nmp_common/src/nmp/common/sdk_factory.py index 57a95891a7..4b79c72bfb 100644 --- a/packages/nmp_common/src/nmp/common/sdk_factory.py +++ b/packages/nmp_common/src/nmp/common/sdk_factory.py @@ -4,177 +4,202 @@ """SDK factory functions for creating NeMo Platform SDK instances.""" import logging +from collections.abc import AsyncGenerator, Generator, Mapping from dataclasses import dataclass -from typing import Any, Callable, Optional, TypeVar, cast +from typing import Any, Generic, Optional, Protocol, TypeVar import httpx from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform_plugin.client.constants import is_workload_identity_token_file_set from nmp.common.auth import Principal, get_principal_auth_headers, principal_from_env -from nmp.common.config import Configuration, PlatformConfig -from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client +from nmp.common.config import get_platform_config +from nmp.common.immutable_http_client import ImmutableDefaultAsyncHttpxClient, ImmutableDefaultHttpxClient from nmp.common.observability import MARK_INTERNAL_REQUEST_HEADERS from nmp.common.observability.otel import get_otel_headers -from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint, resolve_service_endpoint +from nmp.common.platform_endpoint import resolve_platform_endpoint logger = logging.getLogger(__name__) -PlatformSDKT = TypeVar("PlatformSDKT", NeMoPlatform, AsyncNeMoPlatform) +PlatformSDKT = TypeVar("PlatformSDKT", bound=NeMoPlatform | AsyncNeMoPlatform) +_HTTPClientT = TypeVar("_HTTPClientT", httpx.Client, httpx.AsyncClient) -# Test-only: HTTP clients to use for SDK requests in test context. -# Set by test fixtures to route requests through the in-process test transport. -# -# TODO: Remove these module-level variables once all direct get_platform_sdk() / -# get_async_platform_sdk() callers are migrated to use DependencyProvider. See -# architecture/docs/http-client-injection.md for migration path and best practices. -_test_http_client: Optional[httpx.AsyncClient] = None +class _AccessTokenProvider(Protocol): + def get_access_token(self) -> str: ... -def _base_url_from_config() -> str: - return Configuration.get_platform_config().base_url + async def get_access_token_async(self) -> str: ... -def resolve_platform_request_url( - url: str, - *, - platform_config: PlatformConfig, - default_resolver: Callable[[str], httpx.URL], -) -> httpx.URL: - """Resolve the destination URL for an SDK request. - - The generated SDK builds requests from relative paths like - ``/apis/entities/v2/...`` and then calls its private ``_prepare_url`` hook. - Keep the routing policy here, not in that private hook: - - - resolve the SDK URL normally against ``platform.base_url``; - - if the path targets ``/apis/{api_name}/...``, replace only the origin - with ``platform.get_service_url(api_name)``; - - preserve the original path and query string. - """ - request_url = default_resolver(url) - service_pattern = platform_config.create_service_pattern() - if service_pattern is None: - return request_url - match = service_pattern.search(request_url.path) - if match is None: - logger.debug( - "Routing URL to original URL", - extra={"service": "unknown", "path": request_url.path, "host": request_url.host, "port": request_url.port}, - ) - return request_url - - api_name = match.group(1) - svc_endpoint = resolve_service_endpoint(api_name, platform_config) - if svc_endpoint.transport == "uds": - logger.debug( - "Routing URL to UDS service", - extra={ - "service": api_name, - "path": request_url.path, - "transport": svc_endpoint.transport, - }, - ) - return request_url.copy_with(scheme="http", host="nemo-platform.local", port=None) - service_url = httpx.URL(svc_endpoint.connect_base_url) - routed_url = request_url.copy_with( - scheme=service_url.scheme, - host=service_url.host, - port=service_url.port, - ) - logger.debug( - "Routing URL to service URL", - extra={ - "service": api_name, - "path": request_url.path, - "host": routed_url.host, - "port": routed_url.port, - }, - ) - return routed_url +def _base_url_from_config() -> str: + return get_platform_config().base_url @dataclass(frozen=True) -class PlatformRequestRouter: - """Routes SDK requests to the platform gateway or a service-specific origin.""" +class _SDKConnection(Generic[_HTTPClientT]): + base_url: str + http_client: _HTTPClientT - platform_config: PlatformConfig - default_resolver: Callable[[str], httpx.URL] - def resolve(self, url: str) -> httpx.URL: - return resolve_platform_request_url( - url, - platform_config=self.platform_config, - default_resolver=self.default_resolver, - ) +def _sync_sdk_connection(base_url: str | None, http_client: httpx.Client | None) -> _SDKConnection[httpx.Client]: + if http_client is not None: + return _SDKConnection(base_url=base_url or _base_url_from_config(), http_client=http_client) + if base_url is not None: + return _SDKConnection(base_url=base_url, http_client=ImmutableDefaultHttpxClient()) + endpoint = resolve_platform_endpoint() + return _SDKConnection(base_url=endpoint.connect_base_url, http_client=endpoint.sync_sdk_http_client()) -def attach_platform_request_router(sdk: PlatformSDKT) -> PlatformSDKT: - """Attach the platform request router to a generated SDK instance. +def _async_sdk_connection( + base_url: str | None, + http_client: httpx.AsyncClient | None, +) -> _SDKConnection[httpx.AsyncClient]: + if http_client is not None: + return _SDKConnection(base_url=base_url or _base_url_from_config(), http_client=http_client) + if base_url is not None: + return _SDKConnection(base_url=base_url, http_client=ImmutableDefaultAsyncHttpxClient()) + endpoint = resolve_platform_endpoint() + return _SDKConnection(base_url=endpoint.connect_base_url, http_client=endpoint.async_sdk_http_client()) - Stainless sends every request through ``_prepare_url``. Assigning the hook is - the SDK integration point; the routing policy itself lives in - :class:`PlatformRequestRouter`. - """ - router = PlatformRequestRouter( - platform_config=Configuration.get_platform_config(), - default_resolver=sdk._prepare_url, - ) - setattr(sdk, "_nmp_request_router", router) - sdk._prepare_url = router.resolve - return sdk +def with_options_reusing_http_client(base_sdk: PlatformSDKT, **kwargs: Any) -> PlatformSDKT: + """Return ``base_sdk.with_options(...)`` while reusing its underlying HTTP client.""" + return base_sdk.with_options(**kwargs) -def with_options_preserving_request_router(base_sdk: PlatformSDKT, **kwargs: Any) -> PlatformSDKT: - """Return ``base_sdk.with_options(...)`` while preserving platform request routing.""" - scoped_sdk = cast(PlatformSDKT, base_sdk.with_options(**kwargs)) - router = getattr(base_sdk, "_nmp_request_router", None) - if isinstance(router, PlatformRequestRouter): - setattr(scoped_sdk, "_nmp_request_router", router) - scoped_sdk._prepare_url = router.resolve - return scoped_sdk +class _WorkloadIdentityAuth(httpx.Auth): + def __init__(self, provider: _AccessTokenProvider) -> None: + self._provider = provider -def _sync_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.Client | None, -) -> httpx.Client: - if http_client is not None: - return http_client - if endpoint.transport == "uds": - return endpoint.sync_http_client() - return shared_sync_http_client() + def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]: + request.headers["Authorization"] = f"Bearer {self._provider.get_access_token()}" + yield request + async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]: + request.headers["Authorization"] = f"Bearer {await self._provider.get_access_token_async()}" + yield request -def _async_http_client_for_endpoint( - endpoint: PlatformEndpoint, - http_client: httpx.AsyncClient | None, -) -> httpx.AsyncClient: - if http_client is not None: - return http_client - if _test_http_client is not None: - return _test_http_client - if endpoint.transport == "uds": - return endpoint.async_http_client() - return shared_async_http_client() + +class _CustomAuthNeMoPlatform(NeMoPlatform): + def __init__(self, *, custom_auth: httpx.Auth, **kwargs: Any) -> None: + self._custom_auth = custom_auth + super().__init__(**kwargs) + + @property + def custom_auth(self) -> httpx.Auth | None: + return self._custom_auth + + +class _CustomAuthAsyncNeMoPlatform(AsyncNeMoPlatform): + def __init__(self, *, custom_auth: httpx.Auth, **kwargs: Any) -> None: + self._custom_auth = custom_auth + super().__init__(**kwargs) + + @property + def custom_auth(self) -> httpx.Auth | None: + return self._custom_auth + + +@dataclass(frozen=True) +class _ResolvedSDKInitConfig(Generic[_HTTPClientT]): + base_url: str + workspace: str | None + default_headers: Mapping[str, str] | None + http_client: _HTTPClientT + custom_auth: httpx.Auth | None -def _should_bootstrap_workload_identity( +def _should_bootstrap_workload_identity() -> bool: + return is_workload_identity_token_file_set() + + +def _workload_identity_extra_headers(*, internal: bool) -> dict[str, str]: + return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} + + +def _resolve_sdk_init_config( *, + connection: _SDKConnection[_HTTPClientT], as_service: str | None, + internal: bool, on_behalf_of: str | Principal | None, - http_client: httpx.Client | httpx.AsyncClient | None, - endpoint: PlatformEndpoint, -) -> bool: - return ( - as_service is None - and on_behalf_of is None - and http_client is None - and endpoint.transport != "uds" - and is_workload_identity_token_file_set() +) -> _ResolvedSDKInitConfig[_HTTPClientT]: + if not _should_bootstrap_workload_identity(): + headers = _get_default_headers(as_service, internal, on_behalf_of) + return _ResolvedSDKInitConfig( + base_url=connection.base_url, + workspace=None, + default_headers=headers if headers else None, + http_client=connection.http_client, + custom_auth=None, + ) + + from nemo_platform.client.factory import _resolve_bootstrap + + extra_headers = _workload_identity_extra_headers(internal=internal) + bootstrap = _resolve_bootstrap( + config_path=None, + base_url=connection.base_url, + context_name=None, + access_token=None, + extra_headers=extra_headers if extra_headers else None, + ) + custom_auth = None + if bootstrap.token_provider is not None: + custom_auth = _WorkloadIdentityAuth(bootstrap.token_provider) + return _ResolvedSDKInitConfig( + base_url=bootstrap.base_url, + workspace=bootstrap.workspace, + default_headers=bootstrap.default_headers or None, + http_client=connection.http_client, + custom_auth=custom_auth, ) -def _workload_identity_extra_headers(*, internal: bool) -> dict[str, str]: - return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {} +def _sync_sdk_from_init_config(sdk_config: _ResolvedSDKInitConfig[httpx.Client]) -> NeMoPlatform: + if sdk_config.custom_auth is None: + return NeMoPlatform( + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + ) + return _CustomAuthNeMoPlatform( + custom_auth=sdk_config.custom_auth, + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + ) + + +def _async_sdk_from_init_config(sdk_config: _ResolvedSDKInitConfig[httpx.AsyncClient]) -> AsyncNeMoPlatform: + if sdk_config.custom_auth is None: + return AsyncNeMoPlatform( + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + ) + return _CustomAuthAsyncNeMoPlatform( + custom_auth=sdk_config.custom_auth, + base_url=sdk_config.base_url, + workspace=sdk_config.workspace, + default_headers=sdk_config.default_headers, + http_client=sdk_config.http_client, + ) + + +def _task_on_behalf_of_principal() -> Principal | None: + principal = principal_from_env() + return principal.effective_principal if principal is not None else None + + +def _warn_missing_task_principal(*, as_service: str, async_sdk: bool) -> None: + qualifier = "async task SDK" if async_sdk else "task SDK" + logger.warning( + "NMP_PRINCIPAL not set; %s will authenticate as service:%s without on-behalf-of delegation", + qualifier, + as_service, + ) def _get_default_headers( @@ -246,7 +271,7 @@ def get_platform_sdk( base_url: str | None = None, ) -> NeMoPlatform: """ - Returns an instance of the NeMoPlatform SDK configured with the platform's base URL. + Returns a NeMoPlatform SDK configured from explicit arguments or the resolved platform endpoint. Args: as_service: If provided, use service principal headers (service:{name}). @@ -257,32 +282,20 @@ def get_platform_sdk( Use this for controllers and background tasks that make internal API calls. http_client: Optional sync HTTP client to use for requests. on_behalf_of: Optional principal ID to use for on-behalf-of authorization. - base_url: Optional platform base URL. Defaults to configured platform base URL. + base_url: Optional platform base URL. When omitted with no explicit http_client, + the resolved platform endpoint supplies both the base URL and HTTP client. Returns: Configured NeMoPlatform SDK instance. """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( + connection = _sync_sdk_connection(base_url, http_client) + sdk_config = _resolve_sdk_init_config( + connection=connection, as_service=as_service, + internal=internal, on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - headers = _workload_identity_extra_headers(internal=internal) - sdk = NeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - default_headers=headers if headers else None, - ) - return attach_platform_request_router(sdk) - - headers = _get_default_headers(as_service, internal, on_behalf_of) - sdk = NeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=_sync_http_client_for_endpoint(endpoint, http_client), - default_headers=headers if headers else None, ) - return attach_platform_request_router(sdk) + return _sync_sdk_from_init_config(sdk_config) def get_task_sdk(as_service: str, http_client: httpx.Client | None = None) -> NeMoPlatform: @@ -299,23 +312,14 @@ def get_task_sdk(as_service: str, http_client: httpx.Client | None = None) -> Ne Returns: Configured NeMoPlatform SDK with internal + on-behalf-of headers. """ - if http_client is None and is_workload_identity_token_file_set(): - return get_platform_sdk(internal=True) - - if http_client is None: - http_client = resolve_platform_endpoint().sync_sdk_http_client() - - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; task SDK will authenticate as service:%s without on-behalf-of delegation", - as_service, - ) + on_behalf_of = _task_on_behalf_of_principal() + if on_behalf_of is None and not is_workload_identity_token_file_set(): + _warn_missing_task_principal(as_service=as_service, async_sdk=False) return get_platform_sdk( as_service=as_service, internal=True, http_client=http_client, - on_behalf_of=principal.effective_principal if principal else None, + on_behalf_of=on_behalf_of, ) @@ -333,23 +337,14 @@ def get_async_task_sdk(as_service: str, http_client: Optional[httpx.AsyncClient] Returns: Configured AsyncNeMoPlatform SDK with internal + on-behalf-of headers. """ - if http_client is None and is_workload_identity_token_file_set(): - return get_async_platform_sdk(internal=True) - - if http_client is None: - http_client = resolve_platform_endpoint().async_sdk_http_client() - - principal = principal_from_env() - if principal is None: - logger.warning( - "NMP_PRINCIPAL not set; async task SDK will authenticate as service:%s without on-behalf-of delegation", - as_service, - ) + on_behalf_of = _task_on_behalf_of_principal() + if on_behalf_of is None and not is_workload_identity_token_file_set(): + _warn_missing_task_principal(as_service=as_service, async_sdk=True) return get_async_platform_sdk( as_service=as_service, internal=True, http_client=http_client, - on_behalf_of=principal.effective_principal if principal else None, + on_behalf_of=on_behalf_of, ) @@ -361,7 +356,7 @@ def get_async_platform_sdk( base_url: str | None = None, ) -> AsyncNeMoPlatform: """ - Returns an instance of the AsyncNeMoPlatform SDK configured with the platform's base URL. + Returns an AsyncNeMoPlatform SDK configured from explicit arguments or the resolved platform endpoint. Args: as_service: If provided, use service principal headers (service:{name}). @@ -370,46 +365,21 @@ def get_async_platform_sdk( If None and auth is enabled, propagates the current user's auth context. internal: If True, mark all requests from this SDK as internal requests. Use this for controllers and background tasks that make internal API calls. - http_client: Optional HTTP client to use for requests. Used for test injection - via DependencyProvider. See architecture/docs/http-client-injection.md. + http_client: Optional async HTTP client to use for requests. on_behalf_of: Optional principal ID to use for on-behalf-of authorization. - base_url: Optional platform base URL. Defaults to configured platform base URL. + base_url: Optional platform base URL. When omitted with no explicit http_client, + the resolved platform endpoint supplies both the base URL and HTTP client. Returns: Configured AsyncNeMoPlatform SDK instance. """ - endpoint = resolve_platform_endpoint() - if _should_bootstrap_workload_identity( + connection = _async_sdk_connection(base_url, http_client) + sdk_config = _resolve_sdk_init_config( + connection=connection, as_service=as_service, + internal=internal, on_behalf_of=on_behalf_of, - http_client=http_client, - endpoint=endpoint, - ): - headers = _workload_identity_extra_headers(internal=internal) - if _test_http_client is None: - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - default_headers=headers if headers else None, - ) - else: - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=_async_http_client_for_endpoint(endpoint, http_client), - default_headers=headers if headers else None, - ) - return attach_platform_request_router(sdk) - - headers = _get_default_headers(as_service, internal, on_behalf_of) - - # Use explicitly provided http_client (from DependencyProvider) or fall back to - # module-level _test_http_client for backward compatibility with direct callers. - effective_client = _async_http_client_for_endpoint(endpoint, http_client) - - sdk = AsyncNeMoPlatform( - base_url=base_url or endpoint.connect_base_url, - http_client=effective_client, - default_headers=headers if headers else None, ) - return attach_platform_request_router(sdk) + return _async_sdk_from_init_config(sdk_config) def get_request_scoped_sdk( @@ -417,7 +387,7 @@ def get_request_scoped_sdk( ) -> AsyncNeMoPlatform: """Create a request-scoped SDK with current auth and observability headers. - Takes a base SDK (with shared HTTP client) and returns a new SDK instance + Takes a base SDK and returns a new SDK instance that reuses the same HTTP client with the current request's auth headers applied via .with_options(). This is lightweight - the underlying HTTP client is reused. @@ -440,7 +410,7 @@ def get_request_scoped_sdk( # If we have headers to add, create a new SDK with them # This reuses the underlying HTTP client (lightweight operation) if headers: - return with_options_preserving_request_router(base_sdk, set_default_headers=headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=headers) return base_sdk @@ -496,7 +466,7 @@ def get_sdk_on_behalf_of( merged_headers = {**headers, "X-NMP-Principal-On-Behalf-Of": on_behalf_of} merged_headers.pop("X-NMP-Principal-On-Behalf-Of-Groups", None) merged_headers.pop("X-NMP-Principal-On-Behalf-Of-Email", None) - return with_options_preserving_request_router(base_sdk, set_default_headers=merged_headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=merged_headers) def get_entity_parts(name: str, default_workspace: str | None = None) -> tuple[str, str]: @@ -518,8 +488,7 @@ def get_entity_parts(name: str, default_workspace: str | None = None) -> tuple[s class PlatformSDKProvider: """Rich :class:`~nemo_platform_plugin.sdk_provider.SDKProvider` that uses - platform internals (shared HTTP clients, URL routing, OTEL headers, auth - context vars). + platform internals (SDK-owned HTTP clients, OTEL headers, auth context vars). Registered as a ``nemo.sdk_provider`` entry-point so it is discovered automatically when ``nmp-common`` is installed. diff --git a/packages/nmp_common/src/nmp/common/service/__init__.py b/packages/nmp_common/src/nmp/common/service/__init__.py index 084a7a65e8..cedf12467c 100644 --- a/packages/nmp_common/src/nmp/common/service/__init__.py +++ b/packages/nmp_common/src/nmp/common/service/__init__.py @@ -12,11 +12,13 @@ ) from nmp.common.service.deptree import CircularDependencyError, resolve_service_loading_order from nmp.common.service.headers import build_downstream_service_headers +from nmp.common.service.sdk_factory import ServiceSDKFactory __all__ = [ "CircularDependencyError", "DependencyProvider", "Service", + "ServiceSDKFactory", "RouterConfig", "build_downstream_service_headers", "get_entity_client", diff --git a/packages/nmp_common/src/nmp/common/service/base.py b/packages/nmp_common/src/nmp/common/service/base.py index 8a1c076fc3..d1cf10e0b9 100644 --- a/packages/nmp_common/src/nmp/common/service/base.py +++ b/packages/nmp_common/src/nmp/common/service/base.py @@ -10,21 +10,45 @@ from abc import ABC, abstractmethod from contextlib import asynccontextmanager from dataclasses import dataclass -from typing import ClassVar, Dict, Generic, List, Optional, Self, Type, TypeVar, cast, get_args, get_origin +from typing import Any, ClassVar, Dict, Generic, List, Optional, Self, Type, TypeVar, get_args, get_origin import httpx from fastapi import APIRouter, FastAPI from fastapi.openapi.utils import get_openapi -from nemo_platform import AsyncNeMoPlatform, DefaultAsyncHttpxClient +from fastapi.routing import APIRoute +from nemo_platform import AsyncNeMoPlatform from nmp.common.api.utils import register_query_param_schemas from nmp.common.config import Configuration, PlatformConfig, ServiceConfig +from nmp.common.config import get_platform_config as load_platform_config from nmp.common.controller import Controller from nmp.common.entities.client import EntityClient -from nmp.common.platform_endpoint import resolve_service_endpoint +from nmp.common.platform_endpoint import PlatformEndpoint, resolve_platform_endpoint, resolve_service_endpoint +from nmp.common.service.sdk_factory import ServiceSDKFactory logger = logging.getLogger(__name__) +class _ServiceFastAPI(FastAPI): + """FastAPI app with NeMo Platform OpenAPI post-processing.""" + + nmp_openapi_summary: str | None = None + nmp_openapi_description: str | None = None + + def openapi(self) -> dict[str, Any]: + if self.openapi_schema: + return self.openapi_schema + openapi_schema = get_openapi( + title=self.title, + version=self.version, + summary=self.nmp_openapi_summary or self.summary, + description=self.description if self.nmp_openapi_description is None else self.nmp_openapi_description, + routes=self.routes, + tags=self.openapi_tags, + ) + self.openapi_schema = register_query_param_schemas(openapi_schema) + return self.openapi_schema + + @dataclass class RouterConfig: """Configuration for a router including its OpenAPI tag metadata.""" @@ -38,7 +62,7 @@ class RouterConfig: TConfig = TypeVar("TConfig", bound=ServiceConfig) -def _get_config_class_from_generic(cls: type) -> Type[ServiceConfig] | None: +def _get_config_class_from_generic(cls: type) -> Type[TConfig] | None: """Extract the config class from Service[TConfig] generic parameter. Args: @@ -62,27 +86,60 @@ class DependencyProvider: Provides lazy initialization, FastAPI dependency wiring, and cleanup. - The `_http_client` field supports test injection - when set, it's passed to - `get_async_platform_sdk()` to route requests through ASGI transport in tests. - See architecture/docs/http-client-injection.md for details. + The PlatformEndpoint owns service routing for the current PlatformConfig. + Tests and embedded platform assembly can still inject an explicit endpoint + or HTTP client before the SDK factory is created. """ - def __init__(self) -> None: - self._http_client: Optional[httpx.AsyncClient] = None + def __init__(self, service_name: str = "platform") -> None: + self._configured_http_client: Optional[httpx.AsyncClient] = None + self._platform_endpoint: Optional[PlatformEndpoint] = None + self._service_sdk_factory: Optional[ServiceSDKFactory] = None self._sdk_client: Optional[AsyncNeMoPlatform] = None self._platform_config: Optional[PlatformConfig] = None - self._service_name: str = "platform" + self._service_name = service_name + + def configure_service_name(self, service_name: str) -> None: + """Set the service name used for downstream service-principal headers.""" + self._service_name = service_name + + def _get_service_sdk_factory(self) -> ServiceSDKFactory: + if self._service_sdk_factory is None: + self._service_sdk_factory = ServiceSDKFactory( + self.get_platform_endpoint(), + async_http_client=self._configured_http_client, + ) + return self._service_sdk_factory def get_http_client(self) -> httpx.AsyncClient: """Return the httpx.AsyncClient for this provider, creating it lazily. - 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. + The PlatformEndpoint carries resolved routing state. No PlatformConfig + or factory is passed into the HTTP-client constructor. """ - if self._http_client is None: - self._http_client = DefaultAsyncHttpxClient() - return self._http_client + return self._get_service_sdk_factory().get_async_http_client() + + def get_platform_endpoint(self) -> PlatformEndpoint: + """Return the resolved PlatformEndpoint for this provider.""" + if self._platform_endpoint is None: + self._platform_endpoint = resolve_platform_endpoint(self.get_platform_config()) + return self._platform_endpoint + + def configure_platform_endpoint(self, platform_endpoint: PlatformEndpoint) -> None: + """Use an externally resolved PlatformEndpoint before SDK/client creation.""" + if self._service_sdk_factory is not None or self._sdk_client is not None: + raise RuntimeError("Cannot configure DependencyProvider PlatformEndpoint after SDK/client creation") + self._platform_endpoint = platform_endpoint + + def configure_http_client(self, http_client: httpx.AsyncClient) -> None: + """Use an externally managed HTTP client before SDK creation. + + This is used by platform assembly and tests to route all service-owned + SDK calls through the same transport and connection pool. + """ + if self._service_sdk_factory is not None or self._sdk_client is not None: + raise RuntimeError("Cannot configure DependencyProvider HTTP client after SDK/client creation") + self._configured_http_client = http_client def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: """Return the async platform SDK client. @@ -97,16 +154,17 @@ def get_sdk_client(self, as_service: str | None = None) -> AsyncNeMoPlatform: Returns: SDK client - cached instance if as_service is None, new instance otherwise. """ - from nmp.common.sdk_factory import get_async_platform_sdk - # 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 self._get_service_sdk_factory().create_async_platform_sdk( + as_service=as_service, + internal=True, + ) # 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) + self._sdk_client = self._get_service_sdk_factory().create_async_platform_sdk() return self._sdk_client def get_entity_client(self, as_service: str | None = None) -> Optional[EntityClient]: @@ -146,19 +204,20 @@ def _get_entity_sdk_on_behalf_of(self) -> AsyncNeMoPlatform: Uses the cached base SDK and applies per-request headers via .with_options() (lightweight — reuses the HTTP connection pool). """ - from nmp.common.sdk_factory import with_options_preserving_request_router + from nmp.common.sdk_factory import with_options_reusing_http_client from nmp.common.service.headers import build_downstream_service_headers base_sdk = self.get_sdk_client() headers = build_downstream_service_headers(self._service_name) - return with_options_preserving_request_router(base_sdk, set_default_headers=headers) + return with_options_reusing_http_client(base_sdk, set_default_headers=headers) def get_platform_config(self) -> PlatformConfig: """Return the PlatformConfig (lazily initialized).""" if self._platform_config is None: - self._platform_config = Configuration.get_platform_config() - return self._platform_config + self._platform_config = load_platform_config() + platform_config = self._platform_config + return platform_config def get_request_scoped_sdk(self) -> AsyncNeMoPlatform: """Return a request-scoped SDK with current auth and OTEL headers. @@ -192,9 +251,10 @@ async def close(self) -> None: 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() - self._http_client = None + if self._service_sdk_factory is not None: + await self._service_sdk_factory.aclose() + self._service_sdk_factory = None + self._configured_http_client = None if self._sdk_client is not None: await self._sdk_client.close() self._sdk_client = None @@ -256,7 +316,10 @@ def __init__( self.module_name = module_name self._app: Optional[FastAPI] = None self._startup_background_tasks: list[asyncio.Task] = [] - self._dependency_provider = dependency_provider if dependency_provider is not None else DependencyProvider() + self._dependency_provider = ( + dependency_provider if dependency_provider is not None else DependencyProvider(service_name=name) + ) + self._dependency_provider.configure_service_name(name) if dependencies is not None: self._dependencies = list(dependencies) else: @@ -264,9 +327,7 @@ def __init__( # Extract config class from generic type parameter and load config config_class = _get_config_class_from_generic(type(self)) - self._service_config = ( - cast(TConfig | None, Configuration.get_service_config(config_class)) if config_class else None - ) + self._service_config = Configuration.get_service_config(config_class) if config_class else None @property def dependency_provider(self) -> DependencyProvider: @@ -446,13 +507,14 @@ async def lifespan(app: FastAPI): router_configs = self.get_routers() openapi_tags: List[Dict[str, str]] = [{"name": rc.tag, "description": rc.description} for rc in router_configs] - app = FastAPI( + app = _ServiceFastAPI( title=self.title, description=self.description, version=self.version, openapi_tags=openapi_tags, lifespan=lifespan, ) + app.nmp_openapi_summary = f"This is the OpenAPI Schema for the {self.title}." # Store reference to app for use in on_startup self._app = app @@ -473,35 +535,12 @@ async def lifespan(app: FastAPI): # Include service-specific routers, tagging any routes that have no tags yet for rc in router_configs: for route in rc.router.routes: - if hasattr(route, "tags") and not route.tags: + if isinstance(route, APIRoute) and not route.tags: route.tags = [rc.tag] app.include_router(rc.router, prefix=rc.prefix) - # Setup custom OpenAPI schema - self._setup_custom_openapi(app, openapi_tags) - return app - def _setup_custom_openapi(self, app: FastAPI, openapi_tags: List[Dict[str, str]]) -> None: - """Configure custom OpenAPI schema generation.""" - - def custom_openapi(): - if app.openapi_schema: - return app.openapi_schema - openapi_schema = get_openapi( - title=self.title, - version=self.version, - summary=f"This is the OpenAPI Schema for the {self.title}.", - description="", - routes=app.routes, - tags=openapi_tags, - ) - openapi_schema = register_query_param_schemas(openapi_schema) - app.openapi_schema = openapi_schema - return app.openapi_schema - - app.openapi = custom_openapi # type: ignore[method-assign] - # ========================================================================= # Startup and readiness # ========================================================================= diff --git a/packages/nmp_common/src/nmp/common/service/dependencies.py b/packages/nmp_common/src/nmp/common/service/dependencies.py index 8eade0d015..a45892ea7a 100644 --- a/packages/nmp_common/src/nmp/common/service/dependencies.py +++ b/packages/nmp_common/src/nmp/common/service/dependencies.py @@ -9,7 +9,8 @@ from __future__ import annotations -from typing import Callable, TypeVar +from collections.abc import Callable, Mapping +from typing import TypeVar from fastapi import Request from nemo_platform_plugin.dependencies import get_entity_client as get_entity_client @@ -35,12 +36,12 @@ def get_service_config_factory(config_class: type[T]) -> Callable[[Request], T]: """ def _get_config(request: Request) -> T: - registry: dict[type[ServiceConfig], ServiceConfig] = getattr(request.app.state, "service_configs", {}) + registry: Mapping[type[T], T] = getattr(request.app.state, "service_configs", {}) if config_class not in registry: raise RuntimeError( f"Service config {config_class.__name__} not registered. " "Ensure the service is loaded and its config is added to app.state.service_configs." ) - return registry[config_class] # type: ignore[return-value] + return registry[config_class] return _get_config diff --git a/packages/nmp_common/src/nmp/common/service/sdk_factory.py b/packages/nmp_common/src/nmp/common/service/sdk_factory.py new file mode 100644 index 0000000000..7629ac65f7 --- /dev/null +++ b/packages/nmp_common/src/nmp/common/service/sdk_factory.py @@ -0,0 +1,88 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Service-scoped SDK factory. + +Generic SDK creation should not resolve platform service topology. Services +receive a resolved ``PlatformEndpoint`` dependency and this layer passes the +endpoint's already-configured immutable clients into SDK constructors. +""" + +from __future__ import annotations + +import httpx +from httpx._types import TimeoutTypes +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform +from nmp.common.auth import Principal +from nmp.common.platform_endpoint import PlatformEndpoint + + +class ServiceSDKFactory: + """Create SDKs for a service using a resolved platform endpoint.""" + + def __init__( + self, + platform_endpoint: PlatformEndpoint, + *, + sync_http_client: httpx.Client | None = None, + async_http_client: httpx.AsyncClient | None = None, + ) -> None: + self._platform_endpoint = platform_endpoint + self._sync_http_client = sync_http_client + self._async_http_client = async_http_client + + @property + def base_url(self) -> str: + return self._platform_endpoint.connect_base_url + + def get_sync_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.Client: + if self._sync_http_client is None: + self._sync_http_client = self._platform_endpoint.sync_sdk_http_client(timeout=timeout) + return self._sync_http_client + + def get_async_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.AsyncClient: + if self._async_http_client is None: + self._async_http_client = self._platform_endpoint.async_sdk_http_client(timeout=timeout) + return self._async_http_client + + def create_platform_sdk( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + ) -> NeMoPlatform: + from nmp.common.sdk_factory import get_platform_sdk + + return get_platform_sdk( + as_service=as_service, + internal=internal, + http_client=self.get_sync_http_client(), + on_behalf_of=on_behalf_of, + base_url=self.base_url, + ) + + def create_async_platform_sdk( + self, + *, + as_service: str | None = None, + internal: bool = False, + on_behalf_of: str | Principal | None = None, + ) -> AsyncNeMoPlatform: + from nmp.common.sdk_factory import get_async_platform_sdk + + return get_async_platform_sdk( + as_service=as_service, + internal=internal, + http_client=self.get_async_http_client(), + on_behalf_of=on_behalf_of, + base_url=self.base_url, + ) + + async def aclose(self) -> None: + if self._async_http_client is not None: + await self._async_http_client.aclose() + self._async_http_client = None + if self._sync_http_client is not None: + self._sync_http_client.close() + self._sync_http_client = None diff --git a/packages/nmp_common/tests/api/test_query_param_schemas.py b/packages/nmp_common/tests/api/test_query_param_schemas.py index 5b9236e8db..b789424519 100644 --- a/packages/nmp_common/tests/api/test_query_param_schemas.py +++ b/packages/nmp_common/tests/api/test_query_param_schemas.py @@ -4,23 +4,22 @@ """Tests for register_query_param_schemas / clear_query_param_schemas. These schemas are attached to FastAPI endpoints via ``openapi_extra`` and are -not reachable through Pydantic's response-model walk. The runtime -``custom_openapi`` hook has to call ``register_query_param_schemas`` explicitly -or the live ``/openapi.json`` will contain dangling ``$ref``s to the filter -classes. +not reachable through Pydantic's response-model walk. The runtime OpenAPI path +has to call ``register_query_param_schemas`` or the live ``/openapi.json`` will +contain dangling ``$ref``s to the filter classes. """ from typing import Optional import pytest -from fastapi import FastAPI, Query, Request -from fastapi.openapi.utils import get_openapi +from fastapi import APIRouter, Query, Request from fastapi.testclient import TestClient from nmp.common.api.utils import ( clear_query_param_schemas, generate_openapi_extra_params, register_query_param_schemas, ) +from nmp.common.service import RouterConfig, Service from pydantic import BaseModel @@ -67,31 +66,26 @@ def test_clear_resets_registry_between_services(): assert "_DummyFilter" not in spec["components"]["schemas"] -def test_custom_openapi_hook_resolves_filter_ref(): - """End-to-end: a FastAPI app that wires ``register_query_param_schemas`` - into its ``custom_openapi`` hook emits a spec where the filter $ref - resolves — which is exactly the regression the runtime was missing. - """ - app = FastAPI() - - @app.get( - "/items", - openapi_extra=generate_openapi_extra_params(filter_schema=_DummyFilter), - ) - async def list_items(request: Request, page: int = Query(default=1)): - return {"data": []} - - def custom_openapi(): - if app.openapi_schema: - return app.openapi_schema - spec = get_openapi(title="t", version="0", routes=app.routes) - spec = register_query_param_schemas(spec) - app.openapi_schema = spec - return spec - - app.openapi = custom_openapi # type: ignore[method-assign] - - spec = TestClient(app).get("/openapi.json").json() +def test_service_openapi_resolves_filter_ref(): + """End-to-end: service OpenAPI generation resolves generated filter refs.""" + + class _QueryParamService(Service): + def __init__(self): + super().__init__(name="query-param-test", module_name="nmp.test") + + def get_routers(self) -> list[RouterConfig]: + router = APIRouter() + + @router.get( + "/items", + openapi_extra=generate_openapi_extra_params(filter_schema=_DummyFilter), + ) + async def list_items(request: Request, page: int = Query(default=1)): + return {"data": []} + + return [RouterConfig(router, tag="Items", description="Item endpoints")] + + spec = TestClient(_QueryParamService().create_app()).get("/openapi.json").json() assert "_DummyFilter" in spec["components"]["schemas"] param = next(p for p in spec["paths"]["/items"]["get"]["parameters"] if p["name"] == "filter") diff --git a/packages/nmp_common/tests/nmp_common/test_common_service.py b/packages/nmp_common/tests/nmp_common/test_common_service.py index 45d1bcce6f..194c4a7fe7 100644 --- a/packages/nmp_common/tests/nmp_common/test_common_service.py +++ b/packages/nmp_common/tests/nmp_common/test_common_service.py @@ -113,6 +113,7 @@ def test_service_create_app(self): assert app is not None assert app.title == service.title assert app.version == service.version + assert app.openapi()["info"]["description"] == service.description def test_service_app_property_caches(self): """Test app property returns cached instance.""" @@ -181,7 +182,7 @@ def handler(request: httpx.Request) -> httpx.Response: provider = DependencyProvider() provider._platform_config = PlatformConfig(base_url="http://platform.local") async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: - provider._http_client = client + provider.configure_http_client(client) service = MockService(dependency_provider=provider) ready = await service.wait_for_service_ready("entities", timeout=1.0, poll_interval=0) @@ -201,7 +202,7 @@ def handler(request: httpx.Request) -> httpx.Response: provider = DependencyProvider() provider._platform_config = PlatformConfig(base_url="http://platform.local") async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: - provider._http_client = client + provider.configure_http_client(client) service = MockService(dependency_provider=provider) ready = await service.wait_for_service_ready("models", timeout=1.0, poll_interval=0) @@ -217,7 +218,7 @@ def test_init(self): """Test DependencyProvider initialization.""" provider = DependencyProvider() assert provider._sdk_client is None - assert provider._http_client is None + assert provider._configured_http_client is None @pytest.mark.asyncio async def test_close_without_clients(self): @@ -235,6 +236,20 @@ def test_service_has_provider(self): assert service.dependency_provider is not None assert isinstance(service.dependency_provider, DependencyProvider) + def test_service_configures_provider_service_name_for_entity_client_headers(self): + """Test Service configures the provider before entity SDK creation.""" + provider = DependencyProvider() + service = MockService(dependency_provider=provider) + + with patch.object(service.dependency_provider, "get_sdk_client") as mock_sdk: + mock_base_sdk = mock_sdk.return_value + + service.dependency_provider._get_entity_sdk_on_behalf_of() + + call_kwargs = mock_base_sdk.with_options.call_args + headers = call_kwargs.kwargs.get("set_default_headers") or call_kwargs[1].get("set_default_headers") + assert headers["X-NMP-Principal-Id"] == "service:test-service" + class LifecycleService(MockService): def __init__(self): diff --git a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py index efffb0e41a..7ed19536dc 100644 --- a/packages/nmp_common/tests/nmp_common/test_dependency_provider.py +++ b/packages/nmp_common/tests/nmp_common/test_dependency_provider.py @@ -3,43 +3,105 @@ from __future__ import annotations -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call, patch import pytest from fastapi import FastAPI +from nmp.common.config import PlatformConfig from nmp.common.service import DependencyProvider -from nmp.common.service.dependencies import get_entity_client, get_platform_config, get_sdk_client +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_delegates_to_service_sdk_factory() -> None: provider = DependencyProvider() + platform_config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + provider._platform_config = platform_config + platform_endpoint = MagicMock(name="platform_endpoint") client = MagicMock() + factory = MagicMock() + factory.get_async_http_client.return_value = client - with patch("nmp.common.service.base.DefaultAsyncHttpxClient", return_value=client) as factory: + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint) as resolve_endpoint, + patch("nmp.common.service.base.ServiceSDKFactory", return_value=factory) as factory_cls, + ): first = provider.get_http_client() second = provider.get_http_client() assert first is client assert second is client - factory.assert_called_once_with() + resolve_endpoint.assert_called_once_with(platform_config) + factory_cls.assert_called_once_with(platform_endpoint, async_http_client=None) + assert factory.get_async_http_client.call_count == 2 + + +def test_configure_http_client_is_used_for_cached_sdk_client() -> None: + provider = DependencyProvider() + http_client = MagicMock(name="http_client") + sdk = MagicMock(name="sdk") + platform_config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + provider._platform_config = platform_config + platform_endpoint = MagicMock(name="platform_endpoint") + factory = MagicMock() + factory.create_async_platform_sdk.return_value = sdk + + provider.configure_http_client(http_client) + + with ( + patch("nmp.common.service.base.resolve_platform_endpoint", return_value=platform_endpoint), + patch("nmp.common.service.base.ServiceSDKFactory", return_value=factory) as factory_cls, + ): + assert provider.get_sdk_client() is sdk + assert provider.get_sdk_client() is sdk + + factory_cls.assert_called_once_with(platform_endpoint, async_http_client=http_client) + factory.create_async_platform_sdk.assert_called_once_with() + + +def test_configure_platform_endpoint_is_used_for_service_sdk_factory() -> None: + provider = DependencyProvider() + platform_endpoint = MagicMock(name="platform_endpoint") + factory = MagicMock() + factory.get_async_http_client.return_value = MagicMock(name="http_client") + + provider.configure_platform_endpoint(platform_endpoint) + + with patch("nmp.common.service.base.ServiceSDKFactory", return_value=factory) as factory_cls: + provider.get_http_client() + + factory_cls.assert_called_once_with(platform_endpoint, async_http_client=None) + + +def test_configure_http_client_rejects_changes_after_sdk_created() -> None: + provider = DependencyProvider() + sdk = MagicMock(name="sdk") + factory = MagicMock() + factory.create_async_platform_sdk.return_value = sdk + + with patch("nmp.common.service.base.ServiceSDKFactory", return_value=factory): + assert provider.get_sdk_client() is sdk + + with pytest.raises(RuntimeError, match="Cannot configure DependencyProvider HTTP client after SDK/client creation"): + provider.configure_http_client(MagicMock(name="late_http_client")) def test_get_sdk_client_caches_request_sdk_and_creates_fresh_service_sdk() -> None: provider = DependencyProvider() request_sdk = MagicMock(name="request_sdk") service_sdk = MagicMock(name="service_sdk") + factory = MagicMock() + factory.create_async_platform_sdk.side_effect = [request_sdk, service_sdk] - with patch("nmp.common.sdk_factory.get_async_platform_sdk", side_effect=[request_sdk, service_sdk]) as factory: + with patch("nmp.common.service.base.ServiceSDKFactory", return_value=factory): assert provider.get_sdk_client() is request_sdk assert provider.get_sdk_client() is request_sdk assert provider.get_sdk_client(as_service="jobs") is service_sdk - assert factory.call_args_list[0].kwargs == {"http_client": None} - assert factory.call_args_list[1].kwargs == { - "as_service": "jobs", - "internal": True, - "http_client": None, - } + assert factory.create_async_platform_sdk.call_args_list == [call(), call(as_service="jobs", internal=True)] def test_setup_dependencies_registers_fastapi_overrides() -> None: @@ -58,18 +120,20 @@ def test_setup_dependencies_registers_fastapi_overrides() -> None: @pytest.mark.asyncio async def test_close_closes_managed_clients_and_clears_references() -> None: provider = DependencyProvider() - http_client = MagicMock() - http_client.aclose = AsyncMock() + service_sdk_factory = MagicMock() + service_sdk_factory.aclose = AsyncMock() sdk = MagicMock() sdk.close = AsyncMock() - provider._http_client = http_client + provider._service_sdk_factory = service_sdk_factory + provider._configured_http_client = MagicMock() provider._sdk_client = sdk await provider.close() - http_client.aclose.assert_awaited_once_with() + service_sdk_factory.aclose.assert_awaited_once_with() sdk.close.assert_awaited_once_with() - assert provider._http_client is None + assert provider._service_sdk_factory is None + assert provider._configured_http_client is None assert provider._sdk_client is None diff --git a/packages/nmp_common/tests/sdk_factory/test_sdk.py b/packages/nmp_common/tests/sdk_factory/test_sdk.py index e3df60674d..4a32ddf7c8 100644 --- a/packages/nmp_common/tests/sdk_factory/test_sdk.py +++ b/packages/nmp_common/tests/sdk_factory/test_sdk.py @@ -7,12 +7,11 @@ import httpx import pytest +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform from nemo_platform.auth.helpers import NMPOIDCConfig from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR -from nmp.common.config import Configuration, PlatformConfig -from nmp.common.http_clients import shared_async_http_client, shared_sync_http_client +from nmp.common.config import Configuration from nmp.common.sdk_factory import ( - PlatformRequestRouter, get_async_platform_sdk, get_async_task_sdk, get_entity_parts, @@ -20,7 +19,6 @@ get_request_scoped_sdk, get_sdk_on_behalf_of, get_task_sdk, - resolve_platform_request_url, ) @@ -35,24 +33,25 @@ def _workload_oidc_config() -> NMPOIDCConfig: ) -@pytest.fixture(autouse=True) -def _clear_sdk_factory_test_client(): - """Clear SDK factory state before each test so config-based SDK behavior is asserted. +def _apply_sync_auth(sdk: NeMoPlatform, request: httpx.Request) -> None: + assert sdk.custom_auth is not None + auth_flow = sdk.custom_auth.sync_auth_flow(request) + assert next(auth_flow) is request + + +async def _apply_async_auth(sdk: AsyncNeMoPlatform, request: httpx.Request) -> None: + assert sdk.custom_auth is not None + auth_flow = sdk.custom_auth.async_auth_flow(request) + assert await anext(auth_flow) is request - When _test_http_client is set (e.g. by another test's create_test_client), the SDK - is created with base_url='http://testserver' and no request router, which breaks tests - that assert on base_url or service routing. Clearing it keeps tests order-independent - and ensures sdk_factory tests always exercise the config path. - """ - import nmp.common.sdk_factory as sdk_factory_module - old = sdk_factory_module._test_http_client - sdk_factory_module._test_http_client = None +@pytest.fixture(autouse=True) +def _clear_sdk_factory_config(): + """Clear SDK factory config state before each test.""" Configuration.clear_cache() try: yield finally: - sdk_factory_module._test_http_client = old Configuration.clear_cache() @@ -68,7 +67,7 @@ def test_get_platform_sdk(): def test_get_platform_sdk_keeps_platform_base_url_for_local_services(monkeypatch: pytest.MonkeyPatch): - """The SDK base URL remains the platform entrypoint; per-service routing handles local APIs.""" + """The SDK base URL remains the platform entrypoint.""" monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICES", "auth") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") @@ -117,8 +116,8 @@ def capture_request(request: httpx.Request) -> httpx.Response: assert str(captured_requests[0].url) == "http://nemo-platform-api:8080/apis/jobs/v2/workspaces/default/jobs" -def test_get_platform_sdk_routes_local_service_path_to_process_listener(monkeypatch: pytest.MonkeyPatch): - """Requests for APIs hosted in this process bypass the platform entrypoint.""" +def test_get_platform_sdk_does_not_route_local_service_paths(monkeypatch: pytest.MonkeyPatch): + """SDK instances build URLs only from their configured base URL.""" monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICES", "auth") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") @@ -128,21 +127,12 @@ def test_get_platform_sdk_routes_local_service_path_to_process_listener(monkeypa sdk = get_platform_sdk() prepared = sdk._prepare_url("https://nemo-gateway:8080/apis/auth/v2/authz/allow") - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" + assert prepared.scheme == "https" + assert prepared.host == "nemo-gateway" assert prepared.port == 8080 assert prepared.path == "/apis/auth/v2/authz/allow" -def test_get_platform_sdk_uses_uds_endpoint_from_base_url(): - config = PlatformConfig(base_url="unix:///tmp/nemo-platform.sock") # type: ignore[abstract] - - with patch("nmp.common.sdk_factory.Configuration.get_platform_config", return_value=config): - sdk = get_platform_sdk() - - assert sdk.base_url == "http://nemo-platform.local" - - def test_get_platform_sdk_with_service_principal(): """Test get_platform_sdk with as_service parameter.""" sdk = get_platform_sdk(as_service="my-service") @@ -195,7 +185,7 @@ def token_exchange_grant(**kwargs): sdk = get_platform_sdk() try: request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") - sdk._client._event_hooks["request"][0](request) + _apply_sync_auth(sdk, request) finally: sdk.close() @@ -204,47 +194,97 @@ def token_exchange_grant(**kwargs): assert exchange_requests[0]["subject_token"] == "subject-token-from-file" -def test_get_async_platform_sdk(): - """Test get_async_platform_sdk basic functionality (config path: SDK base_url matches platform config).""" - sdk = get_async_platform_sdk() - - assert sdk is not None - assert hasattr(sdk, "base_url") - # Normalize to str: SDK may expose URL object, config may be str; both environments - expected = Configuration.get_platform_config().base_url - assert str(sdk.base_url).rstrip("/") == str(expected).rstrip("/") +def test_get_platform_sdk_uses_workload_identity_with_service_headers_when_token_file_configured( + monkeypatch: pytest.MonkeyPatch, tmp_path +): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "service-access-token", "expires_in": 300} -def test_get_async_platform_sdk_uses_uds_endpoint_from_base_url(): - config = PlatformConfig(base_url="unix:///tmp/nemo-platform.sock") # type: ignore[abstract] + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config()) + monkeypatch.setattr("nemo_platform.auth.workload_exchange.token_exchange_grant", token_exchange_grant) - with patch("nmp.common.sdk_factory.Configuration.get_platform_config", return_value=config): - sdk = get_async_platform_sdk() + sdk = get_platform_sdk(as_service="jobs", internal=True) + try: + request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") + _apply_sync_auth(sdk, request) + finally: + sdk.close() - assert str(sdk.base_url).rstrip("/") == "http://nemo-platform.local" + assert sdk.default_headers["X-NMP-Internal"] == "true" + assert request.headers["Authorization"] == "Bearer service-access-token" + assert "X-NMP-Principal-Id" not in request.headers + assert "X-NMP-Principal-On-Behalf-Of" not in request.headers + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" -@pytest.mark.asyncio -async def test_get_async_platform_sdk_workload_identity_reuses_test_http_client( +def test_get_platform_sdk_uses_workload_identity_with_explicit_sync_http_client( monkeypatch: pytest.MonkeyPatch, tmp_path ): - import nmp.common.sdk_factory as sdk_factory_module - subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + captured_requests: list[httpx.Request] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "injected-client-access-token", "expires_in": 300} + + def capture_request(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "data": [], + "pagination": { + "current_page_size": 0, + "page": 1, + "page_size": 0, + "total_pages": 1, + "total_results": 0, + }, + }, + ) + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config()) + monkeypatch.setattr("nemo_platform.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + + with httpx.Client(transport=httpx.MockTransport(capture_request)) as http_client: + sdk = get_platform_sdk(http_client=http_client) + sdk.jobs.list(workspace="default") - transport = httpx.MockTransport(lambda _request: httpx.Response(200, json={})) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as http_client: - sdk_factory_module._test_http_client = http_client - try: - sdk = get_async_platform_sdk() + assert sdk._client is http_client - assert sdk._client is http_client - assert str(sdk.base_url).rstrip("/") == "http://nmp.example.test" - finally: - sdk_factory_module._test_http_client = None + assert captured_requests[0].headers["Authorization"] == "Bearer injected-client-access-token" + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + + +def test_get_async_platform_sdk(): + """Test get_async_platform_sdk basic functionality (config path: SDK base_url matches platform config).""" + sdk = get_async_platform_sdk() + + assert sdk is not None + assert hasattr(sdk, "base_url") + # Normalize to str: SDK may expose URL object, config may be str; both environments + expected = Configuration.get_platform_config().base_url + assert str(sdk.base_url).rstrip("/") == str(expected).rstrip("/") def test_get_async_platform_sdk_with_service_principal(): @@ -276,6 +316,90 @@ def test_get_async_platform_sdk_internal_flag(): assert sdk.default_headers["X-NMP-Internal"] == "true" +@pytest.mark.asyncio +async def test_get_async_platform_sdk_uses_workload_identity_with_service_headers_when_token_file_configured( + monkeypatch: pytest.MonkeyPatch, tmp_path +): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "async-service-access-token", "expires_in": 300} + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config()) + monkeypatch.setattr("nemo_platform.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + + sdk = get_async_platform_sdk(as_service="jobs", internal=True) + try: + request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") + await _apply_async_auth(sdk, request) + finally: + await sdk.close() + + assert sdk.default_headers["X-NMP-Internal"] == "true" + assert request.headers["Authorization"] == "Bearer async-service-access-token" + assert "X-NMP-Principal-Id" not in request.headers + assert "X-NMP-Principal-On-Behalf-Of" not in request.headers + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + + +@pytest.mark.asyncio +async def test_get_async_platform_sdk_uses_workload_identity_with_explicit_async_http_client( + monkeypatch: pytest.MonkeyPatch, tmp_path +): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") + config_file = tmp_path / "config.yaml" + config_file.write_text("{}\n", encoding="utf-8") + exchange_requests: list[dict] = [] + captured_requests: list[httpx.Request] = [] + + def token_exchange_grant(**kwargs): + exchange_requests.append(kwargs) + return {"access_token": "async-injected-client-access-token", "expires_in": 300} + + def capture_request(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "data": [], + "pagination": { + "current_page_size": 0, + "page": 1, + "page_size": 0, + "total_pages": 1, + "total_results": 0, + }, + }, + ) + + monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) + monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) + monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config()) + monkeypatch.setattr("nemo_platform.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + + async with httpx.AsyncClient(transport=httpx.MockTransport(capture_request)) as http_client: + sdk = get_async_platform_sdk(http_client=http_client) + await sdk.jobs.list(workspace="default") + + assert sdk._client is http_client + + assert captured_requests[0].headers["Authorization"] == "Bearer async-injected-client-access-token" + assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + + def test_on_behalf_of_without_service_principal(): """Test that on_behalf_of works without as_service (propagates user context).""" # When auth is enabled but no context is set, on_behalf_of should still be added @@ -318,60 +442,62 @@ def test_get_task_sdk_without_principal(monkeypatch: pytest.MonkeyPatch): assert "X-NMP-Principal-On-Behalf-Of" not in sdk.default_headers -def test_get_task_sdk_does_not_inherit_shared_client_authorization(monkeypatch: pytest.MonkeyPatch): - """Service-principal task SDK auth must come from SDK headers, not stale shared-client auth.""" +def test_get_task_sdk_creates_fresh_immutable_sdk_client(monkeypatch: pytest.MonkeyPatch): + """Service-principal SDK auth should be SDK-scoped, not stored on the SDK HTTP client.""" monkeypatch.setenv( "NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"}), ) - client = shared_sync_http_client() - old_headers = dict(client.headers) - client.headers["Authorization"] = "Bearer service:jobs" - sdk = None + + sdk = get_task_sdk(as_service="jobs") + other_sdk = get_task_sdk(as_service="jobs") + scoped_sdk = sdk.with_options(set_default_headers={"X-Test": "true"}) try: - sdk = get_task_sdk(as_service="jobs") - assert sdk._client is not client + assert sdk._client is not other_sdk._client + assert scoped_sdk._client is sdk._client assert sdk.default_headers["X-NMP-Principal-Id"] == "service:jobs" assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "creator@example.com" assert "Authorization" not in sdk.default_headers assert "Authorization" not in sdk._client.headers + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + sdk._client.headers["Authorization"] = "Bearer stale" finally: - if sdk is not None: - sdk.close() - client.headers.clear() - client.headers.update(old_headers) - client.headers.pop("Authorization", None) + sdk.close() + other_sdk.close() + scoped_sdk.close() @pytest.mark.asyncio -async def test_get_async_task_sdk_does_not_inherit_shared_client_authorization(monkeypatch: pytest.MonkeyPatch): - """Async task SDK auth must also avoid stale shared-client auth.""" +async def test_get_async_task_sdk_creates_fresh_immutable_sdk_client(monkeypatch: pytest.MonkeyPatch): + """Async service-principal SDK auth should also stay off the SDK HTTP client.""" monkeypatch.setenv( "NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"}), ) - client = shared_async_http_client() - old_headers = dict(client.headers) - client.headers["Authorization"] = "Bearer service:jobs" - sdk = None + + sdk = get_async_task_sdk(as_service="jobs") + other_sdk = get_async_task_sdk(as_service="jobs") + scoped_sdk = sdk.with_options(set_default_headers={"X-Test": "true"}) try: - sdk = get_async_task_sdk(as_service="jobs") - assert sdk._client is not client + assert sdk._client is not other_sdk._client + assert scoped_sdk._client is sdk._client assert sdk.default_headers["X-NMP-Principal-Id"] == "service:jobs" assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "creator@example.com" assert "Authorization" not in sdk.default_headers assert "Authorization" not in sdk._client.headers + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + sdk._client.headers["Authorization"] = "Bearer stale" finally: - if sdk is not None: - await sdk.close() - client.headers.clear() - client.headers.update(old_headers) - client.headers.pop("Authorization", None) + await sdk.close() + await other_sdk.close() + await scoped_sdk.close() -def test_get_task_sdk_uses_workload_identity_when_token_file_configured(monkeypatch: pytest.MonkeyPatch, tmp_path): +def test_get_task_sdk_uses_workload_identity_when_token_file_configured( + monkeypatch: pytest.MonkeyPatch, tmp_path, caplog: pytest.LogCaptureFixture +): """Task SDKs should centralize the workload-token-vs-service-header choice.""" subject_token_file = tmp_path / "workload-token" subject_token_file.write_text("subject-token-from-file\n", encoding="utf-8") @@ -386,15 +512,16 @@ def token_exchange_grant(**kwargs): monkeypatch.setenv("NMP_CONFIG_FILE", str(config_file)) monkeypatch.setenv("NMP_BASE_URL", "http://nmp.example.test") monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) - monkeypatch.setenv("NMP_PRINCIPAL", json.dumps({"id": "creator@example.com", "email": "creator@example.com"})) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) monkeypatch.delenv("NMP_ACCESS_TOKEN", raising=False) monkeypatch.setattr("nemo_platform.client.factory.discover_nmp_config", lambda _base_url: _workload_oidc_config()) monkeypatch.setattr("nemo_platform.auth.workload_exchange.token_exchange_grant", token_exchange_grant) + caplog.set_level(logging.WARNING, logger="nmp.common.sdk_factory") sdk = get_task_sdk(as_service="customizer") try: request = sdk._client.build_request("GET", "http://nmp.example.test/apis/entities/v2/workspaces/default") - sdk._client._event_hooks["request"][0](request) + _apply_sync_auth(sdk, request) finally: sdk.close() @@ -403,6 +530,7 @@ def token_exchange_grant(**kwargs): assert "X-NMP-Principal-Id" not in request.headers assert "X-NMP-Principal-On-Behalf-Of" not in request.headers assert exchange_requests[0]["subject_token"] == "subject-token-from-file" + assert "will authenticate as service:customizer without on-behalf-of delegation" not in caplog.text def test_get_task_sdk_uses_explicit_sync_http_client(monkeypatch: pytest.MonkeyPatch): @@ -415,6 +543,16 @@ def test_get_task_sdk_uses_explicit_sync_http_client(monkeypatch: pytest.MonkeyP assert sdk._client is client +@pytest.mark.asyncio +async def test_get_async_task_sdk_uses_explicit_async_http_client(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + async with httpx.AsyncClient() as client: + sdk = get_async_task_sdk(as_service="customizer", http_client=client) + + assert sdk._client is client + + def test_get_request_scoped_sdk_merges_otel_and_auth_headers(): """Test that get_request_scoped_sdk merges OTEL and auth headers.""" base_sdk = get_async_platform_sdk() @@ -442,54 +580,28 @@ def test_get_request_scoped_sdk_merges_otel_and_auth_headers(): assert scoped_sdk.default_headers["X-NMP-Principal-Groups"] == "group1,group2" -def test_get_request_scoped_sdk_preserves_request_router(monkeypatch: pytest.MonkeyPatch): - """Derived request SDKs must keep the base SDK's path-aware platform request router.""" - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_SERVICES", "entities") - monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") - monkeypatch.setenv("NMP_SERVICE_PORT", "8080") - Configuration.clear_cache() - - try: - base_sdk = get_async_platform_sdk() - - with patch("nmp.common.sdk_factory.get_otel_headers", return_value={}): - with patch( - "nmp.common.sdk_factory.get_principal_auth_headers", - return_value={"X-NMP-Principal-Id": "service:models"}, - ): - scoped_sdk = get_request_scoped_sdk(base_sdk) - - prepared = scoped_sdk._prepare_url("https://nemo-gateway:8080/apis/entities/v2/workspaces") - - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - finally: - Configuration.clear_cache() +def test_get_request_scoped_sdk_reuses_base_sdk_http_client(): + """Derived request SDKs must keep the base SDK's lifecycle-owned HTTP client.""" + base_sdk = get_async_platform_sdk() + with patch("nmp.common.sdk_factory.get_otel_headers", return_value={}): + with patch( + "nmp.common.sdk_factory.get_principal_auth_headers", + return_value={"X-NMP-Principal-Id": "service:models"}, + ): + scoped_sdk = get_request_scoped_sdk(base_sdk) -def test_get_sdk_on_behalf_of_preserves_request_router(monkeypatch: pytest.MonkeyPatch): - """SDKs derived with on-behalf-of headers must still keep platform request routing.""" - monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") - monkeypatch.setenv("NMP_SERVICES", "entities") - monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") - monkeypatch.setenv("NMP_SERVICE_PORT", "8080") - Configuration.clear_cache() + assert scoped_sdk is not base_sdk + assert scoped_sdk._client is base_sdk._client - try: - base_sdk = get_async_platform_sdk(as_service="models", internal=True) - scoped_sdk = get_sdk_on_behalf_of(base_sdk, "user@example.com") - prepared = scoped_sdk._prepare_url("https://nemo-gateway:8080/apis/entities/v2/workspaces") +def test_get_sdk_on_behalf_of_reuses_base_sdk_http_client(): + """SDKs derived with on-behalf-of headers must keep the base SDK's HTTP client.""" + base_sdk = get_async_platform_sdk(as_service="models", internal=True) + scoped_sdk = get_sdk_on_behalf_of(base_sdk, "user@example.com") - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - finally: - Configuration.clear_cache() + assert scoped_sdk is not base_sdk + assert scoped_sdk._client is base_sdk._client def test_get_request_scoped_sdk_returns_base_sdk_when_no_headers(): @@ -656,221 +768,6 @@ def test_get_request_scoped_sdk_service_principal_with_on_behalf_of(): assert scoped_sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "user@example.com" -# --- Dynamic routing (service discovery map) tests --- - - -@pytest.fixture -def platform_config_with_service_discovery(): - """Platform config with service_discovery map for entities and jobs.""" - return PlatformConfig( # type: ignore[abstract] - base_url="http://platform:8080", - service_discovery={ - "entities": "http://entities-service:8080", - "jobs": "http://jobs-service:8080", - }, - ) - - -def test_resolve_platform_request_url_routes_api_path_to_service_url(platform_config_with_service_discovery): - """The named request router policy owns per-service routing.""" - - def default_resolver(url: str) -> httpx.URL: - if url.startswith("/"): - return httpx.URL(f"http://platform:8080{url}") - return httpx.URL(url) - - prepared = resolve_platform_request_url( - "/apis/entities/v2/workspaces?limit=10", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - - assert prepared.scheme == "http" - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.path == "/apis/entities/v2/workspaces" - assert prepared.query == b"limit=10" - - -def test_resolve_platform_request_url_logs_path_without_raw_url( - caplog: pytest.LogCaptureFixture, - platform_config_with_service_discovery, -): - """Routing logs expose the resolved path without query parameters.""" - - def default_resolver(url: str) -> httpx.URL: - if url.startswith("/"): - return httpx.URL(f"http://platform:8080{url}") - return httpx.URL(url) - - caplog.set_level(logging.DEBUG, logger="nmp.common.sdk_factory") - - resolve_platform_request_url( - "/health/ready?token=secret", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - resolve_platform_request_url( - "/apis/entities/v2/workspaces?token=secret", - platform_config=platform_config_with_service_discovery, - default_resolver=default_resolver, - ) - - original_record = next(record for record in caplog.records if record.message == "Routing URL to original URL") - service_record = next(record for record in caplog.records if record.message == "Routing URL to service URL") - - assert not hasattr(original_record, "url") - assert original_record.service == "unknown" - assert original_record.path == "/health/ready" - assert original_record.host == "platform" - assert original_record.port == 8080 - - assert not hasattr(service_record, "url") - assert service_record.service == "entities" - assert service_record.path == "/apis/entities/v2/workspaces" - assert service_record.host == "entities-service" - assert service_record.port == 8080 - - for record in (original_record, service_record): - assert "token=secret" not in str(record.__dict__) - - -def test_platform_request_router_uses_default_resolver_for_non_api_paths(platform_config_with_service_discovery): - """Non-API paths follow the SDK's normal URL preparation.""" - router = PlatformRequestRouter( - platform_config=platform_config_with_service_discovery, - default_resolver=lambda url: httpx.URL(f"http://platform:8080{url}"), - ) - - prepared = router.resolve("/health/ready") - - assert str(prepared) == "http://platform:8080/health/ready" - - -def test_get_platform_sdk_routes_entities_path_to_entities_service( - platform_config_with_service_discovery, -): - """Routes /apis/entities/v2/workspaces to the entities service URL.""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/entities/v2/workspaces" in str(prepared.path) - - -def test_get_platform_sdk_routes_service_path_to_env_override( - monkeypatch: pytest.MonkeyPatch, - platform_config_with_service_discovery, -): - monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-env:9090") - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-env" - assert prepared.port == 9090 - assert prepared.scheme == "http" - - -def test_get_platform_sdk_routes_jobs_path_to_jobs_service( - platform_config_with_service_discovery, -): - """Routes /apis/jobs/v2/workspaces/jobs to the jobs service URL.""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - request_url = "http://platform:8080/apis/jobs/v2/workspaces/jobs" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "jobs-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/jobs/v2/workspaces/jobs" in str(prepared.path) - - -def test_get_platform_sdk_routing_fallback_to_base_url_when_no_match( - platform_config_with_service_discovery, -): - """When the path does not match /apis/{service-name}/ (lowercase+dashes), use the original URL (base).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_platform_sdk() - # Path that does not match /apis/{service-name}/ (e.g. /api/ singular, or no such prefix) - request_url = "http://platform:8080/api/other/v1/thing" - prepared = sdk._prepare_url(request_url) - - # Should pass through to original behavior: same host as request - assert prepared.host == "platform" - assert prepared.port == 8080 - - -def test_get_async_platform_sdk_routes_entities_path_to_entities_service( - platform_config_with_service_discovery, -): - """Routes /apis/entities/v2/workspaces to the entities service URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/apis/entities/v2/workspaces" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "entities-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/entities/v2/workspaces" in str(prepared.path) - - -def test_get_async_platform_sdk_routes_jobs_path_to_jobs_service( - platform_config_with_service_discovery, -): - """Routes /apis/jobs/v2/workspaces/jobs to the jobs service URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/apis/jobs/v2/workspaces/jobs" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "jobs-service" - assert prepared.port == 8080 - assert prepared.scheme == "http" - assert "/apis/jobs/v2/workspaces/jobs" in str(prepared.path) - - -def test_get_async_platform_sdk_routing_fallback_to_base_url_when_no_match( - platform_config_with_service_discovery, -): - """When the path does not match /apis/{service-name}/, use the original URL (async SDK).""" - with patch( - "nmp.common.sdk_factory.Configuration.get_platform_config", - return_value=platform_config_with_service_discovery, - ): - sdk = get_async_platform_sdk() - request_url = "http://platform:8080/api/other/v1/thing" - prepared = sdk._prepare_url(request_url) - - assert prepared.host == "platform" - assert prepared.port == 8080 - - # --- get_entity_parts tests --- diff --git a/packages/nmp_common/tests/service/test_sdk_factory.py b/packages/nmp_common/tests/service/test_sdk_factory.py new file mode 100644 index 0000000000..f7985aa28a --- /dev/null +++ b/packages/nmp_common/tests/service/test_sdk_factory.py @@ -0,0 +1,65 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from nmp.common.platform_endpoint import PlatformEndpoint, parse_platform_endpoint +from nmp.common.service.sdk_factory import ServiceSDKFactory + + +@pytest.fixture +def platform_endpoint() -> PlatformEndpoint: + return parse_platform_endpoint("http://platform:8080") + + +def test_service_sdk_factory_creates_async_sdk_with_configured_client( + platform_endpoint, +) -> None: + http_client = MagicMock(name="http_client") + sdk = MagicMock(name="sdk") + factory = ServiceSDKFactory(platform_endpoint, async_http_client=http_client) + + with patch("nmp.common.sdk_factory.get_async_platform_sdk", return_value=sdk) as sdk_factory: + result = factory.create_async_platform_sdk(as_service="jobs", internal=True) + + assert result is sdk + sdk_factory.assert_called_once_with( + as_service="jobs", + internal=True, + http_client=http_client, + on_behalf_of=None, + base_url="http://platform:8080", + ) + + +def test_service_sdk_factory_constructs_client_from_endpoint() -> None: + platform_endpoint = MagicMock(name="platform_endpoint") + platform_endpoint.connect_base_url = "http://platform:8080" + http_client = MagicMock(name="http_client") + platform_endpoint.async_sdk_http_client.return_value = http_client + factory = ServiceSDKFactory(platform_endpoint) + + first = factory.get_async_http_client() + second = factory.get_async_http_client() + + assert first is http_client + assert second is http_client + platform_endpoint.async_sdk_http_client.assert_called_once_with(timeout=None) + + +@pytest.mark.asyncio +async def test_service_sdk_factory_closes_constructed_clients() -> None: + async_client = MagicMock(name="async_client") + async_client.aclose = AsyncMock() + sync_client = MagicMock(name="sync_client") + factory = ServiceSDKFactory( + parse_platform_endpoint("http://platform:8080"), + sync_http_client=sync_client, + async_http_client=async_client, + ) + + await factory.aclose() + + async_client.aclose.assert_awaited_once_with() + sync_client.close.assert_called_once_with() diff --git a/packages/nmp_common/tests/test_immutable_http_client.py b/packages/nmp_common/tests/test_immutable_http_client.py new file mode 100644 index 0000000000..143d2365cf --- /dev/null +++ b/packages/nmp_common/tests/test_immutable_http_client.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import httpx +import pytest +from nmp.common.immutable_http_client import ImmutableHttpClientMixin + + +def _noop_request_hook(request: httpx.Request) -> None: + pass + + +class _FrozenClient(ImmutableHttpClientMixin, httpx.Client): + def __init__(self) -> None: + super().__init__( + headers={"X-Initial": "true"}, + event_hooks={"request": [_noop_request_hook]}, + ) + self._freeze_http_client() + + +def test_immutable_sdk_client_still_builds_requests() -> None: + with _FrozenClient() as client: + request = client.build_request("GET", "http://nmp.example.test/health") + + assert request.url == "http://nmp.example.test/health" + assert request.headers["X-Initial"] == "true" + + +def test_immutable_sdk_client_blocks_client_configuration_assignment() -> None: + with _FrozenClient() as client: + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.headers = httpx.Headers({"Authorization": "Bearer stale"}) + + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.base_url = httpx.URL("http://other.example.test") + + with pytest.raises(AttributeError, match="SDK HTTP clients are immutable"): + client.params = {"debug": "true"} + + +def test_immutable_sdk_client_blocks_header_mutation() -> None: + with _FrozenClient() as client: + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers["Authorization"] = "Bearer stale" + + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers.update({"Authorization": "Bearer stale"}) + + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.headers.pop("X-Initial") + + +def test_immutable_sdk_client_blocks_cookie_mutation_and_ignores_response_cookies() -> None: + with _FrozenClient() as client: + with pytest.raises(TypeError, match="SDK HTTP clients are immutable"): + client.cookies["session"] = "stale" + + request = client.build_request("GET", "http://nmp.example.test/health") + response = httpx.Response(200, headers={"Set-Cookie": "session=stale"}, request=request) + + client.cookies.extract_cookies(response) + + assert "session" not in client.cookies + + +def test_immutable_sdk_client_blocks_event_hook_mutation() -> None: + with _FrozenClient() as client: + with pytest.raises(AttributeError): + client.event_hooks["request"].append(_noop_request_hook) + + with pytest.raises(TypeError): + client.event_hooks["request"] = [_noop_request_hook] diff --git a/packages/nmp_common/tests/test_platform_endpoint.py b/packages/nmp_common/tests/test_platform_endpoint.py index 2b4911b53e..7157731fa9 100644 --- a/packages/nmp_common/tests/test_platform_endpoint.py +++ b/packages/nmp_common/tests/test_platform_endpoint.py @@ -1,12 +1,31 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import logging +import os +import threading +import time from pathlib import Path from unittest.mock import patch import pytest from nmp.common.config import PlatformConfig -from nmp.common.platform_endpoint import UDS_BASE_URL, parse_platform_endpoint, resolve_service_endpoint +from nmp.common.platform_endpoint import ( + UDS_BASE_URL, + _SyncPlatformEndpointRoutingTransport, + parse_platform_endpoint, + resolve_platform_endpoint, + resolve_service_endpoint, +) + + +@pytest.fixture(autouse=True) +def clear_service_url_env_vars(monkeypatch: pytest.MonkeyPatch) -> None: + for key in tuple(os.environ): + if key == "NMP_BASE_URL": + continue + if key.startswith("NMP_") and key.endswith("_URL"): + monkeypatch.delenv(key, raising=False) def test_parse_tcp_endpoint() -> None: @@ -80,6 +99,175 @@ def test_endpoint_env_family_is_not_part_of_contract(monkeypatch: pytest.MonkeyP assert endpoint.connect_base_url == "http://platform:8080" +@pytest.fixture +def platform_config_with_service_discovery() -> PlatformConfig: + return PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={ + "entities": "http://entities-service:8080", + "jobs": "http://jobs-service:8080", + }, + ) + + +def test_resolve_platform_endpoint_carries_service_routes( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + assert endpoint.connect_base_url == "http://platform:8080" + assert endpoint.service_endpoints["entities"].connect_base_url == "http://entities-service:8080" + assert endpoint.service_endpoints["jobs"].connect_base_url == "http://jobs-service:8080" + + +def test_resolve_platform_endpoint_rejects_malformed_service_route() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={ + "entities": "http://entities-service:8080", + "jobs": "not-a-url", + }, + ) + + with pytest.raises(ValueError, match="Unsupported platform endpoint URL 'not-a-url'"): + resolve_platform_endpoint(config) + + +def test_resolve_platform_endpoint_rejects_malformed_base_url() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="not-a-url", + service_discovery={"entities": "http://entities-service:8080"}, + ) + + with pytest.raises(ValueError, match="Unsupported platform endpoint URL 'not-a-url'"): + resolve_platform_endpoint(config) + + +def test_platform_endpoint_routes_api_path_to_service_url( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?limit=10") + + assert routed.endpoint.connect_base_url == "http://entities-service:8080" + assert routed.url.scheme == "http" + assert routed.url.host == "entities-service" + assert routed.url.port == 8080 + assert routed.url.path == "/apis/entities/v2/workspaces" + assert routed.url.query == b"limit=10" + + +def test_platform_endpoint_preserves_service_url_path_prefix() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={"entities": "http://entities-service:8080/entities-prefix"}, + ) + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?limit=10") + + assert routed.endpoint.connect_base_url == "http://entities-service:8080/entities-prefix" + assert routed.url.scheme == "http" + assert routed.url.host == "entities-service" + assert routed.url.port == 8080 + assert routed.url.path == "/entities-prefix/apis/entities/v2/workspaces" + assert routed.url.query == b"limit=10" + + +def test_platform_endpoint_uses_env_override( + monkeypatch: pytest.MonkeyPatch, + platform_config_with_service_discovery: PlatformConfig, +) -> None: + monkeypatch.setenv("NMP_ENTITIES_URL", "http://entities-env:9090") + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces") + + assert routed.endpoint.connect_base_url == "http://entities-env:9090" + assert routed.url.host == "entities-env" + assert routed.url.port == 9090 + + +def test_platform_endpoint_uses_env_route_for_configured_service(monkeypatch: pytest.MonkeyPatch) -> None: + config = PlatformConfig(base_url="http://platform:8080", services="hello-world") # type: ignore[abstract] + monkeypatch.setenv("NMP_HELLO_WORLD_URL", "http://hello-world-service:8080") + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/hello-world/v2/workspaces/default/hello") + + assert routed.endpoint.connect_base_url == "http://hello-world-service:8080" + assert routed.url.host == "hello-world-service" + + +def test_platform_endpoint_ignores_unknown_service_url_env_var(monkeypatch: pytest.MonkeyPatch) -> None: + config = PlatformConfig(base_url="http://platform:8080") # type: ignore[abstract] + monkeypatch.setenv("NMP_NOT_A_SERVICE_URL", "http://not-a-service:8080") + + endpoint = resolve_platform_endpoint(config) + routed = endpoint.route_request_url("http://platform:8080/apis/not-a-service/v2/workspaces/default") + + assert "not-a-service" not in endpoint.service_endpoints + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert routed.url.host == "platform" + + +def test_platform_endpoint_routes_uds_service() -> None: + config = PlatformConfig( # type: ignore[abstract] + base_url="http://platform:8080", + service_discovery={"entities": "unix:///tmp/entities.sock"}, + ) + endpoint = resolve_platform_endpoint(config) + + routed = endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces") + + assert routed.endpoint.transport == "uds" + assert routed.endpoint.socket_path == Path("/tmp/entities.sock") + assert str(routed.url) == "http://nemo-platform.local/apis/entities/v2/workspaces" + + +def test_platform_endpoint_keeps_non_api_path_on_default_endpoint( + platform_config_with_service_discovery: PlatformConfig, +) -> None: + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + routed = endpoint.route_request_url("http://platform:8080/health/ready") + + assert routed.endpoint.connect_base_url == "http://platform:8080" + assert str(routed.url) == "http://platform:8080/health/ready" + + +def test_platform_endpoint_logs_path_without_raw_url( + caplog: pytest.LogCaptureFixture, + platform_config_with_service_discovery: PlatformConfig, +) -> None: + caplog.set_level(logging.DEBUG, logger="nmp.common.platform_endpoint") + endpoint = resolve_platform_endpoint(platform_config_with_service_discovery) + + endpoint.route_request_url("http://platform:8080/health/ready?token=secret") + endpoint.route_request_url("http://platform:8080/apis/entities/v2/workspaces?token=secret") + + default_record = next( + record for record in caplog.records if record.message == "Routing SDK URL to default endpoint" + ) + service_record = next( + record for record in caplog.records if record.message == "Routing SDK URL to service endpoint" + ) + + assert not hasattr(default_record, "url") + assert default_record.service == "unknown" + assert default_record.path == "/health/ready" + + assert not hasattr(service_record, "url") + assert service_record.service == "entities" + assert service_record.path == "/apis/entities/v2/workspaces" + assert service_record.host == "entities-service" + assert service_record.port == 8080 + + for record in (default_record, service_record): + assert "token=secret" not in str(record.__dict__) + + def test_sync_http_client_omits_unset_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") @@ -101,7 +289,7 @@ def test_sync_http_client_passes_explicit_timeout() -> None: def test_sync_sdk_http_client_uses_sdk_default_for_tcp() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as client: endpoint.sync_sdk_http_client() client.assert_called_once_with() @@ -110,7 +298,7 @@ def test_sync_sdk_http_client_uses_sdk_default_for_tcp() -> None: def test_sync_sdk_http_client_passes_explicit_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as client: endpoint.sync_sdk_http_client(timeout=2.0) client.assert_called_once_with(timeout=2.0) @@ -132,8 +320,8 @@ def test_uds_sync_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> N endpoint = parse_platform_endpoint("unix:///tmp/nemo-platform.sock") with ( - patch("nmp.common.platform_endpoint.DefaultHttpxClient") as sdk_client, - patch("nmp.common.platform_endpoint.httpx.Client") as client, + patch("nmp.common.platform_endpoint.ImmutableDefaultHttpxClient") as sdk_client, + patch("nmp.common.platform_endpoint.ImmutableHttpxClient") as client, ): endpoint.sync_sdk_http_client() @@ -144,6 +332,51 @@ def test_uds_sync_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> N assert "timeout" not in kwargs +def test_sync_routing_transport_reuses_uds_transport_concurrently() -> None: + endpoint = parse_platform_endpoint("unix:///tmp/entities.sock") + barrier = threading.Barrier(8) + constructed_transports: list[object] = [] + closed_transports: list[object] = [] + returned_transports: list[object] = [] + errors: list[BaseException] = [] + + class FakeHTTPTransport: + def __init__(self, *, uds: str | None = None) -> None: + self.uds = uds + if uds is not None: + constructed_transports.append(self) + time.sleep(0.01) + + def close(self) -> None: + if self.uds is not None: + closed_transports.append(self) + + def get_transport() -> None: + try: + barrier.wait(timeout=2.0) + returned_transports.append(routing_transport._transport_for_endpoint(endpoint)) + except BaseException as error: + errors.append(error) + + with patch("nmp.common.platform_endpoint.httpx.HTTPTransport", FakeHTTPTransport): + routing_transport = _SyncPlatformEndpointRoutingTransport( + endpoint=parse_platform_endpoint("http://platform:8080") + ) + threads = [threading.Thread(target=get_transport) for _ in range(8)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=2.0) + + routing_transport.close() + + assert errors == [] + assert all(not thread.is_alive() for thread in threads) + assert len(constructed_transports) == 1 + assert len({id(transport) for transport in returned_transports}) == 1 + assert closed_transports == constructed_transports + + def test_async_http_client_omits_unset_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") @@ -165,7 +398,7 @@ def test_async_http_client_passes_explicit_timeout() -> None: def test_async_sdk_http_client_uses_sdk_default_for_tcp() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as client: endpoint.async_sdk_http_client() client.assert_called_once_with() @@ -174,7 +407,7 @@ def test_async_sdk_http_client_uses_sdk_default_for_tcp() -> None: def test_async_sdk_http_client_passes_explicit_timeout() -> None: endpoint = parse_platform_endpoint("http://127.0.0.1:8080") - with patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as client: + with patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as client: endpoint.async_sdk_http_client(timeout=2.0) client.assert_called_once_with(timeout=2.0) @@ -196,8 +429,8 @@ def test_uds_async_sdk_http_client_keeps_transport_and_omits_unset_timeout() -> endpoint = parse_platform_endpoint("unix:///tmp/nemo-platform.sock") with ( - patch("nmp.common.platform_endpoint.DefaultAsyncHttpxClient") as sdk_client, - patch("nmp.common.platform_endpoint.httpx.AsyncClient") as client, + patch("nmp.common.platform_endpoint.ImmutableDefaultAsyncHttpxClient") as sdk_client, + patch("nmp.common.platform_endpoint.ImmutableAsyncHttpxClient") as client, ): endpoint.async_sdk_http_client() diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py b/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py index ed6ea42c37..e9a01e32d0 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/loader.py @@ -9,13 +9,13 @@ import logging import threading from collections.abc import Callable -from typing import cast from nmp.common.service import Service from nmp.common.service.deptree import resolve_service_loading_order logger = logging.getLogger(__name__) + ControllerRunFunc = Callable[[threading.Event], object] @@ -66,4 +66,4 @@ def load_controller_run_func(controller_name: str, import_path: str) -> Controll raise TypeError(f"Controller {controller_name} must be a callable, got {type(run_func)}") logger.debug("Loaded controller %s", controller_name) - return cast(ControllerRunFunc, run_func) + return run_func diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py b/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py index 1d074b64cb..e004a9e3b7 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/plugin_adapter.py @@ -11,7 +11,8 @@ Wraps a :class:`~nemo_platform_plugin.controller.NemoController` (async) as a platform-native :class:`~nmp.common.controller.Controller` (sync / thread-based). Use :func:`make_controller_run_func` to create the - ``run(stop_signal)`` callable expected by :func:`~nmp.platform_runner.server.create_app`. + ``run(stop_signal)`` callable expected by + :func:`~nmp.platform_runner.server.create_app`. """ from __future__ import annotations @@ -166,7 +167,6 @@ def make_controller_run_func(controller_cls: type[NemoController]) -> Callable[[ The returned callable matches the signature expected by :func:`~nmp.platform_runner.server.create_app` (same as core controller run functions registered in ``AVAILABLE_CONTROLLERS``). - Lifecycle inside the returned function: 1. Instantiate *controller_cls*. diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/run.py b/packages/nmp_platform_runner/src/nmp/platform_runner/run.py index d3a54eac98..4077a8109e 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/run.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/run.py @@ -11,7 +11,6 @@ import threading import time from collections.abc import Callable, Mapping -from typing import cast from nmp.common.config import get_auth_config, get_common_service_config, get_platform_config, get_service_config from nmp.common.observability import initialize_obs, setup_global_instrumentations @@ -69,7 +68,12 @@ def run_controllers_in_threads( """Start controller run functions in daemon threads.""" threads = [] for name, run_func in controller_run_funcs.items(): - thread = threading.Thread(target=run_func, args=(stop_signal,), name=f"controller-{name}", daemon=True) + thread = threading.Thread( + target=run_func, + args=(stop_signal,), + name=f"controller-{name}", + daemon=True, + ) thread.start() threads.append(thread) return threads @@ -215,10 +219,10 @@ def _load_run_functions( t0 = time.perf_counter() value = registry[name] try: - if callable(value): - run_funcs[name] = cast(ControllerRunFunc, value) - else: + if isinstance(value, str): run_funcs[name] = load_controller_run_func(name, value) + else: + run_funcs[name] = value except (ImportError, TypeError, AttributeError, ValueError) as error: logger.error("Failed to load %s %s: %s", kind, name, error) raise SystemExit(1) from error diff --git a/packages/nmp_platform_runner/src/nmp/platform_runner/server.py b/packages/nmp_platform_runner/src/nmp/platform_runner/server.py index eef0e55d2c..78f336ca39 100644 --- a/packages/nmp_platform_runner/src/nmp/platform_runner/server.py +++ b/packages/nmp_platform_runner/src/nmp/platform_runner/server.py @@ -10,9 +10,8 @@ import logging import os import threading -from collections.abc import Callable, Mapping, MutableMapping +from collections.abc import Mapping, MutableMapping from contextlib import asynccontextmanager -from typing import cast import httpx import uvicorn @@ -21,7 +20,6 @@ from fastapi.responses import JSONResponse from nmp.common.auth import AuthorizationMiddleware from nmp.common.config import get_auth_config, get_platform_config -from nmp.common.http_clients import close_shared_http_clients from nmp.common.observability import initialize_obs, setup_fastapi_instrumentations, setup_global_instrumentations from nmp.common.observability.context import create_app_context_dependency from nmp.common.pyleak import detect_blocking @@ -158,6 +156,10 @@ def create_app( ) ) + if http_client is not None: + for service_instance in services: + service_instance.dependency_provider.configure_http_client(http_client) + @asynccontextmanager async def lifespan(app: FastAPI): logger.info("Starting Nemo Platform server") @@ -215,7 +217,6 @@ async def run_platform_seed_and_update_readiness() -> None: for thread in controller_threads: thread.join(timeout=5) - await close_shared_http_clients() logger.info("Shutting down Nemo Platform API server") app = FastAPI( @@ -281,9 +282,9 @@ async def root_handler() -> Response: def _load_run_functions( names: list[str], - registry: Mapping[str, str | Callable[[threading.Event], object]], -) -> dict[str, Callable[[threading.Event], object]]: - run_funcs: dict[str, Callable[[threading.Event], object]] = {} + registry: Mapping[str, str | ControllerRunFunc], +) -> dict[str, ControllerRunFunc]: + run_funcs: dict[str, ControllerRunFunc] = {} for name in names: value = registry[name] if isinstance(value, str): @@ -434,9 +435,9 @@ def create_default_app() -> FastAPI: "Unknown controller %r requested via NMP_CONTROLLERS=%r. Available controllers: %s" % (controller_name, controller_names_env, available) ) - if callable(controller_value): - controller_run_funcs[controller_name] = cast(ControllerRunFunc, controller_value) - else: + if isinstance(controller_value, str): controller_run_funcs[controller_name] = load_controller_run_func(controller_name, controller_value) + else: + controller_run_funcs[controller_name] = controller_value return create_app(services, controller_run_funcs=controller_run_funcs) diff --git a/packages/nmp_platform_runner/tests/test_run.py b/packages/nmp_platform_runner/tests/test_run.py index def553aadc..df2866163b 100644 --- a/packages/nmp_platform_runner/tests/test_run.py +++ b/packages/nmp_platform_runner/tests/test_run.py @@ -73,7 +73,7 @@ def test_run_platform_marks_loaded_services_local_before_starting_controllers(mo monkeypatch.setattr( runner, "_load_run_functions", - lambda names, registry, kind: {"jobs": lambda stop_signal: None} if kind == "controller" else {}, + lambda names, registry, kind: {"jobs": lambda _stop_signal: None} if kind == "controller" else {}, ) monkeypatch.setattr(runner, "_display_banner", lambda **_: None) monkeypatch.setattr(runner, "run_server", lambda services, host, port, socket_path=None: None) diff --git a/packages/nmp_platform_runner/tests/test_server.py b/packages/nmp_platform_runner/tests/test_server.py index 2087de75f1..c32d6f5935 100644 --- a/packages/nmp_platform_runner/tests/test_server.py +++ b/packages/nmp_platform_runner/tests/test_server.py @@ -10,10 +10,12 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient from nmp.common.config import AuthConfig, Configuration +from nmp.common.config import get_platform_config as load_platform_config from nmp.common.config.base import OIDCConfig from nmp.common.service import Service from nmp.platform_runner import config as runner_config @@ -174,13 +176,49 @@ def test_create_app_marks_mounted_services_as_local(monkeypatch): assert platform_cfg.services == "agents" -def test_create_app_mounted_services_drive_sdk_local_routing_without_services_env(monkeypatch): +@pytest.mark.asyncio +async def test_create_app_configures_service_dependency_provider_http_client(monkeypatch): + _patch_platform_app_config(monkeypatch, seed_on_startup=False) + service = PluginService() + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + + try: + server.create_app([service], http_client=http_client) + + assert service.dependency_provider.get_http_client() is http_client + finally: + await http_client.aclose() + + +@pytest.mark.asyncio +async def test_create_app_starts_controller_with_stop_signal(monkeypatch): + _patch_platform_app_config(monkeypatch, seed_on_startup=False) + captured: dict[str, object] = {} + started = threading.Event() + + def run_controller(stop_signal: threading.Event) -> None: + captured["stop_signal"] = stop_signal + started.set() + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: httpx.Response(200))) + app = server.create_app(controller_run_funcs={"test": run_controller}, http_client=http_client) + + try: + with TestClient(app): + assert started.wait(timeout=2) + finally: + await http_client.aclose() + + assert isinstance(captured["stop_signal"], threading.Event) + + +def test_create_app_mounted_services_drive_lifecycle_routing_without_services_env(monkeypatch): monkeypatch.delenv("NMP_SERVICES", raising=False) monkeypatch.setenv("NMP_BASE_URL", "https://nemo-gateway:8080") monkeypatch.setenv("NMP_SERVICE_HOST", "127.0.0.1") monkeypatch.setenv("NMP_SERVICE_PORT", "8080") Configuration.clear_cache() - platform_cfg = Configuration.get_platform_config() + platform_cfg = load_platform_config() try: auth_cfg = _make_auth_config(enabled=False) @@ -188,19 +226,19 @@ def test_create_app_mounted_services_drive_sdk_local_routing_without_services_en monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) import nmp.common.auth.middleware as auth_middleware - from nmp.common.sdk_factory import get_platform_sdk + from nmp.common.platform_endpoint import resolve_platform_endpoint monkeypatch.setattr(auth_middleware, "get_auth_config", lambda: auth_cfg) server.create_app(services=[PluginService()]) - sdk = get_platform_sdk() - prepared = sdk._prepare_url("https://nemo-gateway:8080/apis/agents/v2/example") + endpoint = resolve_platform_endpoint(platform_cfg) + routed = endpoint.route_request_url("https://nemo-gateway:8080/apis/agents/v2/example") assert platform_cfg.services == "agents" - assert prepared.scheme == "http" - assert prepared.host == "127.0.0.1" - assert prepared.port == 8080 + assert routed.url.scheme == "http" + assert routed.url.host == "127.0.0.1" + assert routed.url.port == 8080 finally: Configuration.clear_cache() @@ -360,6 +398,7 @@ def test_create_default_app_raises_for_unknown_controller_from_env(monkeypatch): def _make_platform_config_mock(*, redirect_root_to_studio: bool = True) -> MagicMock: cfg = MagicMock() + cfg.base_url = "http://platform.local" cfg.seed_on_startup = False cfg.redirect_root_to_studio = redirect_root_to_studio return cfg diff --git a/packages/nmp_platform_runner/tests/test_sidecars.py b/packages/nmp_platform_runner/tests/test_sidecars.py index e3950b3a93..4dcfddc54c 100644 --- a/packages/nmp_platform_runner/tests/test_sidecars.py +++ b/packages/nmp_platform_runner/tests/test_sidecars.py @@ -105,6 +105,7 @@ def test_create_app_starts_and_stops_dummy_sidecar_with_lifespan() -> None: patch("nmp.platform_runner.server.get_auth_config", return_value=_auth_config(False)), patch("nmp.common.auth.middleware.get_auth_config", return_value=_auth_config(False)), ): + platform_config.return_value.base_url = "http://platform.local" platform_config.return_value.seed_on_startup = False platform_config.return_value.redirect_root_to_studio = False app = server.create_app( @@ -130,6 +131,7 @@ def test_build_platform_app_loads_dependent_sidecar_into_lifespan(monkeypatch: p patch("nmp.platform_runner.server.get_auth_config", return_value=_auth_config(False)), patch("nmp.common.auth.middleware.get_auth_config", return_value=_auth_config(False)), ): + platform_config.return_value.base_url = "http://platform.local" platform_config.return_value.seed_on_startup = False platform_config.return_value.redirect_root_to_studio = False app = server.build_platform_app(runner_config.PlatformAppConfig(services=["models"], controllers=[]), env={}) @@ -187,14 +189,18 @@ def join(self) -> None: monkeypatch.delenv("VLLM_ENDPOINT", raising=False) monkeypatch.setattr(adapters_main, "get_platform_config", lambda: MagicMock(base_url="http://platform.local")) - monkeypatch.setattr(adapters_main, "get_platform_sdk", lambda **_kwargs: MagicMock()) monkeypatch.setattr(adapters_main.asyncio, "new_event_loop", lambda: MagicMock()) monkeypatch.setattr(adapters_main, "Loop", FakeLoop) monkeypatch.setattr(adapters_main, "TimedLoopWaiter", lambda *_args, **_kwargs: object()) monkeypatch.setattr(adapters_main.ControllerManager, "get_instance", classmethod(lambda _cls: manager)) + monkeypatch.setattr(adapters_main, "get_platform_sdk", MagicMock(return_value=MagicMock())) stop_signal = threading.Event() - thread = threading.Thread(target=adapters_main.run, args=(stop_signal,), daemon=True) + thread = threading.Thread( + target=adapters_main.run, + args=(stop_signal,), + daemon=True, + ) try: thread.start() diff --git a/packages/nmp_testing/src/nmp/testing/client.py b/packages/nmp_testing/src/nmp/testing/client.py index 14f9e7830c..50176cf885 100644 --- a/packages/nmp_testing/src/nmp/testing/client.py +++ b/packages/nmp_testing/src/nmp/testing/client.py @@ -442,12 +442,6 @@ def _add_service( services_to_start = [_create_svc(svc, configs) for svc in services_to_create] services_to_start = order_services_by_dependencies(services_to_start) - # Clear any stale SDK client from previous tests BEFORE creating app. - # This prevents service startup code from using a previous test's http transport. - from nmp.common import sdk_factory as sdk_factory_module - - sdk_factory_module._test_http_client = None - # Create transport and http_client BEFORE the app, so we can inject the client # into create_app() for middleware (AuthorizationMiddleware). We set transport.app # after app creation - this works because no requests are made until setup completes. @@ -475,27 +469,11 @@ async def _pending_asgi_app(scope: Scope, receive: Receive, send: Send) -> None: # Store on app.state so tests can access it via test_client.app.state.access_log app.state.access_log = access_log_instance - # Configure module-level http client as FALLBACK for direct callers of - # get_async_platform_sdk()/get_platform_sdk() that don't use DependencyProvider. - # The primary injection path is through DependencyProvider (see below). These - # module-level variables will be removed once all direct callers are migrated. - # See architecture/docs/http-client-injection.md for details. - sdk_factory_module._test_http_client = async_http_client - stack.callback(lambda: setattr(sdk_factory_module, "_test_http_client", None)) - async_sdk = AsyncNeMoPlatform(base_url="http://testserver", http_client=async_http_client, workspace=workspace) # Create the EntityClient (used for DI and optionally yielded) entity_client = EntityClient(AsyncEntitiesResource(async_sdk)) - # Inject ASGI-transport clients into each service's DependencyProvider. - # This is critical for services that call dependency_provider.get_sdk_client() - # directly (e.g., in on_startup for background tasks like auth policy refresh). - # See architecture/docs/http-client-injection.md for details. - for svc in services_to_start: - svc.dependency_provider._http_client = async_http_client - svc.dependency_provider._sdk_client = async_sdk - # Merge dependency overrides all_overrides = {} if dependency_overrides: diff --git a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_health_child.py b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_health_child.py index 683e7b3720..5e9e57940d 100644 --- a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_health_child.py +++ b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_health_child.py @@ -53,7 +53,11 @@ def controller_run(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -87,7 +91,11 @@ def stubborn_controller(stop_signal: threading.Event) -> None: patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") @@ -107,25 +115,21 @@ def stubborn_controller(stop_signal: threading.Event) -> None: @pytest.mark.integration -def test_lifespan_cleanup_runs_on_app_shutdown() -> None: - """``close_shared_http_clients`` should be called during lifespan teardown.""" - cleanup_called = threading.Event() - +def test_lifespan_shutdown_marks_controller_stop_signal() -> None: + """Lifespan teardown should mark the controller stop signal.""" with ( patch("nmp.platform_runner.server.get_platform_config") as mock_pc, patch("nmp.platform_runner.server.get_auth_config") as mock_ac, patch("nmp.common.auth.middleware.get_auth_config") as mock_ac2, - patch("nmp.platform_runner.server.close_shared_http_clients") as mock_close, ): - mock_pc.return_value = MagicMock(seed_on_startup=False, redirect_root_to_studio=False) + mock_pc.return_value = MagicMock( + base_url="http://platform.local", + seed_on_startup=False, + redirect_root_to_studio=False, + ) mock_ac.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") mock_ac2.return_value = MagicMock(enabled=False, policy_decision_point_provider="embedded") - async def fake_close(): - cleanup_called.set() - - mock_close.side_effect = fake_close - from nmp.platform_runner.server import create_app app = create_app(services=[]) @@ -135,7 +139,7 @@ async def fake_close(): with TestClient(app): pass - assert cleanup_called.is_set(), "close_shared_http_clients was not called during shutdown" + assert app.state.controller_stop_signal.is_set() # --------------------------------------------------------------------------- diff --git a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services_contract.py b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services_contract.py index 747a3b0850..ee5d2309bf 100644 --- a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services_contract.py +++ b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services_contract.py @@ -274,7 +274,7 @@ def get_routers(self): return [] dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": sidecar_run_func} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": sidecar_run_func} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -298,6 +298,7 @@ def get_routers(self): monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) diff --git a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_sidecar_integration.py b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_sidecar_integration.py index ae813ac3f4..5ec0a470a2 100644 --- a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_sidecar_integration.py +++ b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_sidecar_integration.py @@ -42,7 +42,7 @@ def get_routers(self): def _sidecar_with_events(started: threading.Event, stopped: threading.Event) -> Callable[[threading.Event], None]: - """Return a sidecar ``run(stop_signal)`` that signals start/stop via events.""" + """Return a sidecar run function that signals start/stop via events.""" def run(stop_signal: threading.Event) -> None: started.set() @@ -71,7 +71,7 @@ def patched_registry( test sidecar, plus minimal auth/platform config stubs.""" started, stopped = sidecar_events dummy_services: dict[str, Service] = {"models": _DummyService()} - dummy_sidecars: dict[str, Callable] = {"adapters": _sidecar_with_events(started, stopped)} + dummy_sidecars: dict[str, Callable[[threading.Event], None]] = {"adapters": _sidecar_with_events(started, stopped)} monkeypatch.setattr(runner_config, "get_available_services", lambda: dummy_services) monkeypatch.setattr(runner_config, "get_available_controllers", lambda: {}) @@ -95,6 +95,7 @@ def patched_registry( monkeypatch.setattr(server, "get_auth_config", lambda: auth_cfg) monkeypatch.setattr("nmp.common.auth.middleware.get_auth_config", lambda: auth_cfg) platform_cfg = MagicMock() + platform_cfg.base_url = "http://platform.local" platform_cfg.seed_on_startup = False platform_cfg.redirect_root_to_studio = False monkeypatch.setattr(server, "get_platform_config", lambda: platform_cfg) diff --git a/services/core/entities/src/nmp/core/entities/controllers/main.py b/services/core/entities/src/nmp/core/entities/controllers/main.py index a5a7416980..869562ffc9 100644 --- a/services/core/entities/src/nmp/core/entities/controllers/main.py +++ b/services/core/entities/src/nmp/core/entities/controllers/main.py @@ -28,7 +28,7 @@ def handle_signal(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: platform_config = get_platform_config() logger.info("Starting entities controller") @@ -39,10 +39,7 @@ def run(parent_stop_signal: threading.Event | None = None): else: local_stop_signal = parent_stop_signal - nmp_sdk = get_async_platform_sdk( - as_service="entities", - internal=True, - ) + nmp_sdk = get_async_platform_sdk(as_service="entities", internal=True) # Create a single event loop that will be shared for DB init and the cleanup controller, # so SQLAlchemy's async pool is bound to the same loop that later runs queries. diff --git a/services/core/files/src/nmp/core/files/api/endpoint_helpers.py b/services/core/files/src/nmp/core/files/api/endpoint_helpers.py index 7bdc8b7685..2a8f5fcb80 100644 --- a/services/core/files/src/nmp/core/files/api/endpoint_helpers.py +++ b/services/core/files/src/nmp/core/files/api/endpoint_helpers.py @@ -302,7 +302,13 @@ async def resolve_storage_secrets_for_user( auth_client: AuthClient, ) -> dict[str, str]: """Resolve storage secrets using delegated headers on request-scoped SDK.""" - service_sdk = get_async_platform_sdk(as_service="files", internal=True, on_behalf_of=auth_client.principal.id) + service_sdk = get_async_platform_sdk( + as_service="files", + internal=True, + http_client=sdk._client, + on_behalf_of=auth_client.principal.id, + base_url=str(sdk.base_url).rstrip("/"), + ) return await resolve_storage_secrets(storage, workspace, service_sdk) diff --git a/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py b/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py index 3b023cb4ff..78f09fec28 100644 --- a/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py +++ b/services/core/inference-gateway/src/nmp/core/inference_gateway/service.py @@ -10,7 +10,6 @@ import aiohttp from fastapi import Request from fastapi.exceptions import RequestValidationError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service import RouterConfig, Service from starlette import status from starlette.responses import JSONResponse @@ -96,7 +95,7 @@ async def on_startup(self) -> None: from nmp.core.inference_gateway.api.virtual_model_cache import VirtualModelCache from nmp.core.inference_gateway.config import config as inference_gateway_config - sdk = get_async_platform_sdk(as_service="inference-gateway", internal=True) + sdk = self.dependency_provider.get_sdk_client(as_service="inference-gateway") # Initialize caches model_cache = set_global_model_cache(ModelCache(secret_value_ttl=inference_gateway_config.secrets_ttl_sec)) diff --git a/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py b/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py index 3c6cda1236..c8c318b2ed 100644 --- a/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py +++ b/services/core/inference-gateway/src/nmp/core/inference_gateway/testing/fixtures.py @@ -60,29 +60,38 @@ def _enable_post_response_task_tracking(client_context: ClientContext) -> None: _app_from(client_context).state.pending_post_response_tasks = [] +@contextmanager +def _loopback_plugin_sdk_provider(base_url: str) -> Generator[None, None, None]: + """Make plugin-created SDKs target the loopback test app.""" + from unittest.mock import patch + + from nemo_platform_plugin import sdk_provider as sdk_provider_module + from nemo_platform_plugin.sdk_provider import DefaultSDKProvider + + previous_provider = getattr(sdk_provider_module, "_cached_provider", None) + with patch.dict("os.environ", {"NMP_BASE_URL": base_url}): + sdk_provider_module.set_sdk_provider(DefaultSDKProvider()) + try: + yield + finally: + sdk_provider_module.set_sdk_provider(previous_provider) + + @contextmanager def _build_app_context( *extra_services: ServiceFactory, ) -> Generator[ClientContext, None, None]: """Yield an IGW + Models + extras :class:`ClientContext` (module-lived). - Two module-scope hazards are neutralised here: - - 1. The 3-second background ``refresh_model_cache_task``. ``on_startup`` - reads ``refresh_model_cache_interval_sec`` from the module-level - config snapshot (captured at first import), so a ``service_configs`` - override is too late. Patch the snapshot field to 0 *before* - entering ``create_test_client`` and ``on_startup`` never schedules - the loop. - 2. The shared SDK HTTP client's ``aclose``. Plugins like - ``nemo-guardrails`` call ``await sdk.close()`` in ``on_shutdown``, - which would close the shared client for every later test in the - module. Patch ``aclose`` to a no-op for the module's lifetime; - ``ASGITransport`` is in-process so nothing actually leaks. + The 3-second background ``refresh_model_cache_task`` is neutralised here. + ``on_startup`` reads ``refresh_model_cache_interval_sec`` from the + module-level config snapshot (captured at first import), so a + ``service_configs`` override is too late. Patch the snapshot field to 0 + *before* entering ``create_test_client`` and ``on_startup`` never schedules + the loop. """ from unittest.mock import patch - from nmp.common import sdk_factory as sdk_factory_module from nmp.core.inference_gateway import config as igw_config_module from nmp.core.inference_gateway.service import InferenceGatewayService from nmp.core.models.service import ModelsService @@ -95,21 +104,7 @@ def _build_app_context( client_type=ClientContext, igw_mock_provider_mode=False, ) as client_context: - shared_async_client = sdk_factory_module._test_http_client - if shared_async_client is None: - yield client_context - return - - original_aclose = shared_async_client.aclose - - async def _noop_aclose() -> None: - return None - - shared_async_client.aclose = _noop_aclose # type: ignore[method-assign] - try: - yield client_context - finally: - shared_async_client.aclose = original_aclose # type: ignore[method-assign] + yield client_context @pytest.fixture(scope="module") @@ -226,6 +221,8 @@ def _build_loopback_harness( * ``get_platform_config`` is patched at IGW's middleware-registry import site so :meth:`get_openai_compatible_inference_url_and_model` returns URLs reachable from the test process. + * The plugin SDK provider is temporarily rebound to the loopback URL + so middleware can fetch platform entities created by the test SDK. Both patches roll back before the next test runs, so a plain ``igw_plugin_harness`` test sharing the module doesn't observe them. @@ -269,6 +266,7 @@ def _restore_http_client_override() -> None: stack.callback(_restore_http_client_override) stack.enter_context(override_platform_base_url(igw_loopback_base_url)) + stack.enter_context(_loopback_plugin_sdk_provider(igw_loopback_base_url)) harness = cast( IGWLoopbackHarness, diff --git a/services/core/inference-gateway/tests/integration/conftest.py b/services/core/inference-gateway/tests/integration/conftest.py index baeabe10a3..99160567c2 100644 --- a/services/core/inference-gateway/tests/integration/conftest.py +++ b/services/core/inference-gateway/tests/integration/conftest.py @@ -90,6 +90,10 @@ def init(self) -> None: """No-op init for mock backend.""" pass + def shutdown(self) -> None: + """No-op shutdown for mock backend.""" + pass + async def create_model_deployment(self, ctx: Any) -> DeploymentStatusUpdate: """Record call and return configured response.""" self.create_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) @@ -105,9 +109,16 @@ async def get_model_deployment_status(self, ctx: Any) -> DeploymentStatusUpdate: self.status_calls.append(ctx.model_deployment) return self.status_response - async def delete_model_deployment(self, deployment: Any) -> DeploymentStatusUpdate: + async def delete_model_deployment( + self, + workspace: str, + name: str, + *, + deleting_elapsed_seconds: float | None = None, + ) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.delete_calls.append(deployment) + del deleting_elapsed_seconds + self.delete_calls.append((workspace, name)) return self.delete_response @@ -352,7 +363,10 @@ def patched_get_qualified_image(name: str, tag=None, registry=None): return_value=mock_platform_config, ), patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk") as mock_sdk_factory, - patch("nemo_platform_plugin.sdk_provider.get_async_platform_sdk") as mock_sdk, + patch( + "nmp.core.models.controllers.backends.deployments_plugin.backend.get_async_platform_sdk" + ) as mock_backend_sdk, + patch("nemo_deployments_plugin.controller.get_async_platform_sdk") as mock_deployments_sdk, patch("nemo_deployments_plugin.config.DeploymentsConfig.get", return_value=deployments_config), patch( "nemo_platform_plugin.jobs.image.get_qualified_image", @@ -360,23 +374,21 @@ def patched_get_qualified_image(name: str, tag=None, registry=None): ), ): mock_sdk_factory.return_value = test_clients.async_sdk - mock_sdk.return_value = test_clients.async_sdk + mock_backend_sdk.return_value = test_clients.async_sdk + mock_deployments_sdk.return_value = test_clients.async_sdk - controller = ModelsController( + class ModelsControllerWithDeploymentsPlugin(ModelsController): + def step(self) -> None: + super().step() + self._loop.run_until_complete(deployments_controller.reconcile()) + + controller = ModelsControllerWithDeploymentsPlugin( backend_registry=backend_registry, stop_signal=None, ) controller._provider_reconciler.reconcile_model_providers = AsyncMock(return_value=None) controller._loop.run_until_complete(deployments_controller.on_startup()) - original_step = controller.step - - def step_with_deployments_plugin() -> None: - original_step() - controller._loop.run_until_complete(deployments_controller.reconcile()) - - controller.step = step_with_deployments_plugin - model_cache = global_model_cache() yield controller, model_cache, test_clients.sdk, mock_nim_image, docker_test_context, test_clients.async_sdk @@ -416,7 +428,7 @@ async def trigger_cache_refresh( @pytest.hookimpl(tryfirst=True, hookwrapper=True) -def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, None, None]: +def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, Any, None]: """Store test results on the item for fixture access.""" outcome = yield rep = outcome.get_result() diff --git a/services/core/inference-gateway/tests/unit/conftest.py b/services/core/inference-gateway/tests/unit/conftest.py index a81eea9334..1e397be0aa 100644 --- a/services/core/inference-gateway/tests/unit/conftest.py +++ b/services/core/inference-gateway/tests/unit/conftest.py @@ -169,9 +169,8 @@ def app_and_client( ], ) - mocker.patch("nmp.core.inference_gateway.service.get_async_platform_sdk", return_value=mock_nmp_sdk) - service = InferenceGatewayService() + mocker.patch.object(service.dependency_provider, "get_sdk_client", return_value=mock_nmp_sdk) app = service.app app.dependency_overrides[global_http_client] = lambda: mock_proxy_client app.dependency_overrides[global_model_cache] = lambda: model_cache diff --git a/services/core/inference-gateway/tests/unit/test_service.py b/services/core/inference-gateway/tests/unit/test_service.py index 5c1cf8b583..6dbba82cbd 100644 --- a/services/core/inference-gateway/tests/unit/test_service.py +++ b/services/core/inference-gateway/tests/unit/test_service.py @@ -17,7 +17,6 @@ async def test_debug_startup_hydrates_model_entity_metadata(mocker): http_client = Mock() http_client.close = AsyncMock() - mocker.patch("nmp.core.inference_gateway.service.get_async_platform_sdk", return_value=sdk) mocker.patch( "nmp.core.inference_gateway.api.middleware_registry.load_middleware_plugins", AsyncMock(return_value=MiddlewareRegistry()), @@ -38,6 +37,7 @@ async def test_debug_startup_hydrates_model_entity_metadata(mocker): ) service = InferenceGatewayService() + mocker.patch.object(service.dependency_provider, "get_sdk_client", return_value=sdk) await service.on_startup() await service.on_shutdown() diff --git a/services/core/jobs/src/nmp/core/jobs/api/dependencies.py b/services/core/jobs/src/nmp/core/jobs/api/dependencies.py index 1a1a7049df..0e89bf77a0 100644 --- a/services/core/jobs/src/nmp/core/jobs/api/dependencies.py +++ b/services/core/jobs/src/nmp/core/jobs/api/dependencies.py @@ -3,27 +3,22 @@ """FastAPI dependencies for the Jobs API.""" -from fastapi import Depends, Request +from fastapi import Depends from nemo_platform import AsyncNeMoPlatform from nmp.common.entities.client import EntityClient -from nmp.common.sdk_factory import get_async_platform_sdk -from nmp.common.service.dependencies import get_entity_client +from nmp.common.service.dependencies import get_entity_client, get_sdk_client from nmp.core.jobs.app.dispatcher import JobDispatcher -async def get_sdk_with_auth(request: Request) -> AsyncNeMoPlatform: +async def get_sdk_with_auth( + sdk: AsyncNeMoPlatform = Depends(get_sdk_client), +) -> AsyncNeMoPlatform: """Get SDK client with current request's auth headers. - This dependency creates a new SDK instance with the current user's - auth context propagated. This is needed for internal service calls - that require authorization (e.g., creating filesets). - - Args: - request: The FastAPI request object (needed to ensure we're in request context) + The platform overrides get_sdk_client with a request-scoped SDK that + preserves the service provider's HTTP transport. """ - # By taking request as parameter, we ensure this runs in request context - # where auth headers are available via context vars - return get_async_platform_sdk() + return sdk async def dep_dispatcher( diff --git a/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py b/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py index 65f23c9b31..dc9a9443e9 100644 --- a/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py +++ b/services/core/jobs/src/nmp/core/jobs/api/v2/jobs/endpoints.py @@ -31,7 +31,6 @@ PlatformJobStatusResponse, ) from nmp.common.observability import scoped_app_ctx -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service.dependencies import get_sdk_client from nmp.core.jobs.api.dependencies import dep_dispatcher from nmp.core.jobs.api.v2.jobs.schemas import ( @@ -675,7 +674,7 @@ async def download_job_result( job_name=job, workspace=workspace, artifact_url=result.artifact_url, - files_sdk=get_async_platform_sdk(), + files_sdk=dispatcher.sdk, ) background_tasks.add_task(lambda: tmp_dir_path.cleanup_tmp_dir()) return FileResponse(path=tmp_dir_path.path, filename=filename, background=background_tasks) diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/main.py b/services/core/jobs/src/nmp/core/jobs/controllers/main.py index 15343d1eae..bae647053c 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/main.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/main.py @@ -23,7 +23,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: # Create logger after configuration is set up logger = logging.getLogger(__name__) logger.info("Starting jobs controller") diff --git a/services/core/models/src/nmp/core/models/api/dependencies.py b/services/core/models/src/nmp/core/models/api/dependencies.py index b50a2ade03..eacf2da598 100644 --- a/services/core/models/src/nmp/core/models/api/dependencies.py +++ b/services/core/models/src/nmp/core/models/api/dependencies.py @@ -17,16 +17,18 @@ def get_model_entity_service( entity_client: EntityClient = Depends(get_entity_client), + nmp_sdk: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> ModelEntityService: """Dependency to get ModelEntityService instance.""" - return ModelEntityService(entity_client) + return ModelEntityService(entity_client, sdk=nmp_sdk) def get_adapter_entity_service( entity_client: EntityClient = Depends(get_entity_client), + nmp_sdk: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> AdapterEntityService: """Dependency to get AdapterEntityService instance.""" - return AdapterEntityService(entity_client) + return AdapterEntityService(entity_client, sdk=nmp_sdk) def get_model_provider_service( diff --git a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py index c67be6d441..88bf92f11a 100644 --- a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py @@ -11,7 +11,6 @@ from nmp.common.api.parsed_filter import ParsedFilter from nmp.common.entities import ALL_WORKSPACES, ListResponse from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.core.models.api.service.model_entity_service import _adapter_to_adapter_schema, get_fileset_and_files_list from nmp.core.models.constants import parse_model_ref from nmp.core.models.entities import Adapter, Model @@ -24,9 +23,9 @@ class AdapterEntityService: """Service for adapter CRUD, scoped to a workspace, with model reference from path or body.""" - def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform | None = None) -> None: + def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform) -> None: self.entity_client = entity_client - self.sdk = sdk or get_async_platform_sdk() + self.sdk = sdk async def _fetch_all_entities( self, diff --git a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py index c6e1ecb2c8..fc1d6bff03 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py @@ -17,7 +17,6 @@ from nmp.common.auth import AuthClient from nmp.common.entities import ALL_WORKSPACES, ListResponse from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.core.models.api.permissions import can_set_tool_call_plugin, check_fileset_access from nmp.core.models.config import config from nmp.core.models.entities import Adapter, Model, ModelDeploymentConfig @@ -235,9 +234,9 @@ async def validate_tool_call_plugin_allowed(auth_client: AuthClient, workspace: class ModelEntityService: """Service layer for Model Entity operations.""" - def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform | None = None): + def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform): self.entity_client = entity_client - self.sdk = sdk or get_async_platform_sdk() + self.sdk = sdk async def _fetch_all_entities( self, diff --git a/services/core/models/src/nmp/core/models/api/v2/models.py b/services/core/models/src/nmp/core/models/api/v2/models.py index 10ff5dcbca..ae9202067f 100644 --- a/services/core/models/src/nmp/core/models/api/v2/models.py +++ b/services/core/models/src/nmp/core/models/api/v2/models.py @@ -26,7 +26,6 @@ from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, get_auth_client from nmp.common.entities.client import EntityNotFoundError, EntityValidationError -from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service.dependencies import get_sdk_client from nmp.core.models.api.dependencies import get_adapter_entity_service, get_model_entity_service from nmp.core.models.api.permissions import check_fileset_access @@ -156,7 +155,7 @@ async def create_model( # add sdk job creation here for checkpoint metadata if created_model.fileset: - await start_update_model_spec_job(created_model) + await start_update_model_spec_job(created_model, nmp_sdk) return created_model @@ -264,8 +263,7 @@ async def get_model( return model_entity -async def start_update_model_spec_job(model_entity: ModelEntity): - sdk = get_async_platform_sdk(as_service="models", internal=True) +async def start_update_model_spec_job(model_entity: ModelEntity, sdk: AsyncNeMoPlatform) -> None: model_spec_task_config = ModelSpecTaskConfig(workspace=model_entity.workspace, name=model_entity.name) task_spec = PlatformJobSpec( steps=[ @@ -421,7 +419,7 @@ async def update_model( raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to update model entity") if updated_model.fileset and (updated_model.fileset != original_fileset or not updated_model.spec): - await start_update_model_spec_job(updated_model) + await start_update_model_spec_job(updated_model, nmp_sdk) return updated_model diff --git a/services/core/models/src/nmp/core/models/config.py b/services/core/models/src/nmp/core/models/config.py index f4b3319280..57f3bd2cc2 100644 --- a/services/core/models/src/nmp/core/models/config.py +++ b/services/core/models/src/nmp/core/models/config.py @@ -454,7 +454,7 @@ class ModelsConfig(create_service_config_class("models")): # type: ignore # Module-level singleton instances config = get_service_config(ModelsConfig) -backends = merge_backends( +backends: dict[BackendName, BackendConfig] = merge_backends( config.controller.backends, get_default_backends_for_runtime(get_platform_config().runtime), ) diff --git a/services/core/models/src/nmp/core/models/controllers/backends/registry.py b/services/core/models/src/nmp/core/models/controllers/backends/registry.py index 2010772345..72d2728975 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/registry.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/registry.py @@ -3,8 +3,9 @@ """Backend registry for Models Controller service.""" +from collections.abc import Mapping from logging import getLogger -from typing import Dict, Self +from typing import Any, Dict, Protocol, Self from nemo_platform import AsyncNeMoPlatform from nmp.core.models.controllers.backends.backends import ServiceBackend @@ -23,9 +24,19 @@ # Type alias for the backend name BackendName = str + +class BackendFactory(Protocol): + def __call__( + self, + nmp_sdk: AsyncNeMoPlatform, + config: dict[str, Any], + huggingface_model_puller: str, + ) -> ServiceBackend: ... + + # The deployments_plugin backend is resolved lazily because it imports the # optional `nemo_deployments_plugin` package. -backend_classes: Dict[BackendName, type[ServiceBackend]] = {} +backend_classes: Dict[BackendName, BackendFactory] = {} _LAZY_BACKEND_NAMES = frozenset({"deployments_plugin"}) @@ -37,8 +48,8 @@ def _resolve_backend_class( - name: BackendName, available_backends: Dict[BackendName, type[ServiceBackend]] -) -> type[ServiceBackend]: + name: BackendName, available_backends: Mapping[BackendName, BackendFactory] +) -> BackendFactory: """Return the backend class for ``name``, importing optional backends lazily.""" if name in available_backends: return available_backends[name] @@ -83,9 +94,9 @@ def __init__(self, registry: Dict[BackendName, ServiceBackend]) -> None: def from_config( cls, nmp_sdk: AsyncNeMoPlatform, - backend_configs: Dict[BackendName, BackendConfig], + backend_configs: Mapping[BackendName, BackendConfig], huggingface_model_puller: str, - available_backends: Dict[BackendName, type[ServiceBackend]] | None = None, + available_backends: Mapping[BackendName, BackendFactory] | None = None, ) -> Self: """Create a BackendRegistry from backend configurations. diff --git a/services/core/models/src/nmp/core/models/controllers/main.py b/services/core/models/src/nmp/core/models/controllers/main.py index 92824a9ee8..d22321778c 100644 --- a/services/core/models/src/nmp/core/models/controllers/main.py +++ b/services/core/models/src/nmp/core/models/controllers/main.py @@ -11,7 +11,7 @@ from nmp.common.service.api.health import wait_for_service_ready from nmp.core.models.config import backends from nmp.core.models.config import config as models_config -from nmp.core.models.controllers.backends.registry import BackendRegistry +from nmp.core.models.controllers.backends.registry import BackendConfig, BackendRegistry from nmp.core.models.controllers.models_controller import ModelsController stop_signal = threading.Event() @@ -35,7 +35,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: """Run the Models Controller with its control loop.""" global models_controller_monitored @@ -59,9 +59,10 @@ def run(parent_stop_signal: threading.Event | None = None): # Initialize backend registry from configuration logger.info("Initializing backend registry...") logger.debug(f"Models backend configs: {backends}") + backend_configs: dict[str, BackendConfig] = {name: config for name, config in backends.items()} backend_registry = BackendRegistry.from_config( nmp_sdk=nmp_sdk, - backend_configs=backends, + backend_configs=backend_configs, huggingface_model_puller=models_config.huggingface_model_puller, ) logger.info(f"Backend registry initialized with: {', '.join(backend_registry.list_backends())}") diff --git a/services/core/models/src/nmp/core/models/sidecars/adapters/main.py b/services/core/models/src/nmp/core/models/sidecars/adapters/main.py index 40e9ee1ce3..209bdbf988 100644 --- a/services/core/models/src/nmp/core/models/sidecars/adapters/main.py +++ b/services/core/models/src/nmp/core/models/sidecars/adapters/main.py @@ -31,7 +31,11 @@ class AdaptersController(Controller): - def __init__(self, stop_signal: threading.Event | None = None): + def __init__( + self, + *, + stop_signal: threading.Event | None = None, + ) -> None: self.nim_peft_source = os.getenv("NIM_PEFT_SOURCE", "") if not self.nim_peft_source: msg = "NIM_PEFT_SOURCE is not set on the container" @@ -142,12 +146,12 @@ def _update_prompt_tuned_models(self, dirs_to_keep: set[str]): # (AALGO-129): they remain single-workspace for now and continue to # use the bare model_entity.name as their on-disk directory. logger.info(f"Fetching prompt data for {self.workspace}/{self.model_name}") - model_entities: list[ModelEntity] = self._sdk.models.list( + model_entities = self._sdk.models.list( workspace=self.workspace, filter={ "base_model": self.model_name, }, - ) + ).data for model_entity in model_entities: if model_entity.prompt: dirs_to_keep.add(model_entity.name) @@ -448,7 +452,7 @@ def handle_sighup(signum, frame): stop_signal.set() -def run(parent_stop_signal: threading.Event | None = None): +def run(parent_stop_signal: threading.Event | None = None) -> None: """Run the Adapters Controller with its control loop.""" global adapters_controller_monitored diff --git a/services/core/models/tests/integration/conftest.py b/services/core/models/tests/integration/conftest.py index 0ac087d285..b46ff12a34 100644 --- a/services/core/models/tests/integration/conftest.py +++ b/services/core/models/tests/integration/conftest.py @@ -6,7 +6,7 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, Generator, Optional +from typing import Any, Generator, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -160,12 +160,20 @@ def shutdown(self) -> None: async def create_model_deployment(self, ctx: ModelContext) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.create_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) + deployment = ctx.model_deployment + config = ctx.model_deployment_config + assert deployment is not None + assert config is not None + self.create_calls.append((deployment, config, ctx.model_entity)) return self.create_response async def update_model_deployment(self, ctx: ModelContext) -> DeploymentStatusUpdate: """Record call and return configured response.""" - self.update_calls.append((ctx.model_deployment, ctx.model_deployment_config, ctx.model_entity)) + deployment = ctx.model_deployment + config = ctx.model_deployment_config + assert deployment is not None + assert config is not None + self.update_calls.append((deployment, config, ctx.model_entity)) return self.create_response # Update returns same as create async def get_model_deployment_status(self, ctx: ModelContext) -> DeploymentStatusUpdate: @@ -179,6 +187,7 @@ async def get_model_deployment_status(self, ctx: ModelContext) -> DeploymentStat otherwise falls back to default_status_response. """ deployment = ctx.model_deployment + assert deployment is not None self.status_calls.append(deployment) return self.status_responses.get(deployment.name, self.default_status_response) @@ -243,7 +252,7 @@ def controller_with_mock_backend( Yields: Tuple of (controller, mock_backend, sync_sdk) for testing """ - mock_backend = mock_backend_registry.get_backend() + mock_backend = cast(MockServiceBackend, mock_backend_registry.get_backend()) # Create controller with mock backend registry # We need to patch the SDK factory and platform config (used in config and main modules) @@ -445,11 +454,15 @@ def reconcile_stack(models_controller: ModelsController) -> None: return_value=mock_platform_config, ), patch("nmp.core.models.controllers.models_controller.get_async_platform_sdk") as mock_models_sdk, - patch("nemo_platform_plugin.sdk_provider.get_async_platform_sdk") as mock_sdk, + patch( + "nmp.core.models.controllers.backends.deployments_plugin.backend.get_async_platform_sdk" + ) as mock_backend_sdk, + patch("nemo_deployments_plugin.controller.get_async_platform_sdk") as mock_deployments_sdk, patch("nemo_deployments_plugin.config.DeploymentsConfig.get", return_value=deployments_config), ): mock_models_sdk.return_value = test_clients.async_sdk - mock_sdk.return_value = test_clients.async_sdk + mock_backend_sdk.return_value = test_clients.async_sdk + mock_deployments_sdk.return_value = test_clients.async_sdk models_controller = ModelsController( backend_registry=backend_registry, @@ -479,7 +492,7 @@ def reconcile_stack(models_controller: ModelsController) -> None: @pytest.hookimpl(tryfirst=True, hookwrapper=True) -def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, None, None]: +def pytest_runtest_makereport(item: pytest.Item, call: pytest.CallInfo[None]) -> Generator[None, Any, None]: """Store test results on the item for fixture access.""" outcome = yield rep = outcome.get_result() diff --git a/services/core/models/tests/unit/api/test_models_api.py b/services/core/models/tests/unit/api/test_models_api.py index 3b0a9c2a03..30ecc45077 100644 --- a/services/core/models/tests/unit/api/test_models_api.py +++ b/services/core/models/tests/unit/api/test_models_api.py @@ -411,18 +411,16 @@ def test_create_model_entity_validation_error_returns_422(client, mock_model_ent @pytest.mark.asyncio -async def test_model_spec_job_transport_failure_does_not_fail_persisted_model(sample_model_entity): +async def test_start_update_model_spec_job_swallows_nemo_transport_error(sample_model_entity): request = httpx.Request("POST", "http://test/apis/jobs/v2/workspaces/nvidia/jobs") + sdk = MagicMock() jobs = MagicMock() jobs.create_job = AsyncMock( side_effect=NemoTransportError(httpx.ConnectError("Connection refused", request=request)) ) - with ( - patch("nmp.core.models.api.v2.models.get_async_platform_sdk"), - patch("nmp.core.models.api.v2.models.client_from_platform", return_value=jobs), - ): - await start_update_model_spec_job(sample_model_entity) + with patch("nmp.core.models.api.v2.models.client_from_platform", return_value=jobs): + await start_update_model_spec_job(sample_model_entity, sdk) jobs.create_job.assert_awaited_once() diff --git a/services/core/models/tests/unit/sidecars/test_adapters_controller.py b/services/core/models/tests/unit/sidecars/test_adapters_controller.py index a1503eb9f1..3d4079181b 100644 --- a/services/core/models/tests/unit/sidecars/test_adapters_controller.py +++ b/services/core/models/tests/unit/sidecars/test_adapters_controller.py @@ -648,7 +648,7 @@ def test_step_gc_removes_stale_dir_after_adapter_workspace_change(self, controll controller._sdk.files.list.return_value = mock_files_response # No prompt-tuned models in this scenario. - controller._sdk.models.list.return_value = [] + controller._sdk.models.list.return_value = MagicMock(data=[]) controller.step() @@ -823,7 +823,7 @@ def test_step_unloads_removed_adapter_before_delete(self, controller, tmp_path): adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models with patch.object(controller, "_vllm_api_call", return_value=(200, "")) as api: controller.step() @@ -843,7 +843,7 @@ def test_step_keeps_removed_adapter_dir_when_vllm_unreachable(self, controller, adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models # vLLM unreachable: both the kept adapter's load and the stale one's unload # hit a transport error. @@ -864,7 +864,7 @@ def test_step_deletes_removed_adapter_dir_when_vllm_answers_non_200(self, contro adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": @@ -887,7 +887,7 @@ def test_step_keeps_removed_adapter_dir_when_vllm_server_error(self, controller, adapter = _make_adapter("kept-adapter", "default/fs", updated_at=None, workspace="default") self._model_entity(controller, adapter) - controller._sdk.models.list.return_value = [] # no prompt-tuned models + controller._sdk.models.list.return_value = MagicMock(data=[]) # no prompt-tuned models def _responses(route, payload): if route == "/v1/unload_lora_adapter": diff --git a/services/guardrails/src/nmp/guardrails/api/dependencies.py b/services/guardrails/src/nmp/guardrails/api/dependencies.py index cfdac1eaa2..b9f8f2ddd8 100644 --- a/services/guardrails/src/nmp/guardrails/api/dependencies.py +++ b/services/guardrails/src/nmp/guardrails/api/dependencies.py @@ -4,13 +4,13 @@ """API dependencies for the Guardrails service.""" import logging +from collections.abc import AsyncIterator from functools import lru_cache from typing import Annotated from fastapi import Depends from nemo_platform import AsyncNeMoPlatform from nmp.common.entities.client import EntityClient -from nmp.common.http_clients import shared_async_http_client from nmp.common.service.dependencies import get_entity_client from nmp.guardrails.app.services.configs.registry import ConfigRegistry from nmp.guardrails.app.services.rails.registry import RailsRegistry @@ -48,16 +48,17 @@ def get_rails_service( # Dependency for NeMo Platform -def get_nemo_platform() -> AsyncNeMoPlatform: +async def get_nemo_platform() -> AsyncIterator[AsyncNeMoPlatform]: nim_endpoint_url = settings.nim_endpoint_settings.base_url # Remove the /v1 from the end of the URL if it exists # This is necessary because the NeMo Platform API SDK expects the base URL to not have the /v1 suffix unlike OpenAI SDK if nim_endpoint_url.endswith("/v1"): nim_endpoint_url = nim_endpoint_url[: -len("/v1")] - return AsyncNeMoPlatform( - inference_base_url=nim_endpoint_url, - http_client=shared_async_http_client(), - ) + sdk = AsyncNeMoPlatform(inference_base_url=nim_endpoint_url) + try: + yield sdk + finally: + await sdk.close() RailsServiceDep = Annotated[RailsService, Depends(get_rails_service)] diff --git a/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py b/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py index 22778a0d00..d00a015e1f 100644 --- a/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py +++ b/services/hello-world/tests/integration/tasks/test_workload_workspace_get_task.py @@ -3,7 +3,7 @@ import json from types import SimpleNamespace -from typing import cast +from unittest.mock import Mock import httpx import pytest @@ -34,27 +34,20 @@ def platform_base_url(): Configuration.clear_override(PlatformConfig) -class _StubWorkspaces: - def __init__(self) -> None: - self.requested: list[str] = [] - - def retrieve(self, workspace: str) -> SimpleNamespace: - self.requested.append(workspace) - return SimpleNamespace(name=workspace) - - -class _StubSDK: - def __init__(self) -> None: - self.workspaces = _StubWorkspaces() +def _stub_sdk() -> tuple[NeMoPlatform, Mock]: + sdk = Mock(spec=NeMoPlatform) + retrieve = Mock(return_value=SimpleNamespace(name="workload-read-target")) + sdk.workspaces.retrieve = retrieve + return sdk, retrieve def test_workload_workspace_get_uses_task_sdk_factory(monkeypatch): - sdk = _StubSDK() + sdk, retrieve = _stub_sdk() sdk_factory_calls: list[str] = [] - def get_task_sdk(*, as_service: str) -> NeMoPlatform: + def get_task_sdk(*, as_service: str): sdk_factory_calls.append(as_service) - return cast(NeMoPlatform, sdk) + return sdk monkeypatch.setenv(TASK_CONFIG_ENVVAR, '{"workspace":"workload-read-target"}') monkeypatch.setattr("nmp.hello_world.tasks.workload_workspace_get.run.get_task_sdk", get_task_sdk) @@ -63,20 +56,20 @@ def get_task_sdk(*, as_service: str) -> NeMoPlatform: assert exit_code == 0 assert sdk_factory_calls == ["jobs"] - assert sdk.workspaces.requested == ["workload-read-target"] + retrieve.assert_called_once_with("workload-read-target") def test_workload_workspace_get_uses_injected_sdk_without_workload_token(monkeypatch): - sdk = _StubSDK() + sdk, retrieve = _stub_sdk() monkeypatch.setenv(TASK_CONFIG_ENVVAR, '{"workspace":"workload-read-target"}') monkeypatch.delenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, raising=False) monkeypatch.delenv("NEMO_WORKLOAD_TOKEN", raising=False) monkeypatch.delenv("NEMO_WORKLOAD_TOKEN_FILE", raising=False) - exit_code = task_run(sdk=cast(NeMoPlatform, sdk)) + exit_code = task_run(sdk=sdk) assert exit_code == 0 - assert sdk.workspaces.requested == ["workload-read-target"] + retrieve.assert_called_once_with("workload-read-target") @respx.mock diff --git a/services/studio/src/nmp/studio/service.py b/services/studio/src/nmp/studio/service.py index 7a05c52823..5a395e190b 100644 --- a/services/studio/src/nmp/studio/service.py +++ b/services/studio/src/nmp/studio/service.py @@ -9,9 +9,9 @@ from pathlib import Path from typing import ClassVar, List +import httpx from fastapi import FastAPI, Request, status from fastapi.responses import HTMLResponse -from nmp.common.http_clients import shared_async_http_client from nmp.common.service import RouterConfig, Service from nmp.studio import coding_agents from nmp.studio.config import StudioConfig @@ -53,9 +53,15 @@ class StudioService(Service[StudioConfig]): dependencies: ClassVar[list[str]] = [] - def __init__(self): + def __init__(self, telemetry_http_client: httpx.AsyncClient | None = None): """Initialize the studio service.""" super().__init__(name="studio", module_name="nmp.studio") + self._telemetry_http_client = telemetry_http_client or httpx.AsyncClient() + + async def on_shutdown(self) -> None: + """Close service-owned telemetry proxy resources.""" + await self._telemetry_http_client.aclose() + await super().on_shutdown() @property def title(self) -> str: @@ -144,7 +150,7 @@ async def _proxy_telemetry(self, request: Request, telemetry_path: str = "") -> target_url = self._build_telemetry_target_url(collector_url, telemetry_path, request.url.query) try: - upstream_response = await shared_async_http_client().request( + upstream_response = await self._telemetry_http_client.request( method=request.method, url=target_url, content=await request.body(), diff --git a/services/studio/tests/unit/test_service.py b/services/studio/tests/unit/test_service.py index ab367809ce..b03dffb9a0 100644 --- a/services/studio/tests/unit/test_service.py +++ b/services/studio/tests/unit/test_service.py @@ -27,11 +27,21 @@ class FakeTelemetryClient: def __init__(self, response: FakeTelemetryResponse | None = None): self.response = response or FakeTelemetryResponse() self.calls: list[dict] = [] + self.close_calls = 0 + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc_value, traceback) -> None: + pass async def request(self, **kwargs): self.calls.append(kwargs) return self.response + async def aclose(self) -> None: + self.close_calls += 1 + class TestStudioService: """Tests for the StudioService class.""" @@ -70,24 +80,21 @@ class TestTelemetryProxy: def _client( self, config: StudioConfig, - monkeypatch: pytest.MonkeyPatch, fake_client: FakeTelemetryClient | None = None, ) -> tuple[TestClient, FakeTelemetryClient]: app = FastAPI() telemetry_client = fake_client or FakeTelemetryClient() - monkeypatch.setattr("nmp.studio.service.shared_async_http_client", lambda: telemetry_client) - StudioService().with_config(config).configure_app(app) + StudioService(telemetry_http_client=telemetry_client).with_config(config).configure_app(app) return TestClient(app), telemetry_client - def test_post_strips_studio_telemetry_prefix_and_proxies_request(self, monkeypatch: pytest.MonkeyPatch): + def test_post_strips_studio_telemetry_prefix_and_proxies_request(self): """Test that /studio/telemetry/* proxies to the collector without the route prefix.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post( @@ -108,15 +115,14 @@ def test_post_strips_studio_telemetry_prefix_and_proxies_request(self, monkeypat assert call["headers"]["X-Real-IP"] == "testclient" assert call["headers"]["X-Forwarded-For"] == "testclient" - def test_post_only_forwards_whitelisted_telemetry_headers(self, monkeypatch: pytest.MonkeyPatch): + def test_post_only_forwards_whitelisted_telemetry_headers(self): """Test that browser credentials and metadata are not forwarded to the collector.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post( @@ -146,15 +152,14 @@ def test_post_only_forwards_whitelisted_telemetry_headers(self, monkeypatch: pyt "X-Forwarded-For": "testclient", } - def test_post_strips_root_telemetry_prefix_and_proxies_request(self, monkeypatch: pytest.MonkeyPatch): + def test_post_strips_root_telemetry_prefix_and_proxies_request(self): """Test that /telemetry/* keeps parity with the old nginx route.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.post("/telemetry/v1/logs", headers={"origin": origin}) @@ -162,15 +167,14 @@ def test_post_strips_root_telemetry_prefix_and_proxies_request(self, monkeypatch assert response.status_code == 200 assert telemetry_client.calls[0]["url"] == "http://collector:4318/v1/logs" - def test_options_returns_preflight_response_without_proxying(self, monkeypatch: pytest.MonkeyPatch): + def test_options_returns_preflight_response_without_proxying(self): """Test that CORS preflight requests are handled locally.""" origin = "http://studio.test" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, - ), - monkeypatch, + ) ) response = client.options("/studio/telemetry/v1/traces", headers={"origin": origin}) @@ -183,11 +187,10 @@ def test_options_returns_preflight_response_without_proxying(self, monkeypatch: assert response.headers["access-control-max-age"] == "1728000" assert telemetry_client.calls == [] - def test_disabled_telemetry_returns_404(self, monkeypatch: pytest.MonkeyPatch): + def test_disabled_telemetry_returns_404(self): """Test that disabled telemetry preserves the old nginx 404 behavior.""" client, telemetry_client = self._client( - StudioConfig(telemetry_enabled=False, otel={"collector_url": "http://collector:4318"}), - monkeypatch, + StudioConfig(telemetry_enabled=False, otel={"collector_url": "http://collector:4318"}) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://testserver"}) @@ -195,14 +198,13 @@ def test_disabled_telemetry_returns_404(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 404 assert telemetry_client.calls == [] - def test_disallowed_origin_returns_403(self, monkeypatch: pytest.MonkeyPatch): + def test_disallowed_origin_returns_403(self): """Test that disallowed origins preserve the old nginx 403 behavior.""" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": ["http://studio.test"]}, - ), - monkeypatch, + ) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://not-allowed.test"}) @@ -210,14 +212,13 @@ def test_disallowed_origin_returns_403(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 403 assert telemetry_client.calls == [] - def test_same_origin_request_is_allowed(self, monkeypatch: pytest.MonkeyPatch): + def test_same_origin_request_is_allowed(self): """Test that same-origin Studio deployments work without hard-coded host config.""" client, telemetry_client = self._client( StudioConfig( telemetry_enabled=True, otel={"collector_url": "http://collector:4318", "allowed_origins": []}, - ), - monkeypatch, + ) ) response = client.post("/studio/telemetry/v1/traces", headers={"origin": "http://testserver"}) @@ -225,6 +226,47 @@ def test_same_origin_request_is_allowed(self, monkeypatch: pytest.MonkeyPatch): assert response.status_code == 200 assert telemetry_client.calls[0]["url"] == "http://collector:4318/v1/traces" + def test_post_reuses_service_owned_telemetry_client(self, monkeypatch: pytest.MonkeyPatch): + """Test that proxied telemetry requests reuse one service-scoped client.""" + origin = "http://studio.test" + app = FastAPI() + created_clients: list[FakeTelemetryClient] = [] + + def create_client() -> FakeTelemetryClient: + client = FakeTelemetryClient() + created_clients.append(client) + return client + + monkeypatch.setattr("nmp.studio.service.httpx.AsyncClient", create_client) + StudioService().with_config( + StudioConfig( + telemetry_enabled=True, + otel={"collector_url": "http://collector:4318", "allowed_origins": [origin]}, + ) + ).configure_app(app) + client = TestClient(app) + + first_response = client.post("/studio/telemetry/v1/traces", headers={"origin": origin}) + second_response = client.post("/studio/telemetry/v1/logs", headers={"origin": origin}) + + assert first_response.status_code == 200 + assert second_response.status_code == 200 + assert len(created_clients) == 1 + assert [call["url"] for call in created_clients[0].calls] == [ + "http://collector:4318/v1/traces", + "http://collector:4318/v1/logs", + ] + + @pytest.mark.asyncio + async def test_shutdown_closes_telemetry_client(self): + """Test that the service closes its telemetry HTTP client during shutdown.""" + telemetry_client = FakeTelemetryClient() + service = StudioService(telemetry_http_client=telemetry_client) + + await service.on_shutdown() + + assert telemetry_client.close_calls == 1 + class TestStaticFilesPath: """Tests for static_files_path configuration."""