Skip to content

Commit 4da1682

Browse files
committed
address feedback
Signed-off-by: Ryan S <267728323+ironcommit@users.noreply.github.com>
1 parent 0370bfc commit 4da1682

11 files changed

Lines changed: 256 additions & 11 deletions

File tree

packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,12 +16,12 @@
1616
from urllib.parse import urlparse
1717

1818
import httpx
19-
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
20-
from nmp.common.auth import (
19+
from nemo_platform_plugin.client.constants import (
2120
DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE as _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE,
2221
)
23-
from nmp.common.auth import (
22+
from nemo_platform_plugin.client.constants import (
2423
JWT_WORKLOAD_SUBJECT_TOKEN_TYPE,
24+
WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR,
2525
subject_token_type_for_exchange,
2626
)
2727

packages/nemo_platform_ext/tests/auth/test_workload_exchange.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,12 +3,15 @@
33

44
"""Tests for RFC 8693 workload identity token exchange."""
55

6+
import ast
67
import json
78
import time
89
from base64 import urlsafe_b64encode
10+
from pathlib import Path
911
from unittest.mock import MagicMock, patch
1012

1113
import httpx
14+
import nemo_platform_ext.auth.workload_exchange as workload_exchange_module
1215
import pytest
1316
from nemo_platform_ext.auth.workload_exchange import (
1417
ACCESS_TOKEN_TYPE,
@@ -24,6 +27,19 @@
2427
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
2528

2629

30+
def test_workload_exchange_module_has_no_nmp_common_dependency():
31+
assert workload_exchange_module.__file__ is not None
32+
tree = ast.parse(Path(workload_exchange_module.__file__).read_text(encoding="utf-8"))
33+
for node in ast.walk(tree):
34+
if isinstance(node, ast.Import):
35+
imported = [alias.name for alias in node.names]
36+
elif isinstance(node, ast.ImportFrom):
37+
imported = [node.module or ""]
38+
else:
39+
continue
40+
assert not any(name == "nmp.common" or name.startswith("nmp.common.") for name in imported)
41+
42+
2743
def _make_jwt(claims: dict) -> str:
2844
header = {"alg": "RS256", "typ": "JWT"}
2945
h = urlsafe_b64encode(json.dumps(header).encode()).rstrip(b"=").decode()

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

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,17 @@
66
import os
77

88
WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE"
9+
JWT_WORKLOAD_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt"
10+
DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = "urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof"
11+
OPAQUE_DOCKER_PROOF_PREFIX = "nmp_obo_v1"
912

1013

1114
def is_workload_identity_token_file_set() -> bool:
1215
return bool(os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR))
16+
17+
18+
def subject_token_type_for_exchange(subject_token: str) -> str:
19+
"""Return the RFC 8693 subject_token_type for a workload identity subject token."""
20+
if subject_token.startswith(f"{OPAQUE_DOCKER_PROOF_PREFIX}."):
21+
return DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE
22+
return JWT_WORKLOAD_SUBJECT_TOKEN_TYPE

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

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,15 +30,15 @@
3030
from urllib.parse import urlparse
3131

3232
import httpx
33-
from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR
34-
from nemo_platform_plugin.client.tls import client_verify_from_env
35-
from nmp.common.auth import (
33+
from nemo_platform_plugin.client.constants import (
3634
DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE as _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE,
3735
)
38-
from nmp.common.auth import (
36+
from nemo_platform_plugin.client.constants import (
3937
JWT_WORKLOAD_SUBJECT_TOKEN_TYPE,
38+
WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR,
4039
subject_token_type_for_exchange,
4140
)
41+
from nemo_platform_plugin.client.tls import client_verify_from_env
4242

4343
logger = logging.getLogger(__name__)
4444

packages/nemo_platform_plugin/tests/test_client_auth.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,14 @@
55

66
from __future__ import annotations
77

8+
import ast
89
import asyncio
910
import time
11+
from pathlib import Path
1012
from unittest.mock import patch
1113

1214
import httpx
15+
import nemo_platform_plugin.client.oidc as oidc_module
1316
import pytest
1417
import respx
1518
import yaml
@@ -45,6 +48,24 @@
4548
get_nemo_client,
4649
)
4750

51+
52+
def _assert_no_nmp_common_import(module_file: str | None) -> None:
53+
assert module_file is not None
54+
tree = ast.parse(Path(module_file).read_text(encoding="utf-8"))
55+
for node in ast.walk(tree):
56+
if isinstance(node, ast.Import):
57+
imported = [alias.name for alias in node.names]
58+
elif isinstance(node, ast.ImportFrom):
59+
imported = [node.module or ""]
60+
else:
61+
continue
62+
assert not any(name == "nmp.common" or name.startswith("nmp.common.") for name in imported)
63+
64+
65+
def test_oidc_module_has_no_nmp_common_dependency():
66+
_assert_no_nmp_common_import(oidc_module.__file__)
67+
68+
4869
# ---------------------------------------------------------------------------
4970
# StaticToken
5071
# ---------------------------------------------------------------------------

sdk/python/nemo-platform/src/nemo_platform/auth/workload_exchange.py

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

sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/auth/test_workload_exchange.py

Lines changed: 16 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/authenticate.py

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import jwt
1010
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
1111
from nmp.common.auth.bearer import MalformedBearerTokenError, parse_bearer_authorization_header
12-
from nmp.common.auth.jwt import TokenClaims
12+
from nmp.common.auth.jwt import ActorClaims, TokenClaims
1313
from nmp.common.auth.token_resolver import ResolvedBearerToken, ResolvedTokenKind, resolve_bearer_token
1414
from nmp.common.config import AuthConfig, get_auth_config
1515
from nmp.core.auth.api.v2.workload_token_exchange import (
@@ -80,6 +80,25 @@ def _scopes_from_claims(claims: dict[str, object]) -> list[str]:
8080
return []
8181

8282

83+
def _actor_from_claims(claims: dict[str, object]) -> ActorClaims | None:
84+
actor_claims = claims.get("act")
85+
if not isinstance(actor_claims, dict):
86+
return None
87+
88+
actor_subject = actor_claims.get("sub")
89+
if not isinstance(actor_subject, str):
90+
return None
91+
92+
actor_subject = actor_subject.strip()
93+
if not actor_subject:
94+
return None
95+
96+
return ActorClaims(
97+
subject=actor_subject,
98+
groups=_groups_from_claim(actor_claims.get("groups", [])),
99+
)
100+
101+
83102
def _stamp_principal_headers(response: Response, resolved: ResolvedBearerToken) -> None:
84103
for header_name, header_value in resolved.principal_headers().items():
85104
response.headers[header_name] = header_value
@@ -138,6 +157,7 @@ async def _validate_workload_access_token(
138157
groups=_groups_from_claim(claims.get("groups", [])),
139158
scopes=_scopes_from_claims(claims),
140159
raw_claims=claims,
160+
actor=_actor_from_claims(claims),
141161
)
142162
except jwt.PyJWTError:
143163
return None

services/core/auth/tests/test_authenticate.py

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -197,6 +197,59 @@ def test_authenticate_workload_access_token_returns_principal_headers(tmp_path):
197197
assert response.headers["X-NMP-Scopes"] == "openid email groups"
198198

199199

200+
def test_authenticate_delegated_workload_access_token_returns_obo_principal_headers(tmp_path):
201+
private_key_file = tmp_path / "private.pem"
202+
private_key_file.write_bytes(_private_key_pem())
203+
config = AuthConfig(
204+
enabled=True,
205+
token_signing=TokenSigningConfig(
206+
issuer="http://testserver/apis/auth",
207+
key_id="test-workload",
208+
private_key_file=str(private_key_file),
209+
),
210+
oidc=OIDCConfig(
211+
workload_token_exchange_enabled=True,
212+
workload_audience="nemo-platform",
213+
),
214+
)
215+
signing_key = WorkloadTokenExchangeService().workload_signing_key(config)
216+
now = datetime.now(tz=UTC)
217+
token = jwt.encode(
218+
{
219+
"iss": "http://testserver/apis/auth",
220+
"sub": "submitter@example.com",
221+
"email": "submitter@example.com",
222+
"aud": "nemo-platform",
223+
"iat": now,
224+
"nbf": now,
225+
"exp": now + timedelta(minutes=5),
226+
"scope": "openid email groups",
227+
"groups": "workspace-editors",
228+
"act": {
229+
"sub": "system:serviceaccount:nemo:job",
230+
"groups": ["system:serviceaccounts", "nemo-jobs"],
231+
},
232+
},
233+
signing_key.private_key,
234+
algorithm="RS256",
235+
headers={"kid": signing_key.kid},
236+
)
237+
with _test_client(config) as client:
238+
response = client.get(
239+
"/authenticate",
240+
headers={"Authorization": f"Bearer {token}"},
241+
)
242+
243+
assert response.status_code == 200
244+
assert response.json()["token_kind"] == "workload_access_token"
245+
assert response.headers["X-NMP-Principal-Id"] == "system:serviceaccount:nemo:job"
246+
assert response.headers["X-NMP-Principal-Groups"] == "system:serviceaccounts,nemo-jobs"
247+
assert response.headers["X-NMP-Principal-On-Behalf-Of"] == "submitter@example.com"
248+
assert response.headers["X-NMP-Principal-On-Behalf-Of-Email"] == "submitter@example.com"
249+
assert response.headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "workspace-editors"
250+
assert response.headers["X-NMP-Scopes"] == "openid email groups"
251+
252+
200253
def test_authenticate_workload_subject_token_uses_resolver_callback(tmp_path):
201254
config = AuthConfig(
202255
enabled=True,

services/core/jobs/src/nmp/core/jobs/controllers/backends/docker.py

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -561,6 +561,51 @@ def _revoke_workload_delegation_after_failed_start(self, delegation_name: str |
561561
extra={"delegation_name": delegation_name},
562562
)
563563

564+
def _workload_delegation_name_for_step(self, step: PlatformJobStepWithContext) -> str:
565+
return docker_delegation_name(
566+
workload_workspace=step.workspace,
567+
job_id=step.job,
568+
attempt_id=step.attempt_id,
569+
step_id=step.id,
570+
)
571+
572+
@staticmethod
573+
def _workload_delegation_name_from_container(container: Container) -> str | None:
574+
labels = getattr(container, "labels", None) or {}
575+
if DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL not in labels:
576+
return None
577+
578+
workload_workspace = labels.get(JOB_WORKSPACE_ID_LABEL)
579+
job_id = labels.get(JOB_ID_LABEL)
580+
attempt_id = labels.get(JOB_ATTEMPT_ID_LABEL)
581+
step_id = labels.get(JOB_STEP_ID_LABEL)
582+
if not isinstance(workload_workspace, str) or not workload_workspace:
583+
return None
584+
if not isinstance(job_id, str) or not job_id:
585+
return None
586+
if not isinstance(attempt_id, str) or not attempt_id:
587+
return None
588+
if not isinstance(step_id, str) or not step_id:
589+
return None
590+
591+
return docker_delegation_name(
592+
workload_workspace=workload_workspace,
593+
job_id=job_id,
594+
attempt_id=attempt_id,
595+
step_id=step_id,
596+
)
597+
598+
def _revoke_workload_delegation_after_terminal(self, delegation_name: str | None) -> None:
599+
if not delegation_name:
600+
return
601+
try:
602+
self._revoke_workload_delegation(delegation_name)
603+
except Exception:
604+
logger.exception(
605+
"Failed to revoke Docker workload delegation after terminal status",
606+
extra={"delegation_name": delegation_name},
607+
)
608+
564609
def _write_workload_identity_subject_token(self, volume_name: str, token: str) -> None:
565610
storage_config = self._execution_profile_config.storage
566611
permissions_image = (
@@ -1917,6 +1962,8 @@ def create_step_update(self, step: PlatformJobStepWithContext, container: Contai
19171962
error_details = {}
19181963
if status == PlatformJobStatus.ERROR:
19191964
error_details["message"] = status_details.get("message", "Job encountered an error")
1965+
if status in PlatformJobStatus.terminals() and self._should_enable_workload_identity_for_step(step):
1966+
self._revoke_workload_delegation_after_terminal(self._workload_delegation_name_for_step(step))
19201967

19211968
logger.debug(
19221969
"Docker container status mapped to platform status",
@@ -2166,8 +2213,10 @@ def cleanup_single_container(self, container: Container) -> None:
21662213
job = self.get_label_from_container(container, JOB_ID_LABEL)
21672214
task = self.get_label_from_container(container, JOB_TASK_ID_LABEL)
21682215
exit_code = container.attrs.get("State", {}).get("ExitCode", 0)
2216+
delegation_name = self._workload_delegation_name_from_container(container)
21692217

21702218
self._stop_workload_identity_refresher(container.name)
2219+
self._revoke_workload_delegation_after_terminal(delegation_name)
21712220
self.cleanup_container(container)
21722221
logger.debug(
21732222
"Cleaned up container",

0 commit comments

Comments
 (0)