Skip to content
Merged
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
Original file line number Diff line number Diff line change
@@ -1,6 +1,12 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Shared client constants."""
"""Shared client constants and env checks."""

import os

WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE"


def is_workload_identity_token_file_set() -> bool:
return bool(os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR))
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
from typing import Any, Protocol, TypeVar, runtime_checkable

from nemo_platform import AsyncNeMoPlatform, NeMoPlatform
from nemo_platform_plugin.client.constants import is_workload_identity_token_file_set

_SDKT = TypeVar("_SDKT", NeMoPlatform, AsyncNeMoPlatform)

Expand Down Expand Up @@ -157,6 +158,10 @@ def _on_behalf_of_headers(principal: dict[str, Any]) -> dict[str, str]:
return headers


def _workload_identity_headers(*, internal: bool) -> dict[str, str]:
return {_INTERNAL_REQUEST_HEADER: "true"} if internal else {}


class DefaultSDKProvider:
"""Env-var-based provider that ships with the plugin package.

Expand All @@ -166,6 +171,12 @@ 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",
Expand All @@ -189,6 +200,12 @@ def get_task_sdk(self, service_name: str) -> NeMoPlatform:
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",
Expand Down Expand Up @@ -217,6 +234,10 @@ def _make_sdk(
internal: bool = False,
on_behalf_of: str | None = None,
) -> _SDKT:
if as_service is None and on_behalf_of is None and 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)
return cls(base_url=self._base_url(), default_headers=headers or None)

Expand Down
39 changes: 39 additions & 0 deletions packages/nemo_platform_plugin/tests/test_sdk_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

import pytest
from nemo_platform import AsyncNeMoPlatform, NeMoPlatform
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
from nemo_platform_plugin.sdk_provider import (
DefaultSDKProvider,
SDKProvider,
Expand Down Expand Up @@ -127,6 +128,25 @@ 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):
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_task_sdk("evaluator")
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_as_service(self, monkeypatch):
monkeypatch.setenv("NMP_BASE_URL", "http://test:9090")
monkeypatch.delenv("NMP_PRINCIPAL", raising=False)
Expand Down Expand Up @@ -197,6 +217,25 @@ def test_parity_with_sync_task_sdk(self, monkeypatch):
provider = DefaultSDKProvider()
assert _xnmp(provider.get_async_task_sdk("evaluator")) == _xnmp(provider.get_task_sdk("evaluator"))

@pytest.mark.asyncio
async def test_async_task_sdk_uses_workload_identity_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"]}),
)

sdk = DefaultSDKProvider().get_async_task_sdk("evaluator")
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:
await sdk.close()


# ---------------------------------------------------------------------------
# Provider resolution
Expand Down
15 changes: 15 additions & 0 deletions packages/nmp_common/src/nmp/common/platform_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@

import httpx
from httpx._types import TimeoutTypes
from nemo_platform import DefaultAsyncHttpxClient, DefaultHttpxClient
from nmp.common.config import PlatformConfig

UDS_BASE_URL = "http://nemo-platform.local"
Expand Down Expand Up @@ -47,6 +48,20 @@ 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:
if self.transport == "uds":
return self.sync_http_client(timeout=timeout)
if timeout is None:
return DefaultHttpxClient()
return DefaultHttpxClient(timeout=timeout)

def async_sdk_http_client(self, *, timeout: TimeoutTypes | None = None) -> httpx.AsyncClient:
if self.transport == "uds":
return self.async_http_client(timeout=timeout)
if timeout is None:
return DefaultAsyncHttpxClient()
return DefaultAsyncHttpxClient(timeout=timeout)


def resolve_platform_endpoint(platform_config: PlatformConfig | None = None) -> PlatformEndpoint:
"""Resolve the default platform endpoint from ``NMP_BASE_URL`` / config."""
Expand Down
70 changes: 68 additions & 2 deletions packages/nmp_common/src/nmp/common/sdk_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

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
Expand Down Expand Up @@ -156,6 +157,26 @@ def _async_http_client_for_endpoint(
return shared_async_http_client()


def _should_bootstrap_workload_identity(
*,
as_service: str | None,
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()
)


def _workload_identity_extra_headers(*, internal: bool) -> dict[str, str]:
return MARK_INTERNAL_REQUEST_HEADERS.copy() if internal else {}


def _get_default_headers(
as_service: str | None = None, internal: bool = False, on_behalf_of: str | Principal | None = None
) -> dict[str, str]:
Expand Down Expand Up @@ -241,8 +262,21 @@ def get_platform_sdk(
Returns:
Configured NeMoPlatform SDK instance.
"""
headers = _get_default_headers(as_service, internal, on_behalf_of)
endpoint = resolve_platform_endpoint()
if _should_bootstrap_workload_identity(
as_service=as_service,
on_behalf_of=on_behalf_of,
http_client=http_client,
endpoint=endpoint,
):
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)

Comment thread
ironcommit marked this conversation as resolved.
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),
Expand All @@ -265,6 +299,12 @@ 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(
Expand Down Expand Up @@ -293,6 +333,12 @@ 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(
Expand Down Expand Up @@ -331,8 +377,28 @@ def get_async_platform_sdk(
Returns:
Configured AsyncNeMoPlatform SDK instance.
"""
headers = _get_default_headers(as_service, internal, on_behalf_of)
endpoint = resolve_platform_endpoint()
if _should_bootstrap_workload_identity(
as_service=as_service,
on_behalf_of=on_behalf_of,
http_client=http_client,
endpoint=endpoint,
):
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.
Expand Down
67 changes: 61 additions & 6 deletions packages/nmp_common/tests/auth/test_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,14 @@ async def discovery_endpoint():
async def hf_download_endpoint(workspace: str, name: str, revision: str, path: str):
return {"workspace": workspace, "name": name, "path": path}

@app.post("/apis/files/v2/workspaces/{workspace}/filesets/{name}/otlp/v1/logs")
async def files_otlp_logs_upload(workspace: str, name: str):
return {"workspace": workspace, "name": name}

@app.post("/apis/files/v2/workspaces/{workspace}/filesets/{name}/otlp/v1/logs/query")
async def files_otlp_logs_query(workspace: str, name: str):
return {"workspace": workspace, "name": name}

Configuration.set_override(auth_config)
app.add_middleware(AuthorizationMiddleware, service_name="test-service")

Expand Down Expand Up @@ -426,12 +434,8 @@ def test_service_principal_uses_pdp(self, auth_config_enabled):
mock_authorize.assert_called_once()


class TestHfEndpointAuth:
"""Tests for HuggingFace-compatible endpoint authentication.

HF endpoints accept service principal tokens via Bearer header (HF_TOKEN),
allowing huggingface-hub clients to authenticate with HF_TOKEN=service:<name>.
"""
class TestCompatibilityAuth:
"""Tests for compatibility auth paths used by non-standard clients."""

@pytest.mark.parametrize(
"token,expected_status",
Expand Down Expand Up @@ -475,6 +479,57 @@ def test_hf_endpoint_authorizes_as_bearer_service_principal(self, auth_config_en
auth_client = mock_authorize.call_args.args[0]
assert auth_client.principal.id == "service:models"

def test_files_otlp_logs_upload_accepts_service_principal_header(self, auth_config_enabled):
"""Job log uploads accept the launcher fallback service principal header."""
app = create_test_app(auth_config_enabled)
client = TestClient(app, raise_server_exceptions=False)

with patch.object(AuthClient, "authorize_request", autospec=True) as mock_authorize:
mock_authorize.return_value = MagicMock(allowed=True)

response = client.post(
"/apis/files/v2/workspaces/my-workspace/filesets/job-fileset-1/otlp/v1/logs",
headers={"X-NMP-Principal-Id": "service:jobs"},
)

assert response.status_code == 200
auth_client = mock_authorize.call_args.args[0]
assert auth_client.principal.id == "service:jobs"

def test_files_otlp_logs_upload_does_not_accept_service_bearer_token(self, auth_config_enabled):
"""Job log uploads use X-NMP service headers, not raw service bearer tokens."""
app = create_test_app(auth_config_enabled)
client = TestClient(app, raise_server_exceptions=False)

with patch("nmp.common.auth.jwt.JWTValidator.validate_token") as mock_validate:
mock_validate.return_value = None
with patch("nmp.common.auth.client.AuthClient.authorize_request") as mock_authorize:
response = client.post(
"/apis/files/v2/workspaces/my-workspace/filesets/job-fileset-1/otlp/v1/logs",
headers={"Authorization": "Bearer service:jobs"},
)

assert response.status_code == 401
mock_validate.assert_called_once()
mock_authorize.assert_not_called()

def test_regular_endpoint_does_not_accept_service_bearer_token(self, auth_config_enabled):
"""Raw service bearer tokens are not accepted on arbitrary routes."""
app = create_test_app(auth_config_enabled)
client = TestClient(app, raise_server_exceptions=False)

with patch("nmp.common.auth.jwt.JWTValidator.validate_token") as mock_validate:
mock_validate.return_value = None
with patch("nmp.common.auth.client.AuthClient.authorize_request") as mock_authorize:
response = client.get(
"/test",
headers={"Authorization": "Bearer service:jobs"},
)

assert response.status_code == 401
mock_validate.assert_called_once()
mock_authorize.assert_not_called()


class TestInternalServiceOnlyRoutes:
"""IAM role-bindings and nested Entities APIs require service principals."""
Expand Down
Loading