Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 18 additions & 14 deletions packages/nemo_platform_ext/tests/local/test_health_child.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down Expand Up @@ -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")

Expand All @@ -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=[])
Expand All @@ -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()


# ---------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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: {})
Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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: {})
Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
-----------------------------
Expand Down Expand Up @@ -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.

Expand All @@ -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(
Expand All @@ -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(
Expand Down Expand Up @@ -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)

Expand Down
45 changes: 4 additions & 41 deletions packages/nemo_platform_plugin/tests/client/test_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Loading
Loading