Skip to content
Closed
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
Expand Up @@ -17,6 +17,7 @@
DefaultAsyncHttpxClient,
DefaultHttpxClient,
NotGiven,
Omit,
__version__,
not_given,
)
Expand Down Expand Up @@ -74,7 +75,7 @@ def __init__(
access_token: str | None = None,
timeout: float | Timeout | None | NotGiven = not_given,
max_retries: int = DEFAULT_MAX_RETRIES,
default_headers: Mapping[str, str] | None = None,
default_headers: Mapping[str, str | Omit] | None = None,
default_query: Mapping[str, object] | None = None,
# Configure a custom httpx client.
# We provide a `DefaultHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`.
Expand Down Expand Up @@ -232,8 +233,8 @@ def copy(
timeout: float | Timeout | None | NotGiven = not_given,
http_client: httpx.Client | None = None,
max_retries: int | NotGiven = not_given,
default_headers: Mapping[str, str] | None = None,
set_default_headers: Mapping[str, str] | None = None,
default_headers: Mapping[str, str | Omit] | None = None,
set_default_headers: Mapping[str, str | Omit] | None = None,
default_query: Mapping[str, object] | None = None,
set_default_query: Mapping[str, object] | None = None,
_extra_kwargs: Mapping[str, Any] = {},
Expand Down Expand Up @@ -296,7 +297,7 @@ def __init__(
access_token: str | None = None,
timeout: float | Timeout | None | NotGiven = not_given,
max_retries: int = DEFAULT_MAX_RETRIES,
default_headers: Mapping[str, str] | None = None,
default_headers: Mapping[str, str | Omit] | None = None,
default_query: Mapping[str, object] | None = None,
# Configure a custom httpx client.
# We provide a `DefaultAsyncHttpxClient` class that you can pass to retain the default values we use for `limits`, `timeout` & `follow_redirects`.
Expand Down Expand Up @@ -476,8 +477,8 @@ def copy(
timeout: float | Timeout | None | NotGiven = not_given,
http_client: httpx.AsyncClient | None = None,
max_retries: int | NotGiven = not_given,
default_headers: Mapping[str, str] | None = None,
set_default_headers: Mapping[str, str] | None = None,
default_headers: Mapping[str, str | Omit] | None = None,
set_default_headers: Mapping[str, str | Omit] | None = None,
default_query: Mapping[str, object] | None = None,
set_default_query: Mapping[str, object] | None = None,
_extra_kwargs: Mapping[str, Any] = {},
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -53,17 +53,19 @@
import logging
import os
import threading
from collections.abc import Awaitable, Callable, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, Literal, Mapping, Protocol
from typing import Any, Literal, Protocol

import httpx
from nemo_platform import (
DefaultAsyncHttpxClient,
DefaultHttpxClient,
NeMoPlatform,
NotGiven,
Omit,
not_given,
)
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
Expand Down Expand Up @@ -97,6 +99,12 @@ def get_access_token(self) -> str: ...
async def get_access_token_async(self) -> str: ...


_SyncRequestHook = Callable[[httpx.Request], None]
_AsyncRequestHook = Callable[[httpx.Request], Awaitable[None]]
_SyncHttpClientFactory = Callable[[_SyncRequestHook, str | Literal[True]], httpx.Client]
_AsyncHttpClientFactory = Callable[[_AsyncRequestHook, str | Literal[True]], httpx.AsyncClient]


@dataclass(frozen=True)
class ClientInitConfig:
"""Everything the SDK client constructor needs after config resolution.
Expand All @@ -108,7 +116,7 @@ class ClientInitConfig:

base_url: str
workspace: str | None
default_headers: Mapping[str, str] | None = None
default_headers: Mapping[str, str | Omit] | None = None
http_client: httpx.Client | httpx.AsyncClient | None = None
client_verify: str | Literal[True] = True

Expand All @@ -119,7 +127,7 @@ class _ResolvedBootstrap:

base_url: str
workspace: str | None
default_headers: dict[str, str]
default_headers: dict[str, str | Omit]
token_provider: _AccessTokenProvider | None # None for non-OAuth users
client_verify: str | Literal[True]
certificate_authority: str | None = None
Expand Down Expand Up @@ -271,7 +279,7 @@ async def get_access_token_async(self) -> str:
# ---------------------------------------------------------------------------


def _make_auth_event_hook(provider: _AccessTokenProvider):
def _make_auth_event_hook(provider: _AccessTokenProvider) -> _SyncRequestHook:
"""Create a **sync** httpx request event hook that injects the Bearer token.

Called before every SDK HTTP request. ``provider.get_access_token()``
Expand All @@ -286,7 +294,7 @@ def inject_auth(request: httpx.Request) -> None:
return inject_auth


def _make_async_auth_event_hook(provider: _AccessTokenProvider):
def _make_async_auth_event_hook(provider: _AccessTokenProvider) -> _AsyncRequestHook:
"""Create an **async** httpx request event hook for AsyncNeMoPlatform.

The actual refresh still runs in a worker thread (via
Expand All @@ -300,7 +308,10 @@ async def inject_auth(request: httpx.Request) -> None:
return inject_auth


def _headers_with_seeded_auth(headers: Mapping[str, str], provider: _AccessTokenProvider) -> dict[str, str]:
def _headers_with_seeded_auth(
headers: Mapping[str, str | Omit],
provider: _AccessTokenProvider,
) -> dict[str, str | Omit]:
seeded_headers = dict(headers)
if isinstance(provider, _LazyWorkloadTokenExchangeProvider):
token = provider.get_cached_access_token()
Expand Down Expand Up @@ -502,7 +513,7 @@ def _resolve_bootstrap(
base_url: str | httpx.URL | None,
context_name: str | None,
access_token: str | None,
extra_headers: Mapping[str, str] | None,
extra_headers: Mapping[str, str | Omit] | None,
) -> _ResolvedBootstrap:
"""Resolve the full client bootstrap: config, OIDC discovery, token provider.

Expand All @@ -529,7 +540,7 @@ def _resolve_bootstrap(
base_url = str(resolved.cluster.base_url)
certificate_authority = resolved.cluster.certificate_authority
client_verify = client_verify_from_env(certificate_authority)
headers: dict[str, str] = dict(extra_headers) if extra_headers else {}
headers: dict[str, str | Omit] = dict(extra_headers) if extra_headers else {}

workload_identity_token_file = _workload_identity_token_file_from_env()
if workload_identity_token_file is not None and access_token is None and not os.environ.get("NMP_ACCESS_TOKEN"):
Expand Down Expand Up @@ -629,7 +640,8 @@ def build_client_init_kwargs(
base_url: str | httpx.URL | None = None,
context_name: str | None = None,
access_token: str | None = None,
extra_headers: Mapping[str, str] | None = None,
extra_headers: Mapping[str, str | Omit] | None = None,
http_client_factory: _SyncHttpClientFactory | None = None,
) -> ClientInitConfig:
"""Build constructor kwargs for a **sync** NeMoPlatform client.

Expand Down Expand Up @@ -659,10 +671,14 @@ def build_client_init_kwargs(
# The event hook will overwrite it with a fresh token on each request.
headers = _headers_with_seeded_auth(bootstrap.default_headers, bootstrap.token_provider)
hook = _make_auth_event_hook(bootstrap.token_provider)
http_client = DefaultHttpxClient(
event_hooks={"request": [hook], "response": []},
follow_redirects=True,
verify=bootstrap.client_verify,
http_client = (
http_client_factory(hook, bootstrap.client_verify)
if http_client_factory is not None
else DefaultHttpxClient(
event_hooks={"request": [hook], "response": []},
follow_redirects=True,
verify=bootstrap.client_verify,
)
)
return ClientInitConfig(
base_url=bootstrap.base_url,
Expand All @@ -679,7 +695,8 @@ def build_async_client_init_kwargs(
base_url: str | httpx.URL | None = None,
context_name: str | None = None,
access_token: str | None = None,
extra_headers: Mapping[str, str] | None = None,
extra_headers: Mapping[str, str | Omit] | None = None,
http_client_factory: _AsyncHttpClientFactory | None = None,
) -> ClientInitConfig:
"""Build constructor kwargs for an **async** AsyncNeMoPlatform client.

Expand All @@ -704,10 +721,14 @@ def build_async_client_init_kwargs(

headers = _headers_with_seeded_auth(bootstrap.default_headers, bootstrap.token_provider)
hook = _make_async_auth_event_hook(bootstrap.token_provider)
http_client = DefaultAsyncHttpxClient(
event_hooks={"request": [hook], "response": []},
follow_redirects=True,
verify=bootstrap.client_verify,
http_client = (
http_client_factory(hook, bootstrap.client_verify)
if http_client_factory is not None
else DefaultAsyncHttpxClient(
event_hooks={"request": [hook], "response": []},
follow_redirects=True,
verify=bootstrap.client_verify,
)
)
return ClientInitConfig(
base_url=bootstrap.base_url,
Expand All @@ -726,7 +747,7 @@ def create_client(
access_token: str | None = None,
timeout: float | httpx.Timeout | None | NotGiven = not_given,
max_retries: int = 2,
extra_headers: Mapping[str, str] | None = None,
extra_headers: Mapping[str, str | Omit] | None = None,
) -> NeMoPlatform:
"""Create a NeMoPlatform client from the nmp config.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -281,8 +281,12 @@ def _build_headers(
headers["X-NMP-Principal-Email"] = principal["email"]
if principal.get("groups"):
headers["X-NMP-Principal-Groups"] = ",".join(principal["groups"])
if principal.get("on_behalf_of"):
headers.update(_on_behalf_of_headers(principal))

if on_behalf_of is not None:
headers.pop("X-NMP-Principal-On-Behalf-Of-Email", None)
headers.pop("X-NMP-Principal-On-Behalf-Of-Groups", None)
headers["X-NMP-Principal-On-Behalf-Of"] = on_behalf_of

return headers
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,10 +23,10 @@
FileResultSerializer,
JobRouteOption,
PlatformJobResultRoute,
PlatformJobSpec,
job_route_factory,
)
from nemo_platform_plugin.jobs.routes import _rebase_job_collection_routes
from nemo_platform_plugin.jobs.spec import PlatformJobSpec
from nemo_platform_plugin.service import NemoService, RouterSpec
from pydantic import BaseModel

Expand Down
45 changes: 44 additions & 1 deletion packages/nemo_platform_plugin/tests/test_sdk_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,31 @@ def test_get_platform_sdk_on_behalf_of(self, monkeypatch):

assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "user@ex.com"

def test_get_platform_sdk_propagates_env_principal_on_behalf_of(self, monkeypatch):
monkeypatch.setenv("NMP_BASE_URL", "http://test:9090")
monkeypatch.setenv(
"NMP_PRINCIPAL",
json.dumps(
{
"id": "service:evaluator",
"email": "evaluator@service.test",
"groups": ["system:serviceaccounts"],
"on_behalf_of": "creator@ex.com",
"on_behalf_of_email": "creator@ex.com",
"on_behalf_of_groups": ["workspace-editors", "ml-team"],
}
),
)

sdk = DefaultSDKProvider().get_platform_sdk()

assert sdk.default_headers["X-NMP-Principal-Id"] == "service:evaluator"
assert sdk.default_headers["X-NMP-Principal-Email"] == "evaluator@service.test"
assert sdk.default_headers["X-NMP-Principal-Groups"] == "system:serviceaccounts"
assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of"] == "creator@ex.com"
assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of-Email"] == "creator@ex.com"
assert sdk.default_headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "workspace-editors,ml-team"


# ---------------------------------------------------------------------------
# get_async_task_sdk — the async sibling of get_task_sdk
Expand Down Expand Up @@ -246,9 +271,27 @@ class _CustomProvider:
def get_task_sdk(self, service_name: str) -> NeMoPlatform:
return NeMoPlatform(base_url="http://custom:1234")

def get_platform_sdk(self, **kwargs) -> NeMoPlatform:
def get_async_task_sdk(self, service_name: str) -> AsyncNeMoPlatform:
return AsyncNeMoPlatform(base_url="http://custom:1234")

def get_platform_sdk(
self,
*,
as_service: str | None = None,
internal: bool = False,
on_behalf_of: str | None = None,
) -> NeMoPlatform:
return NeMoPlatform(base_url="http://custom:1234")

def get_async_platform_sdk(
self,
*,
as_service: str | None = None,
internal: bool = False,
on_behalf_of: str | None = None,
) -> AsyncNeMoPlatform:
return AsyncNeMoPlatform(base_url="http://custom:1234")


class _FakeEntryPoint:
def __init__(
Expand Down
2 changes: 1 addition & 1 deletion packages/nmp_common/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ dependencies = [
"base58>=2.1.1",
"pydantic>=2.10.3",
"pydantic-settings>=2.8.1",
"fastapi[standard]>=0.115.4",
"fastapi[standard]>=0.137.0",
"lark>=1.1.0",
"nemo-platform-plugin",
"pyyaml>=6.0.2",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -49,14 +49,11 @@ def __init__(self, config: AuthConfig, http_client: httpx.AsyncClient | None = N

def _get_sdk(self) -> AsyncNeMoPlatform:
if self._sdk is None:
# Import lazily to avoid an auth -> SDK factory import cycle. The
# factory attaches PlatformRequestRouter, which owns service
# discovery and transport selection for /apis/auth requests.
from nmp.common.sdk_factory import get_async_platform_sdk, with_options_preserving_request_router
# Import lazily to avoid an auth -> SDK factory import cycle.
from nmp.common.sdk_factory import get_async_platform_sdk

sdk = get_async_platform_sdk(http_client=self._http_client)
self._sdk = with_options_preserving_request_router(
sdk,
self._sdk = sdk.with_options(
max_retries=0,
_extra_kwargs={"_strict_response_validation": True},
)
Expand Down
Loading
Loading