Skip to content

Commit 0370bfc

Browse files
committed
update
Signed-off-by: Ryan S <267728323+ironcommit@users.noreply.github.com>
1 parent 3744e08 commit 0370bfc

9 files changed

Lines changed: 171 additions & 37 deletions

File tree

openapi/ga/individual/platform.openapi.yaml

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

openapi/ga/openapi.yaml

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

openapi/openapi.yaml

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py

Lines changed: 70 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@
3838
StaticToken,
3939
TokenProvider,
4040
)
41+
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
4142
from nemo_platform_plugin.client.errors import (
4243
NemoResponseValidationError,
4344
NemoTransportError,
@@ -72,6 +73,59 @@
7273
logger = logging.getLogger(__name__)
7374

7475
DEFAULT_TIMEOUT = 60.0
76+
_AUTHORIZATION_HEADER = "Authorization"
77+
_PRINCIPAL_ID_HEADER = "X-NMP-Principal-Id"
78+
79+
80+
def _has_header(headers: Mapping[str, str] | None, name: str) -> bool:
81+
if not headers:
82+
return False
83+
normalized = name.lower()
84+
return any(header.lower() == normalized for header in headers)
85+
86+
87+
@overload
88+
def _resolve_implicit_workload_auth(
89+
*,
90+
base_url: str,
91+
auth: TokenProvider | str | None,
92+
default_headers: Mapping[str, str] | None,
93+
allow_env_bootstrap: bool,
94+
) -> TokenProvider | str | None: ...
95+
96+
97+
@overload
98+
def _resolve_implicit_workload_auth(
99+
*,
100+
base_url: str,
101+
auth: TokenProvider | AsyncTokenProvider | str | None,
102+
default_headers: Mapping[str, str] | None,
103+
allow_env_bootstrap: bool,
104+
) -> TokenProvider | AsyncTokenProvider | str | None: ...
105+
106+
107+
def _resolve_implicit_workload_auth(
108+
*,
109+
base_url: str,
110+
auth: TokenProvider | AsyncTokenProvider | str | None,
111+
default_headers: Mapping[str, str] | None,
112+
allow_env_bootstrap: bool,
113+
) -> TokenProvider | AsyncTokenProvider | str | None:
114+
if auth is not None or not allow_env_bootstrap:
115+
return auth
116+
if _has_header(default_headers, _AUTHORIZATION_HEADER) or _has_header(default_headers, _PRINCIPAL_ID_HEADER):
117+
return None
118+
119+
subject_token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)
120+
if not subject_token_file:
121+
return None
122+
123+
from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider
124+
125+
return resolve_workload_exchange_provider(
126+
base_url=base_url,
127+
subject_token_file=Path(subject_token_file),
128+
)
75129

76130

77131
@cache
@@ -420,6 +474,12 @@ def __init__(
420474
defers to the transport's timeout, giving one we build ourselves
421475
:data:`DEFAULT_TIMEOUT`; ``httpx.Timeout(None)`` waits indefinitely.
422476
"""
477+
auth = _resolve_implicit_workload_auth(
478+
base_url=base_url,
479+
auth=auth,
480+
default_headers=default_headers,
481+
allow_env_bootstrap=http_client is None,
482+
)
423483
super().__init__(
424484
base_url=base_url,
425485
workspace=workspace,
@@ -671,6 +731,12 @@ def __init__(
671731
url_resolver: Callable[[str], str | httpx.URL] | None = None,
672732
) -> None:
673733
"""Create a client. See :meth:`NemoClient.__init__` for *timeout*."""
734+
auth = _resolve_implicit_workload_auth(
735+
base_url=base_url,
736+
auth=auth,
737+
default_headers=default_headers,
738+
allow_env_bootstrap=http_client is None,
739+
)
674740
super().__init__(
675741
base_url=base_url,
676742
workspace=workspace,
@@ -911,8 +977,7 @@ def _client_from_config(
911977
"""Shared implementation for NemoClient.from_config / AsyncNemoClient.from_config."""
912978
from nemo_platform_plugin.client.config.config import Config
913979
from nemo_platform_plugin.client.config.models import ConfigParams, OAuthUser
914-
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
915-
from nemo_platform_plugin.client.oidc_factory import resolve_oidc_provider, resolve_workload_exchange_provider
980+
from nemo_platform_plugin.client.oidc_factory import resolve_oidc_provider
916981

917982
resolved_path = Path(config_path) if isinstance(config_path, str) else config_path
918983
overrides: ConfigParams | None = None
@@ -928,13 +993,9 @@ def _client_from_config(
928993

929994
auth: TokenProvider | str | None = None
930995
workload_identity_token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)
996+
use_implicit_workload_auth = bool(workload_identity_token_file and not explicit_access_token)
931997

932-
if workload_identity_token_file and not explicit_access_token:
933-
auth = resolve_workload_exchange_provider(
934-
base_url=str(ctx.cluster.base_url),
935-
subject_token_file=Path(workload_identity_token_file),
936-
)
937-
elif isinstance(ctx.user, OAuthUser):
998+
if not use_implicit_workload_auth and isinstance(ctx.user, OAuthUser):
938999
auth = resolve_oidc_provider(
9391000
base_url=str(ctx.cluster.base_url),
9401001
context_name=ctx.context_name,
@@ -944,7 +1005,7 @@ def _client_from_config(
9441005
config_path=actual_config_path,
9451006
explicit_access_token=explicit_access_token,
9461007
)
947-
elif ctx.user:
1008+
elif not use_implicit_workload_auth and ctx.user:
9481009
client_config = ctx.user.get_client_config()
9491010
raw_headers = client_config.get("default_headers")
9501011
if isinstance(raw_headers, dict):

packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py

Lines changed: 2 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -16,13 +16,9 @@
1616
import json
1717
import logging
1818
import os
19-
from pathlib import Path
2019
from typing import Any
2120

2221
from nemo_platform_plugin.client.client import AsyncNemoClient, NemoClient
23-
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
24-
from nemo_platform_plugin.client.oidc import WorkloadTokenExchangeProvider
25-
from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider
2622

2723
logger = logging.getLogger(__name__)
2824

@@ -82,23 +78,6 @@ def _base_url() -> str:
8278
return os.environ.get("NMP_BASE_URL", "http://localhost:8080")
8379

8480

85-
def _workload_exchange_auth_from_env(
86-
base_url: str,
87-
headers: dict[str, str],
88-
) -> WorkloadTokenExchangeProvider | None:
89-
if headers.get("X-NMP-Principal-Id"):
90-
return None
91-
92-
subject_token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)
93-
if not subject_token_file:
94-
return None
95-
96-
return resolve_workload_exchange_provider(
97-
base_url=base_url,
98-
subject_token_file=Path(subject_token_file),
99-
)
100-
101-
10281
def get_nemo_client(
10382
*,
10483
as_service: str | None = None,
@@ -112,8 +91,7 @@ def get_nemo_client(
11291
"""
11392
headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of)
11493
base_url = _base_url()
115-
auth = _workload_exchange_auth_from_env(base_url, headers)
116-
return NemoClient(base_url=base_url, auth=auth, default_headers=headers or None)
94+
return NemoClient(base_url=base_url, default_headers=headers or None)
11795

11896

11997
def get_async_nemo_client(
@@ -129,5 +107,4 @@ def get_async_nemo_client(
129107
"""
130108
headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of)
131109
base_url = _base_url()
132-
auth = _workload_exchange_auth_from_env(base_url, headers)
133-
return AsyncNemoClient(base_url=base_url, auth=auth, default_headers=headers or None)
110+
return AsyncNemoClient(base_url=base_url, default_headers=headers or None)

packages/nemo_platform_plugin/tests/test_client_auth.py

Lines changed: 67 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -130,6 +130,52 @@ def test_no_auth_no_header(self):
130130
assert route.called
131131
assert "Authorization" not in route.calls[0].request.headers
132132

133+
def test_constructor_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path):
134+
subject_token_file = tmp_path / "workload-token"
135+
subject_token_file.write_text("subject-token\n", encoding="utf-8")
136+
provider = object()
137+
monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file))
138+
139+
with patch(
140+
"nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider",
141+
return_value=provider,
142+
) as resolve_provider:
143+
client = NemoClient(base_url="https://nemo.example.com")
144+
145+
assert client._auth is provider
146+
resolve_provider.assert_called_once_with(
147+
base_url="https://nemo.example.com",
148+
subject_token_file=subject_token_file,
149+
)
150+
151+
def test_constructor_workload_exchange_does_not_override_authorization_header(self, monkeypatch, tmp_path):
152+
subject_token_file = tmp_path / "workload-token"
153+
subject_token_file.write_text("subject-token\n", encoding="utf-8")
154+
monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file))
155+
156+
with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider:
157+
client = NemoClient(
158+
base_url="https://nemo.example.com",
159+
default_headers={"Authorization": "Bearer explicit-token"},
160+
)
161+
162+
assert client._auth is None
163+
resolve_provider.assert_not_called()
164+
165+
def test_constructor_workload_exchange_does_not_override_principal_header(self, monkeypatch, tmp_path):
166+
subject_token_file = tmp_path / "workload-token"
167+
subject_token_file.write_text("subject-token\n", encoding="utf-8")
168+
monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file))
169+
170+
with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider:
171+
client = NemoClient(
172+
base_url="https://nemo.example.com",
173+
default_headers={"X-NMP-Principal-Id": "service:jobs"},
174+
)
175+
176+
assert client._auth is None
177+
resolve_provider.assert_not_called()
178+
133179

134180
# ---------------------------------------------------------------------------
135181
# AsyncNemoClient auth parameter
@@ -176,6 +222,24 @@ def test_sync_provider_works_in_async_client(self):
176222
assert route.called
177223
assert route.calls[0].request.headers["Authorization"] == "Bearer sync-token"
178224

225+
def test_constructor_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path):
226+
subject_token_file = tmp_path / "workload-token"
227+
subject_token_file.write_text("subject-token\n", encoding="utf-8")
228+
provider = object()
229+
monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file))
230+
231+
with patch(
232+
"nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider",
233+
return_value=provider,
234+
) as resolve_provider:
235+
client = AsyncNemoClient(base_url="https://nemo.example.com")
236+
237+
assert client._auth is provider
238+
resolve_provider.assert_called_once_with(
239+
base_url="https://nemo.example.com",
240+
subject_token_file=subject_token_file,
241+
)
242+
179243

180244
# ---------------------------------------------------------------------------
181245
# OIDCTokenProvider
@@ -802,7 +866,7 @@ def test_sync_client_uses_workload_exchange_provider_from_env(self, monkeypatch,
802866
monkeypatch.delenv("NMP_PRINCIPAL", raising=False)
803867

804868
with patch(
805-
"nemo_platform_plugin.client_provider.resolve_workload_exchange_provider",
869+
"nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider",
806870
return_value=provider,
807871
) as resolve_provider:
808872
client = get_nemo_client()
@@ -822,7 +886,7 @@ def test_async_client_uses_workload_exchange_provider_from_env(self, monkeypatch
822886
monkeypatch.delenv("NMP_PRINCIPAL", raising=False)
823887

824888
with patch(
825-
"nemo_platform_plugin.client_provider.resolve_workload_exchange_provider",
889+
"nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider",
826890
return_value=provider,
827891
) as resolve_provider:
828892
client = get_async_nemo_client()
@@ -839,7 +903,7 @@ def test_workload_exchange_provider_does_not_override_explicit_service(self, mon
839903
monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file))
840904
monkeypatch.delenv("NMP_PRINCIPAL", raising=False)
841905

842-
with patch("nemo_platform_plugin.client_provider.resolve_workload_exchange_provider") as resolve_provider:
906+
with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider:
843907
client = get_nemo_client(as_service="jobs")
844908

845909
assert client._auth is None

sdk/python/nemo-platform/.nmpcontext/openapi.yaml

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

services/core/auth/src/nmp/core/auth/api/v2/workload_token_exchange.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -99,14 +99,17 @@
9999
"description": "Requested audience for the issued access token.",
100100
},
101101
"resource": {
102+
"type": "string",
102103
"description": "Unsupported for the jobs workload token exchange profile.",
103104
"not": {},
104105
},
105106
"actor_token": {
107+
"type": "string",
106108
"description": "Unsupported for the jobs workload token exchange profile.",
107109
"not": {},
108110
},
109111
"actor_token_type": {
112+
"type": "string",
110113
"description": "Unsupported for the jobs workload token exchange profile.",
111114
"not": {},
112115
},

0 commit comments

Comments
 (0)