diff --git a/contrib/auth/README.md b/contrib/auth/README.md index 7962cc9073..a43180e1f7 100644 --- a/contrib/auth/README.md +++ b/contrib/auth/README.md @@ -12,6 +12,15 @@ adaptation: group binding, not as internal `service:*` principals - document provider-specific setup in a local `README.md` +Managed job OBO tests use the provider for user and controller authentication, +but the workload-to-submitter binding is NeMo Platform auth state. Jobs receive +`NMP_WORKLOAD_IDENTITY_TOKEN_FILE`, exchange that subject token through +`/apis/auth/token`, and receive a NeMo Platform token whose top-level subject +is the job submitter and whose RFC 8693 `act.sub` is the workload actor. Provider +manifests may include workload-provider token grants for contract tests, but +managed Docker job OBO must not depend on provider-specific fields such as +`jti`. + Open-source providers with `mode: compose-ci` are intended for the shared auth matrix. Reference-only providers stay documented and manifest-driven but are excluded from the local Compose-backed matrix. diff --git a/contrib/auth/authentik/README.md b/contrib/auth/authentik/README.md index 2aa86ec0c4..96ee4d6aec 100644 --- a/contrib/auth/authentik/README.md +++ b/contrib/auth/authentik/README.md @@ -3,9 +3,10 @@ This directory contains a local Authentik-backed NeMo Platform example. Use it to validate three user-visible flows: -- log in to NeMo with Authentik -- call NeMo APIs through the Authentik gateway -- run a NeMo job whose workload exchanges a real Authentik workload subject token +- log in to NeMo Platform with Authentik +- call NeMo Platform APIs through the Authentik gateway +- run a NeMo Platform job whose workload exchanges a managed workload proof token for a + delegated NeMo Platform access token All credentials in this example are for local development only. @@ -85,8 +86,10 @@ The 2-minute CLI access-token lifetime is a local demo/testing setting so token refresh is easy to observe. Do not use it as a production default; use a longer value such as `hours=1` outside the refresh demonstration. -In the Docker Compose runtime, Authentik issues the demo workload subject token, -but it does not accept the RFC 8693 token exchange grant directly. The Docker -backend refreshes the Authentik subject token file, the SDK posts that token to -the NeMo auth service, and the gateway trusts the NeMo auth service JWKS for -exchanged workload access tokens. +In the Docker Compose runtime, Authentik authenticates users and controller +service principals, but managed Docker job OBO uses a NeMo Platform-owned +opaque workload proof token. The Docker backend writes that proof token into the +job token file, the SDK posts it to the NeMo Platform auth service, and the +gateway trusts the NeMo Platform auth service JWKS for exchanged workload access +tokens. Docker OBO does not depend on IdP `jti` claims or IdP-issued workload +subject tokens. diff --git a/contrib/auth/authentik/compose/implementation-details.md b/contrib/auth/authentik/compose/implementation-details.md index f0c7ebe0a1..de721e8303 100644 --- a/contrib/auth/authentik/compose/implementation-details.md +++ b/contrib/auth/authentik/compose/implementation-details.md @@ -17,12 +17,12 @@ the parent directory: - `../config/platform-compose-authentik.yaml` as the NeMo Platform config. - `../gateway/envoy.yaml` as the local gateway config. -- `../helm/files/blueprints` as the |product-name| blueprint source. +- `../helm/files/blueprints` as the NeMo Platform blueprint source. - `../.generated` for local generated keys and certificates. -The shared tutorial does not build NeMo images for Compose. It runs -`${IMAGE_REGISTRY:-my-registry}/nmp-api:${BAKE_TAG:-local}` for both the NeMo -API service and workload jobs submitted by the tutorial. +The shared tutorial does not build NeMo Platform images for Compose. It runs +`${IMAGE_REGISTRY:-my-registry}/nmp-api:${BAKE_TAG:-local}` for both the +NeMo Platform API service and workload jobs submitted by the tutorial. ## Services @@ -34,12 +34,12 @@ The stack contains: - `gateway-tls-init`: a small init container that copies local TLS material into the named `gateway-tls` volume with permissions suitable for Envoy. - `authentik-blueprint-init`: a one-shot init container that applies the shared - |product-name| blueprint before the gateway starts. + NeMo Platform blueprint before the gateway starts. - `authentik-postgres`: PostgreSQL for Authentik. - `authentik-redis`: Redis for Authentik. - `authentik-server` and `authentik-worker`: Authentik itself. -`nemo` is only on the internal network. Host and workload traffic reaches NeMo +`nemo` is only on the internal network. Host and workload traffic reaches NeMo Platform through the `gateway` service, which also joins the workload network as `nemo-gateway`. @@ -56,15 +56,15 @@ share local keys: The workload-token private key is mounted into `nemo` at `/var/run/secrets/nemo-platform/workload-token-signing/private-key.pem`. `platform-compose-authentik.yaml` points -`auth.token_signing.private_key_file` at that mounted path. The NeMo auth -service uses the private key to sign workload-exchange access tokens and Scoped -Access Key JWTs, and Envoy validates those tokens through the NeMo auth service -JWKS endpoints. +`auth.token_signing.private_key_file` at that mounted path. The NeMo Platform +auth service uses the private key to sign workload-exchange access tokens and +Scoped Access Key JWTs, and Envoy validates those tokens through the +NeMo Platform auth service JWKS endpoints. The gateway TLS files are copied into the `gateway-tls` named volume by `gateway-tls-init`. The `gateway` service uses that volume to serve HTTPS, and the `nemo` service mounts the same volume read-only so Python HTTP clients -inside NeMo trust the demo gateway certificate. +inside NeMo Platform trust the demo gateway certificate. All generated keys and certificates in this example are for local development only. @@ -91,7 +91,7 @@ The `nemo-setup` service account and app-password in the blueprint exist solely for automated auth-idp contract tests. They are not part of the browser login flow or the workload identity pattern. -## NeMo Compose Configuration +## NeMo Platform Compose Configuration `platform-compose-authentik.yaml` configures NeMo Platform for this topology: @@ -101,8 +101,11 @@ flow or the workload identity pattern. the `nemo` container. - Host-side CLI login uses the port-forward-like public gateway URL `https://127.0.0.1:18080`. -- Workload subject tokens come from Authentik's workload OIDC provider. -- Exchanged workload access tokens come from NeMo's `/apis/auth/token` endpoint. +- Authentik provides user and controller service-principal authentication. The + Docker managed-job OBO binding is stored in NeMo Platform auth delegation + state, not in Authentik. +- Exchanged workload access tokens come from NeMo Platform's `/apis/auth/token` + endpoint. The Docker jobs executor mounts the `gateway-tls` volume into workload containers and sets `SSL_CERT_FILE` and `REQUESTS_CA_BUNDLE` so workload code @@ -112,23 +115,23 @@ trusts the local gateway certificate. Envoy is the public entrypoint for the Compose example. It routes: -- NeMo paths such as `/.well-known/nemo-platform/`, `/apis/`, `/health/`, - `/status`, and `/studio/` to `nemo`. +- NeMo Platform paths such as `/.well-known/nemo-platform/`, `/apis/`, + `/health/`, `/status`, and `/studio/` to `nemo`. - `/health/gateway/ready` to an Envoy-owned readiness check that verifies both - NeMo and Authentik through their upstream clusters. + NeMo Platform and Authentik through their upstream clusters. - Authentik paths to `authentik-server`. Before authentication, Envoy removes incoming `X-NMP-Principal-*` and `X-NMP-Scopes` headers so a client cannot spoof identity or scopes. For -protected `/apis/` requests, Envoy calls NeMo's +protected `/apis/` requests, Envoy calls NeMo Platform's `/apis/auth/authenticate` endpoint with the presented bearer token. The auth -service validates Authentik OIDC tokens, NeMo workload-exchange access tokens, -and NeMo Scoped Access Keys, then returns trusted `X-NMP-Principal-*` and -`X-NMP-Scopes` headers for Envoy to forward upstream. +service validates Authentik OIDC tokens, NeMo Platform workload-exchange access +tokens, and NeMo Platform Scoped Access Keys, then returns trusted +`X-NMP-Principal-*` and `X-NMP-Scopes` headers for Envoy to forward upstream. The gateway callout is required for dynamic or revocable Scoped Access Keys because Envoy JWKS validation can only prove token signature, issuer, audience, -and time claims. It cannot check NeMo's access-key lifecycle state. Compose +and time claims. It cannot check NeMo Platform's access-key lifecycle state. Compose keeps `auth.access_keys.enabled=true` so Scoped Access Keys can be created and validated; Envoy performs the bearer-to-header mapping before the request reaches service middleware. @@ -142,7 +145,7 @@ job request should not include `NMP_WORKLOAD_IDENTITY_TOKEN_FILE`, `NEMO_WORKLOAD_TOKEN`, or `NEMO_WORKLOAD_TOKEN_FILE`. When a managed Docker workload starts, the backend creates a dedicated workload -identity volume, writes an Authentik subject token to: +identity volume and writes a NeMo Platform-owned Docker workload proof token to: ```text /var/run/secrets/nemo-platform/workload/token @@ -155,9 +158,21 @@ NMP_WORKLOAD_IDENTITY_TOKEN_FILE=/var/run/secrets/nemo-platform/workload/token ``` The SDK reads that file and sends an RFC 8693 token exchange request to the -NeMo auth service through the gateway. The NeMo auth service validates the -Authentik subject token, mints a NeMo-signed access token, and returns it to the -workload. The workload uses that exchanged token for normal NeMo API calls. +NeMo Platform auth service through the gateway. The Docker backend registered +an internal workload delegation row before the container started. The +NeMo Platform auth service validates the proof token, checks the matching row, +mints a NeMo Platform-signed delegated access token, and returns it to the +workload. The access token uses the captured job submitter as the top-level +subject and the Docker workload as the RFC 8693 `act.sub` actor. + +Docker supports one proof-token mechanism in this flow. The file contains a +private opaque proof token whose secret is stored only as a hash in the +delegation row. + +Docker job OBO therefore does not require Authentik to issue a workload token +and does not depend on an IdP `jti` claim. The Authentik workload-provider +configuration in the manifest is retained for direct provider-token contract +tests, not for the managed Docker job exchange loop. The useful end-to-end validation is the workload job in the shared tutorial: the job uses the exchanged token to call the NeMo Platform API and read the diff --git a/contrib/auth/authentik/config/platform-compose-authentik.yaml b/contrib/auth/authentik/config/platform-compose-authentik.yaml index 03d84acc59..8e791d37dc 100644 --- a/contrib/auth/authentik/config/platform-compose-authentik.yaml +++ b/contrib/auth/authentik/config/platform-compose-authentik.yaml @@ -61,10 +61,6 @@ jobs: additional_volume_mounts: - volume_name: "authentik_gateway_tls" mount_path: "/etc/nmp/gateway-tls" - workload_identity: - token_endpoint: "https://nemo-gateway:8080/application/o/token/" - username: "svc-nemo" - password_env_var: "AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD" executor_defaults: docker: cleanup_completed_jobs_immediately: false diff --git a/contrib/auth/authentik/gateway/envoy.yaml b/contrib/auth/authentik/gateway/envoy.yaml index 1dd8597bc2..439b249ea5 100644 --- a/contrib/auth/authentik/gateway/envoy.yaml +++ b/contrib/auth/authentik/gateway/envoy.yaml @@ -239,6 +239,9 @@ static_resources: - exact: x-nmp-principal-id - exact: x-nmp-principal-email - exact: x-nmp-principal-groups + - exact: x-nmp-principal-on-behalf-of + - exact: x-nmp-principal-on-behalf-of-email + - exact: x-nmp-principal-on-behalf-of-groups - exact: x-nmp-scopes allowed_client_headers: patterns: diff --git a/contrib/auth/authentik/helm/templates/_envoy-config.tpl b/contrib/auth/authentik/helm/templates/_envoy-config.tpl index 1ccd850966..c13868d58e 100644 --- a/contrib/auth/authentik/helm/templates/_envoy-config.tpl +++ b/contrib/auth/authentik/helm/templates/_envoy-config.tpl @@ -245,6 +245,9 @@ static_resources: - exact: x-nmp-principal-id - exact: x-nmp-principal-email - exact: x-nmp-principal-groups + - exact: x-nmp-principal-on-behalf-of + - exact: x-nmp-principal-on-behalf-of-email + - exact: x-nmp-principal-on-behalf-of-groups - exact: x-nmp-scopes allowed_client_headers: patterns: diff --git a/contrib/auth/authentik/kubernetes/implementation-details.md b/contrib/auth/authentik/kubernetes/implementation-details.md index 017c797b7e..c7668ad7c2 100644 --- a/contrib/auth/authentik/kubernetes/implementation-details.md +++ b/contrib/auth/authentik/kubernetes/implementation-details.md @@ -34,11 +34,11 @@ upgrade command. The chart creates or reuses these additional local-demo Secrets during Helm rendering: -- `shared-postgresql` for the shared PostgreSQL superuser, Authentik, and NeMo +- `shared-postgresql` for the shared PostgreSQL superuser, Authentik, and NeMo Platform database passwords. - `shared-postgresql-nemo` for the NeMo Platform external database password. - `nemo-platform-envoy-tls` for the demo Envoy TLS certificate and CA. -- `nemo-workload-token-signing-key` for the NeMo-issued workload access token +- `nemo-workload-token-signing-key` for the NeMo Platform-issued workload access token signing key. The chart generates `Secret/nemo-platform-envoy-tls` during Helm rendering and @@ -59,13 +59,14 @@ mounted file. Use `--wait --wait-for-jobs` when installing the chart so Helm only returns after the blueprint has been applied. -## NeMo Kubernetes Override +## NeMo Platform Kubernetes Override The umbrella chart passes the Kubernetes-specific NeMo Platform configuration through `nemo-platform.platformConfig` values. It also configures `nemo-platform.envoyProxy.configOverride` so the NeMo Platform chart's Envoy deployment keeps the Authentik path split and validates both Authentik-issued -tokens, NeMo workload-exchange tokens, and NeMo Scoped Access Key JWTs. +tokens, NeMo Platform workload-exchange tokens, and NeMo Platform Scoped +Access Key JWTs. Kubernetes projected service account token expiration defaults to `600` seconds in the jobs backend. Override it through the NeMo Platform chart values if you @@ -80,23 +81,36 @@ workload pods and injects: NMP_WORKLOAD_IDENTITY_TOKEN_FILE=/var/run/secrets/nemo-platform/workload/token ``` -The SDK reads that file and sends an RFC 8693 token exchange request to the NeMo -auth service over HTTPS. The chart mounts `ca.crt` from +The SDK reads that file and sends an RFC 8693 token exchange request to the +NeMo Platform auth service over HTTPS. The chart mounts `ca.crt` from `Secret/nemo-platform-envoy-tls` into Kubernetes workload pods and sets `SSL_CERT_FILE` and `REQUESTS_CA_BUNDLE` so in-pod Python HTTP clients verify the demo Envoy certificate. Host-side `nemo` commands should use `NMP_CLIENT_SSL_CERT_FILE` instead so unrelated tools keep their normal trust store. -The NeMo auth service validates projected service account tokens with the -TokenReview API and returns a NeMo-signed JWT trusted by the NeMo Platform -Envoy. The useful end-to-end validation is the workload job in the tutorial: -the job pod uses the exchanged token to call the NeMo Platform API and read the -workspace. +The NeMo Platform auth service validates projected service account tokens with +the TokenReview API and returns a NeMo Platform-signed JWT trusted by the +NeMo Platform Envoy. The TokenReview response must include exactly one +`authentication.kubernetes.io/pod-uid` value. The jobs controller observes the +created Pod, registers an internal delegation row keyed by that Pod UID, +service account subject, and audience, and auth later looks up that row from +the verified TokenReview metadata. The useful end-to-end validation is the +workload job in the tutorial: the job pod uses the exchanged token to call the +NeMo Platform API and read the workspace. + +Auth only needs RBAC to create TokenReview requests. It does not need to read +Pods or Jobs for token exchange, and the chart grants no Pod or Job read +permissions to the auth/API service account for this path. Kubernetes labels, +annotations, and owner references help the jobs controller reconcile and clean +up backend resources, but they are not token-exchange authorization inputs. The workload job request should not include workload auth environment variables. The Kubernetes jobs backend owns `NMP_WORKLOAD_IDENTITY_TOKEN_FILE`; -users must not set `NEMO_WORKLOAD_TOKEN` or `NEMO_WORKLOAD_TOKEN_FILE`. +users must not set `NMP_PRINCIPAL`, `NEMO_WORKLOAD_TOKEN`, or +`NEMO_WORKLOAD_TOKEN_FILE`. When workload OBO is enabled, the workload receives +only the subject token file path and obtains its delegated NeMo Platform access +token by calling `/apis/auth/token`. To inspect the projected token mount for a submitted job: @@ -110,8 +124,8 @@ kubectl --context "${KUBE_CONTEXT}" -n "${NAMESPACE}" describe pod \ ## Shared Token Signing Key -Workload identity token exchange requires the NeMo auth service to sign the -access token it mints from a Kubernetes projected service account subject +Workload identity token exchange requires the NeMo Platform auth service to +sign the access token it mints from a Kubernetes projected service account subject token. The chart creates `Secret/nemo-workload-token-signing-key` by default, mounts `private-key.pem` into the NeMo Platform API pod at `/etc/nmp/workload-token/private-key.pem`, and sets @@ -120,9 +134,9 @@ configuration is used for workload-exchange access tokens and Scoped Access Key JWTs. The matching public keys are served from `/apis/auth/jwks`. The Helm-rendered Envoy config authenticates protected `/apis/` requests by -calling `/apis/auth/authenticate` on the NeMo API service. It does not use Envoy -`claim_to_headers` for Scoped Access Keys. This keeps future revocation and -dynamic-key checks inside the auth service, where access-key records can be +calling `/apis/auth/authenticate` on the NeMo Platform API service. It does not +use Envoy `claim_to_headers` for Scoped Access Keys. This keeps future +revocation and dynamic-key checks inside the auth service, where access-key records can be looked up before Envoy forwards trusted principal headers. Scoped Access Keys remain disabled in the checked-in chart values by default. diff --git a/contrib/auth/authentik/manifest.yaml b/contrib/auth/authentik/manifest.yaml index d108f59b57..13a2d6e106 100644 --- a/contrib/auth/authentik/manifest.yaml +++ b/contrib/auth/authentik/manifest.yaml @@ -68,8 +68,8 @@ test_runtimes: - platform_access_keys - workspace_rbac - workload_job + - managed_workload_job_obo - device_flow - - docker_subject_token_refresh - id: authentik-kubernetes backend: kubernetes command: k8s @@ -83,5 +83,6 @@ test_runtimes: - platform_access_keys - workspace_rbac - workload_job + - managed_workload_job_obo - device_flow - kubernetes_token_review diff --git a/docs/auth/deployment/configuration.mdx b/docs/auth/deployment/configuration.mdx index 8a2eaca2bd..925b8d3e33 100644 --- a/docs/auth/deployment/configuration.mdx +++ b/docs/auth/deployment/configuration.mdx @@ -198,27 +198,10 @@ and optional `scope`. `workload_token_endpoint` is optional; when it is unset, the SDK uses `token_endpoint`. This is useful when host CLI login and workload containers need different network-reachable IdP URLs. -For Docker-backed job runtimes, the executor may also need a controller-side -subject-token issuer. Configure only Docker-specific fields under the Docker -executor profile: - -```yaml -jobs: - executors: - - provider: cpu - profile: workload - backend: docker - config: - workload_identity: - token_endpoint: "https://idp.example.com/oauth/token" - client_id: "nemo-platform-workload" - username: "svc-nemo" - password_env_var: "WORKLOAD_IDENTITY_PASSWORD" - scope: "openid email groups" -``` - -The password value is read from the controller process environment using -`password_env_var`. The config does not support an inline `password` field. +For Docker-backed job runtimes, `auth.oidc.workload_token_exchange_enabled` +controls workload identity. When enabled, job steps with a delegation auth +context receive a NeMo Platform opaque workload proof token at +`NMP_WORKLOAD_IDENTITY_TOKEN_FILE`. For Kubernetes-backed job runtimes, the jobs backend can project a Kubernetes service account token into each workload pod for token exchange. The projected diff --git a/docs/auth/deployment/credential-propagation.mdx b/docs/auth/deployment/credential-propagation.mdx index 78a388af8d..d5c645128c 100644 --- a/docs/auth/deployment/credential-propagation.mdx +++ b/docs/auth/deployment/credential-propagation.mdx @@ -15,18 +15,23 @@ principal as an API credential. The flow: 1. The user submits a job via the API. The control plane authorizes job creation - and records the job metadata. + and records the submitter auth context on the job metadata. 2. When workload identity token exchange is enabled, the managed backend creates a workload with `NMP_WORKLOAD_IDENTITY_TOKEN_FILE` set to a subject-token file - path. -3. The backend owns that subject-token file. Kubernetes uses a projected service - account token volume, while Docker uses a controller-managed volume that is - refreshed by the jobs controller. + path. This is the only platform credential mounted into the workload for job + OBO. +3. The backend owns that subject-token file. Kubernetes uses a projected, + pod-bound service account token volume, while Docker uses a + controller-managed volume containing a Docker workload proof token. 4. The SDK in the workload detects `NMP_WORKLOAD_IDENTITY_TOKEN_FILE`, reads the subject token, discovers the workload token exchange endpoint, and exchanges the subject token for a NeMo Platform access token using OAuth 2.0 Token Exchange (RFC 8693). -5. API calls from the workload use the exchanged access token. When that access +5. The auth service validates the workload proof, fetches the matching internal + workload delegation row, and returns a delegated NeMo access token. The token + top-level `sub` is the captured submitter, and the RFC 8693 `act.sub` claim + is the workload actor. +6. API calls from the workload use the exchanged access token. When that access token nears expiry, the SDK rereads the subject-token file and performs another exchange. @@ -38,26 +43,43 @@ Model](/documentation/access-control/security-model#job-credential-propagation). ## Workload Identity Token Files -For external IdP-backed workload identity, managed backends inject +For workload identity, managed backends inject `NMP_WORKLOAD_IDENTITY_TOKEN_FILE`. The file contains a workload identity -subject token, not a final NeMo API access token. +subject token or proof token, not a final NeMo API access token. The SDK reads that file immediately before OAuth 2.0 Token Exchange (RFC 8693), -exchanges the subject token at the discovered IdP token endpoint, caches the -returned access token until near expiry, and then rereads the file for the next -exchange. +exchanges the subject token at the discovered workload token exchange endpoint, +caches the returned access token until near expiry, and then rereads the file +for the next exchange. Backend ownership: - Kubernetes uses a projected service account token volume. Kubelet rotates the - file. + file. The auth service validates the presented token with Kubernetes + TokenReview and uses the verified + `authentication.kubernetes.io/pod-uid` reference from TokenReview metadata to + look up the internal delegation row. - Docker uses a dedicated controller-managed workload identity volume. The - Docker backend writes and refreshes the file. - -Users must not provide `NEMO_WORKLOAD_TOKEN`, `NEMO_WORKLOAD_TOKEN_FILE`, or -`NMP_WORKLOAD_IDENTITY_TOKEN_FILE` in managed job requests. Direct SDK users may -set `NMP_WORKLOAD_IDENTITY_TOKEN_FILE` only when they own the refreshed subject -token file. + Docker backend writes a NeMo opaque proof token bound to an internal + delegation row. The token uses the private Docker subject-token type and is + checked against the hash stored in the delegation row. Docker does not use + TokenReview or external IdP-issued workload proof tokens. + +Users must not provide `NMP_PRINCIPAL`, `NEMO_WORKLOAD_TOKEN`, +`NEMO_WORKLOAD_TOKEN_FILE`, or `NMP_WORKLOAD_IDENTITY_TOKEN_FILE` in managed job +requests. Direct SDK users may set `NMP_WORKLOAD_IDENTITY_TOKEN_FILE` only when +they own the refreshed subject token file. + +Workload delegation rows are internal auth state. NeMo Platform does not expose +workload-delegation-specific API endpoints, generated public SDK resources, CLI +commands, or UI flows. Platform administrators can still inspect generic entity +store records through the normal entity administration surface. + +Kubernetes labels, annotations, owner references, and Docker labels are +reconciliation and cleanup aids only. They are not token-exchange authorization +inputs. Token exchange authorizes against the verified subject token, the +deterministic delegation row name, the stored submitter auth context, audience, +expiry/revocation state, and any verifier-confirmed bound reference. This follows the same shape as common cloud SDKs: AWS uses `AWS_WEB_IDENTITY_TOKEN_FILE`, Azure uses `AZURE_FEDERATED_TOKEN_FILE`, and diff --git a/docs/cli/reference.mdx b/docs/cli/reference.mdx index 1681a493da..94d525e3b9 100644 --- a/docs/cli/reference.mdx +++ b/docs/cli/reference.mdx @@ -4047,6 +4047,11 @@ Filter jobs by workspace, project, name, status, source, created_at, and updated Get all currently configured execution profiles. +Returns the capability-filtered merge from jobs config. In local standalone the +controller may prune the shared list further after registry boot; in split +topologies the API advertises its own merge result (not controller process +memory). + **Usage:** ```shell diff --git a/docs/set-up/config-reference.mdx b/docs/set-up/config-reference.mdx index 94bb3245d8..a83fe291d8 100644 --- a/docs/set-up/config-reference.mdx +++ b/docs/set-up/config-reference.mdx @@ -299,26 +299,8 @@ jobs: networking: # Docker network for the job container | default: 'host' job_container_network: host - # Docker workload identity subject-token issuer configuration. - workload_identity: - # Enable Docker workload identity token-file injection. Defaults to auth.oidc.workload_token_exchange_enabled. - enabled: - # OAuth token endpoint used by the Docker demo issuer. Defaults to auth.oidc.token_endpoint. - token_endpoint: - # OAuth client ID used by the Docker demo issuer. Defaults to auth.oidc.workload_client_id or auth.oidc.client_id. - client_id: - # OAuth client secret for the Docker demo issuer. - client_secret: - # Username for the Docker demo issuer password grant. - username: - # Controller environment variable that contains the Docker demo issuer password grant shared secret. | default: 'AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD' - password_env_var: AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD - # OAuth scope for the Docker demo issuer. - scope: - # Fallback subject-token lifetime when the Docker demo issuer response omits expires_in. - subject_token_ttl_seconds: 600 - # Seconds before subject-token expiry when the Docker refresher issues a replacement token. | default: 60 - refresh_margin_seconds: 60 + # Docker workload identity configuration. + workload_identity: {} # Default Kubernetes execution profile configuration kubernetes_job: # default: 1800 diff --git a/openapi/ga/individual/platform.openapi.yaml b/openapi/ga/individual/platform.openapi.yaml index 33a0d3c022..d503c2233b 100644 --- a/openapi/ga/individual/platform.openapi.yaml +++ b/openapi/ga/individual/platform.openapi.yaml @@ -126,6 +126,7 @@ paths: type: string enum: - urn:ietf:params:oauth:token-type:jwt + - urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof description: Token type identifier for the subject token. requested_token_type: type: string @@ -136,6 +137,18 @@ paths: audience: type: string description: Requested audience for the issued access token. + resource: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token_type: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. scope: type: string description: Space-separated scopes requested for the issued access @@ -10443,9 +10456,7 @@ components: - $ref: '#/components/schemas/DockerJobNetworkConfig' description: Docker networking configuration workload_identity: - allOf: - - $ref: '#/components/schemas/DockerWorkloadIdentityConfig' - description: Docker workload identity subject-token issuer configuration. + $ref: '#/components/schemas/DockerWorkloadIdentityConfig' type: object title: DockerJobExecutionProfileConfig description: Configuration for Docker Job execution profile. @@ -10516,59 +10527,11 @@ components: - mount_path title: DockerVolumeMount DockerWorkloadIdentityConfig: - properties: - enabled: - title: Enabled - description: Enable Docker workload identity token-file injection. Defaults - to auth.oidc.workload_token_exchange_enabled. - type: boolean - token_endpoint: - title: Token Endpoint - description: OAuth token endpoint used by the Docker demo issuer. Defaults - to auth.oidc.token_endpoint. - type: string - client_id: - title: Client Id - description: OAuth client ID used by the Docker demo issuer. Defaults to - auth.oidc.workload_client_id or auth.oidc.client_id. - type: string - client_secret: - format: password - title: Client Secret - description: OAuth client secret for the Docker demo issuer. - writeOnly: true - type: string - username: - title: Username - description: Username for the Docker demo issuer password grant. - type: string - password_env_var: - type: string - title: Password Env Var - description: Controller environment variable that contains the Docker demo - issuer password grant shared secret. - default: AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD - scope: - title: Scope - description: OAuth scope for the Docker demo issuer. - type: string - subject_token_ttl_seconds: - type: integer - minimum: 1.0 - title: Subject Token Ttl Seconds - description: Fallback subject-token lifetime when the Docker demo issuer - response omits expires_in. - refresh_margin_seconds: - type: integer - minimum: 0.0 - title: Refresh Margin Seconds - description: Seconds before subject-token expiry when the Docker refresher - issues a replacement token. - default: 60 + properties: {} additionalProperties: false type: object title: DockerWorkloadIdentityConfig - description: Docker-only subject token issuer configuration for workload identity. + description: Docker workload identity configuration. E2EJobExecutionProfile: properties: provider: @@ -19544,7 +19507,7 @@ components: type: string title: Error description: OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, - invalid_request, invalid_grant, invalid_scope, or invalid_target. + invalid_request, invalid_scope, or invalid_target. error_description: title: Error Description description: Human-readable ASCII text providing additional information diff --git a/openapi/ga/openapi.yaml b/openapi/ga/openapi.yaml index 33a0d3c022..d503c2233b 100644 --- a/openapi/ga/openapi.yaml +++ b/openapi/ga/openapi.yaml @@ -126,6 +126,7 @@ paths: type: string enum: - urn:ietf:params:oauth:token-type:jwt + - urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof description: Token type identifier for the subject token. requested_token_type: type: string @@ -136,6 +137,18 @@ paths: audience: type: string description: Requested audience for the issued access token. + resource: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token_type: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. scope: type: string description: Space-separated scopes requested for the issued access @@ -10443,9 +10456,7 @@ components: - $ref: '#/components/schemas/DockerJobNetworkConfig' description: Docker networking configuration workload_identity: - allOf: - - $ref: '#/components/schemas/DockerWorkloadIdentityConfig' - description: Docker workload identity subject-token issuer configuration. + $ref: '#/components/schemas/DockerWorkloadIdentityConfig' type: object title: DockerJobExecutionProfileConfig description: Configuration for Docker Job execution profile. @@ -10516,59 +10527,11 @@ components: - mount_path title: DockerVolumeMount DockerWorkloadIdentityConfig: - properties: - enabled: - title: Enabled - description: Enable Docker workload identity token-file injection. Defaults - to auth.oidc.workload_token_exchange_enabled. - type: boolean - token_endpoint: - title: Token Endpoint - description: OAuth token endpoint used by the Docker demo issuer. Defaults - to auth.oidc.token_endpoint. - type: string - client_id: - title: Client Id - description: OAuth client ID used by the Docker demo issuer. Defaults to - auth.oidc.workload_client_id or auth.oidc.client_id. - type: string - client_secret: - format: password - title: Client Secret - description: OAuth client secret for the Docker demo issuer. - writeOnly: true - type: string - username: - title: Username - description: Username for the Docker demo issuer password grant. - type: string - password_env_var: - type: string - title: Password Env Var - description: Controller environment variable that contains the Docker demo - issuer password grant shared secret. - default: AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD - scope: - title: Scope - description: OAuth scope for the Docker demo issuer. - type: string - subject_token_ttl_seconds: - type: integer - minimum: 1.0 - title: Subject Token Ttl Seconds - description: Fallback subject-token lifetime when the Docker demo issuer - response omits expires_in. - refresh_margin_seconds: - type: integer - minimum: 0.0 - title: Refresh Margin Seconds - description: Seconds before subject-token expiry when the Docker refresher - issues a replacement token. - default: 60 + properties: {} additionalProperties: false type: object title: DockerWorkloadIdentityConfig - description: Docker-only subject token issuer configuration for workload identity. + description: Docker workload identity configuration. E2EJobExecutionProfile: properties: provider: @@ -19544,7 +19507,7 @@ components: type: string title: Error description: OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, - invalid_request, invalid_grant, invalid_scope, or invalid_target. + invalid_request, invalid_scope, or invalid_target. error_description: title: Error Description description: Human-readable ASCII text providing additional information diff --git a/openapi/openapi.yaml b/openapi/openapi.yaml index 33a0d3c022..d503c2233b 100644 --- a/openapi/openapi.yaml +++ b/openapi/openapi.yaml @@ -126,6 +126,7 @@ paths: type: string enum: - urn:ietf:params:oauth:token-type:jwt + - urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof description: Token type identifier for the subject token. requested_token_type: type: string @@ -136,6 +137,18 @@ paths: audience: type: string description: Requested audience for the issued access token. + resource: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token_type: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. scope: type: string description: Space-separated scopes requested for the issued access @@ -10443,9 +10456,7 @@ components: - $ref: '#/components/schemas/DockerJobNetworkConfig' description: Docker networking configuration workload_identity: - allOf: - - $ref: '#/components/schemas/DockerWorkloadIdentityConfig' - description: Docker workload identity subject-token issuer configuration. + $ref: '#/components/schemas/DockerWorkloadIdentityConfig' type: object title: DockerJobExecutionProfileConfig description: Configuration for Docker Job execution profile. @@ -10516,59 +10527,11 @@ components: - mount_path title: DockerVolumeMount DockerWorkloadIdentityConfig: - properties: - enabled: - title: Enabled - description: Enable Docker workload identity token-file injection. Defaults - to auth.oidc.workload_token_exchange_enabled. - type: boolean - token_endpoint: - title: Token Endpoint - description: OAuth token endpoint used by the Docker demo issuer. Defaults - to auth.oidc.token_endpoint. - type: string - client_id: - title: Client Id - description: OAuth client ID used by the Docker demo issuer. Defaults to - auth.oidc.workload_client_id or auth.oidc.client_id. - type: string - client_secret: - format: password - title: Client Secret - description: OAuth client secret for the Docker demo issuer. - writeOnly: true - type: string - username: - title: Username - description: Username for the Docker demo issuer password grant. - type: string - password_env_var: - type: string - title: Password Env Var - description: Controller environment variable that contains the Docker demo - issuer password grant shared secret. - default: AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD - scope: - title: Scope - description: OAuth scope for the Docker demo issuer. - type: string - subject_token_ttl_seconds: - type: integer - minimum: 1.0 - title: Subject Token Ttl Seconds - description: Fallback subject-token lifetime when the Docker demo issuer - response omits expires_in. - refresh_margin_seconds: - type: integer - minimum: 0.0 - title: Refresh Margin Seconds - description: Seconds before subject-token expiry when the Docker refresher - issues a replacement token. - default: 60 + properties: {} additionalProperties: false type: object title: DockerWorkloadIdentityConfig - description: Docker-only subject token issuer configuration for workload identity. + description: Docker workload identity configuration. E2EJobExecutionProfile: properties: provider: @@ -19544,7 +19507,7 @@ components: type: string title: Error description: OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, - invalid_request, invalid_grant, invalid_scope, or invalid_target. + invalid_request, invalid_scope, or invalid_target. error_description: title: Error Description description: Human-readable ASCII text providing additional information diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py index 2e9fd8a332..bdffa3cf54 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/auth/workload_exchange.py @@ -16,7 +16,14 @@ from urllib.parse import urlparse import httpx -from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.constants import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE as _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, +) +from nemo_platform_plugin.client.constants import ( + JWT_WORKLOAD_SUBJECT_TOKEN_TYPE, + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + subject_token_type_for_exchange, +) from nemo_platform_ext.auth.token_provider import DEFAULT_REFRESH_MARGIN_SECONDS, TokenSet from nemo_platform_ext.client.tls import client_verify_from_env @@ -24,7 +31,8 @@ logger = logging.getLogger(__name__) TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" -JWT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" +JWT_TOKEN_TYPE = JWT_WORKLOAD_SUBJECT_TOKEN_TYPE +DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE ACCESS_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token" @@ -87,7 +95,7 @@ def token_exchange_grant( "grant_type": TOKEN_EXCHANGE_GRANT_TYPE, "client_id": client_id, "subject_token": subject_token, - "subject_token_type": JWT_TOKEN_TYPE, + "subject_token_type": subject_token_type_for_exchange(subject_token), "requested_token_type": ACCESS_TOKEN_TYPE, } if audience: diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/jobs/__init__.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/jobs/__init__.py index e0e6cfb653..cb367ab453 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/jobs/__init__.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/jobs/__init__.py @@ -400,7 +400,12 @@ def list_execution_profiles_jobs( columns: OutputColumnsOption = None, stream: StreamOutputOption = False, ) -> None: - """Get all currently configured execution profiles.""" + """Get all currently configured execution profiles. + + Returns the capability-filtered merge from jobs config. In local standalone the + controller may prune the shared list further after registry boot; in split + topologies the API advertises its own merge result (not controller process + memory).""" state: CLIContext = ctx.obj output_format = state.get_output_format(output_format) validate_stream_output_format(output_format, stream) diff --git a/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py b/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py index f64f226776..475ee61c64 100644 --- a/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py +++ b/packages/nemo_platform_ext/tests/auth/test_workload_exchange.py @@ -3,14 +3,19 @@ """Tests for RFC 8693 workload identity token exchange.""" +import ast import json import time from base64 import urlsafe_b64encode +from pathlib import Path from unittest.mock import MagicMock, patch +import httpx +import nemo_platform_ext.auth.workload_exchange as workload_exchange_module import pytest from nemo_platform_ext.auth.workload_exchange import ( ACCESS_TOKEN_TYPE, + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, JWT_TOKEN_TYPE, TOKEN_EXCHANGE_GRANT_TYPE, WorkloadTokenExchangeError, @@ -22,6 +27,19 @@ from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +def test_workload_exchange_module_has_no_nmp_common_dependency(): + assert workload_exchange_module.__file__ is not None + tree = ast.parse(Path(workload_exchange_module.__file__).read_text(encoding="utf-8")) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + imported = [alias.name for alias in node.names] + elif isinstance(node, ast.ImportFrom): + imported = [node.module or ""] + else: + continue + assert not any(name == "nmp.common" or name.startswith("nmp.common.") for name in imported) + + def _make_jwt(claims: dict) -> str: header = {"alg": "RS256", "typ": "JWT"} h = urlsafe_b64encode(json.dumps(header).encode()).rstrip(b"=").decode() @@ -74,6 +92,19 @@ def test_token_exchange_grant_sends_rfc8693_request(mock_post): } +@patch("nemo_platform_ext.auth.workload_exchange.httpx.post") +def test_token_exchange_grant_uses_docker_opaque_subject_token_type(mock_post): + mock_post.return_value = httpx.Response(200, json={"access_token": "exchanged-token", "expires_in": 300}) + + token_exchange_grant( + token_endpoint="https://idp.example.com/token", + client_id="nemo-platform-workload", + subject_token="nmp_obo_v1.delegation.secret", + ) + + assert mock_post.call_args.kwargs["data"]["subject_token_type"] == DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE + + @patch("nemo_platform_ext.auth.workload_exchange.httpx.post") def test_token_exchange_grant_rejects_http_non_loopback_endpoint_before_sending_subject_token(mock_post): with pytest.raises(ValueError, match="must use HTTPS"): diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py index 65e994b9c4..e938012bd7 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/client.py @@ -38,6 +38,7 @@ StaticToken, TokenProvider, ) +from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR from nemo_platform_plugin.client.errors import ( NemoResponseValidationError, NemoTransportError, @@ -72,6 +73,59 @@ logger = logging.getLogger(__name__) DEFAULT_TIMEOUT = 60.0 +_AUTHORIZATION_HEADER = "Authorization" +_PRINCIPAL_ID_HEADER = "X-NMP-Principal-Id" + + +def _has_header(headers: Mapping[str, str] | None, name: str) -> bool: + if not headers: + return False + normalized = name.lower() + return any(header.lower() == normalized for header in headers) + + +@overload +def _resolve_implicit_workload_auth( + *, + base_url: str, + auth: TokenProvider | str | None, + default_headers: Mapping[str, str] | None, + allow_env_bootstrap: bool, +) -> TokenProvider | str | None: ... + + +@overload +def _resolve_implicit_workload_auth( + *, + base_url: str, + auth: TokenProvider | AsyncTokenProvider | str | None, + default_headers: Mapping[str, str] | None, + allow_env_bootstrap: bool, +) -> TokenProvider | AsyncTokenProvider | str | None: ... + + +def _resolve_implicit_workload_auth( + *, + base_url: str, + auth: TokenProvider | AsyncTokenProvider | str | None, + default_headers: Mapping[str, str] | None, + allow_env_bootstrap: bool, +) -> TokenProvider | AsyncTokenProvider | str | None: + if auth is not None or not allow_env_bootstrap: + return auth + if _has_header(default_headers, _AUTHORIZATION_HEADER) or _has_header(default_headers, _PRINCIPAL_ID_HEADER): + return None + + subject_token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR) + if not subject_token_file: + return None + + from nemo_platform_plugin.client.oidc_factory import resolve_workload_exchange_provider + + return resolve_workload_exchange_provider( + base_url=base_url, + subject_token_file=Path(subject_token_file), + ) @cache @@ -420,6 +474,12 @@ def __init__( defers to the transport's timeout, giving one we build ourselves :data:`DEFAULT_TIMEOUT`; ``httpx.Timeout(None)`` waits indefinitely. """ + auth = _resolve_implicit_workload_auth( + base_url=base_url, + auth=auth, + default_headers=default_headers, + allow_env_bootstrap=http_client is None, + ) super().__init__( base_url=base_url, workspace=workspace, @@ -671,6 +731,12 @@ def __init__( url_resolver: Callable[[str], str | httpx.URL] | None = None, ) -> None: """Create a client. See :meth:`NemoClient.__init__` for *timeout*.""" + auth = _resolve_implicit_workload_auth( + base_url=base_url, + auth=auth, + default_headers=default_headers, + allow_env_bootstrap=http_client is None, + ) super().__init__( base_url=base_url, workspace=workspace, @@ -911,8 +977,7 @@ def _client_from_config( """Shared implementation for NemoClient.from_config / AsyncNemoClient.from_config.""" from nemo_platform_plugin.client.config.config import Config from nemo_platform_plugin.client.config.models import ConfigParams, OAuthUser - from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR - from nemo_platform_plugin.client.oidc_factory import resolve_oidc_provider, resolve_workload_exchange_provider + from nemo_platform_plugin.client.oidc_factory import resolve_oidc_provider resolved_path = Path(config_path) if isinstance(config_path, str) else config_path overrides: ConfigParams | None = None @@ -928,13 +993,9 @@ def _client_from_config( auth: TokenProvider | str | None = None workload_identity_token_file = os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR) + use_implicit_workload_auth = bool(workload_identity_token_file and not explicit_access_token) - if workload_identity_token_file and not explicit_access_token: - auth = resolve_workload_exchange_provider( - base_url=str(ctx.cluster.base_url), - subject_token_file=Path(workload_identity_token_file), - ) - elif isinstance(ctx.user, OAuthUser): + if not use_implicit_workload_auth and isinstance(ctx.user, OAuthUser): auth = resolve_oidc_provider( base_url=str(ctx.cluster.base_url), context_name=ctx.context_name, @@ -944,7 +1005,7 @@ def _client_from_config( config_path=actual_config_path, explicit_access_token=explicit_access_token, ) - elif ctx.user: + elif not use_implicit_workload_auth and ctx.user: client_config = ctx.user.get_client_config() raw_headers = client_config.get("default_headers") if isinstance(raw_headers, dict): diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py index 6b9ea83424..d39f09b665 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/constants.py @@ -6,7 +6,17 @@ import os WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE" +JWT_WORKLOAD_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" +DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = "urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof" +OPAQUE_DOCKER_PROOF_PREFIX = "nmp_obo_v1" def is_workload_identity_token_file_set() -> bool: return bool(os.environ.get(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR)) + + +def subject_token_type_for_exchange(subject_token: str) -> str: + """Return the RFC 8693 subject_token_type for a workload identity subject token.""" + if subject_token.startswith(f"{OPAQUE_DOCKER_PROOF_PREFIX}."): + return DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE + return JWT_WORKLOAD_SUBJECT_TOKEN_TYPE diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py index ed396993a3..1319655690 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client/oidc.py @@ -30,7 +30,14 @@ from urllib.parse import urlparse import httpx -from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.constants import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE as _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, +) +from nemo_platform_plugin.client.constants import ( + JWT_WORKLOAD_SUBJECT_TOKEN_TYPE, + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + subject_token_type_for_exchange, +) from nemo_platform_plugin.client.tls import client_verify_from_env logger = logging.getLogger(__name__) @@ -124,7 +131,8 @@ def generate_unsigned_jwt( DEFAULT_OAUTH_SCOPES = "openid profile email offline_access" TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" -JWT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" +JWT_TOKEN_TYPE = JWT_WORKLOAD_SUBJECT_TOKEN_TYPE +DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE ACCESS_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token" @@ -309,7 +317,7 @@ def token_exchange_grant( "grant_type": TOKEN_EXCHANGE_GRANT_TYPE, "client_id": client_id, "subject_token": subject_token, - "subject_token_type": JWT_TOKEN_TYPE, + "subject_token_type": subject_token_type_for_exchange(subject_token), "requested_token_type": ACCESS_TOKEN_TYPE, } if audience: diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py index a719dc9ed3..18cedfacd2 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/client_provider.py @@ -90,7 +90,8 @@ def get_nemo_client( ``NMP_PRINCIPAL`` from the environment. """ headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return NemoClient(base_url=_base_url(), default_headers=headers or None) + base_url = _base_url() + return NemoClient(base_url=base_url, default_headers=headers or None) def get_async_nemo_client( @@ -105,4 +106,5 @@ def get_async_nemo_client( ``NMP_PRINCIPAL`` from the environment. """ headers = _build_headers(as_service=as_service, internal=internal, on_behalf_of=on_behalf_of) - return AsyncNemoClient(base_url=_base_url(), default_headers=headers or None) + base_url = _base_url() + return AsyncNemoClient(base_url=base_url, default_headers=headers or None) diff --git a/packages/nemo_platform_plugin/tests/test_client_auth.py b/packages/nemo_platform_plugin/tests/test_client_auth.py index ddb0c81c04..510153f69d 100644 --- a/packages/nemo_platform_plugin/tests/test_client_auth.py +++ b/packages/nemo_platform_plugin/tests/test_client_auth.py @@ -5,11 +5,14 @@ from __future__ import annotations +import ast import asyncio import time +from pathlib import Path from unittest.mock import patch import httpx +import nemo_platform_plugin.client.oidc as oidc_module import pytest import respx import yaml @@ -26,6 +29,7 @@ from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR from nemo_platform_plugin.client.oidc import ( ACCESS_TOKEN_TYPE, + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, JWT_TOKEN_TYPE, TOKEN_EXCHANGE_GRANT_TYPE, NMPOIDCConfig, @@ -44,6 +48,24 @@ get_nemo_client, ) + +def _assert_no_nmp_common_import(module_file: str | None) -> None: + assert module_file is not None + tree = ast.parse(Path(module_file).read_text(encoding="utf-8")) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + imported = [alias.name for alias in node.names] + elif isinstance(node, ast.ImportFrom): + imported = [node.module or ""] + else: + continue + assert not any(name == "nmp.common" or name.startswith("nmp.common.") for name in imported) + + +def test_oidc_module_has_no_nmp_common_dependency(): + _assert_no_nmp_common_import(oidc_module.__file__) + + # --------------------------------------------------------------------------- # StaticToken # --------------------------------------------------------------------------- @@ -129,6 +151,52 @@ def test_no_auth_no_header(self): assert route.called assert "Authorization" not in route.calls[0].request.headers + def test_constructor_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + provider = object() + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + + with patch( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + return_value=provider, + ) as resolve_provider: + client = NemoClient(base_url="https://nemo.example.com") + + assert client._auth is provider + resolve_provider.assert_called_once_with( + base_url="https://nemo.example.com", + subject_token_file=subject_token_file, + ) + + def test_constructor_workload_exchange_does_not_override_authorization_header(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + + with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider: + client = NemoClient( + base_url="https://nemo.example.com", + default_headers={"Authorization": "Bearer explicit-token"}, + ) + + assert client._auth is None + resolve_provider.assert_not_called() + + def test_constructor_workload_exchange_does_not_override_principal_header(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + + with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider: + client = NemoClient( + base_url="https://nemo.example.com", + default_headers={"X-NMP-Principal-Id": "service:jobs"}, + ) + + assert client._auth is None + resolve_provider.assert_not_called() + # --------------------------------------------------------------------------- # AsyncNemoClient auth parameter @@ -175,6 +243,24 @@ def test_sync_provider_works_in_async_client(self): assert route.called assert route.calls[0].request.headers["Authorization"] == "Bearer sync-token" + def test_constructor_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + provider = object() + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + + with patch( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + return_value=provider, + ) as resolve_provider: + client = AsyncNemoClient(base_url="https://nemo.example.com") + + assert client._auth is provider + resolve_provider.assert_called_once_with( + base_url="https://nemo.example.com", + subject_token_file=subject_token_file, + ) + # --------------------------------------------------------------------------- # OIDCTokenProvider @@ -335,6 +421,20 @@ def test_token_exchange_grant_sends_rfc8693_request(self, monkeypatch): verify=True, ) + def test_token_exchange_grant_uses_docker_opaque_subject_token_type(self, monkeypatch): + monkeypatch.delenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, raising=False) + + with patch("nemo_platform_plugin.client.oidc.httpx.post") as mock_post: + mock_post.return_value = httpx.Response(200, json={"access_token": "exchanged-token", "expires_in": 300}) + + token_exchange_grant( + token_endpoint="https://idp.example.com/token", + client_id="nemo-platform-workload", + subject_token="nmp_obo_v1.delegation.secret", + ) + + assert mock_post.call_args.kwargs["data"]["subject_token_type"] == DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE + def test_token_exchange_grant_uses_nemo_scoped_ca_bundle(self, monkeypatch): monkeypatch.setenv(NMP_CLIENT_SSL_CERT_FILE_ENVVAR, "/tmp/nemo-ca.pem") @@ -778,6 +878,58 @@ def test_returns_async_client(self): client = get_async_nemo_client(as_service="test-svc") assert isinstance(client, AsyncNemoClient) + def test_sync_client_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + provider = object() + monkeypatch.setenv("NMP_BASE_URL", "https://nemo.example.com") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with patch( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + return_value=provider, + ) as resolve_provider: + client = get_nemo_client() + + assert client._auth is provider + resolve_provider.assert_called_once_with( + base_url="https://nemo.example.com", + subject_token_file=subject_token_file, + ) + + def test_async_client_uses_workload_exchange_provider_from_env(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + provider = object() + monkeypatch.setenv("NMP_BASE_URL", "https://nemo.example.com") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with patch( + "nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider", + return_value=provider, + ) as resolve_provider: + client = get_async_nemo_client() + + assert client._auth is provider + resolve_provider.assert_called_once_with( + base_url="https://nemo.example.com", + subject_token_file=subject_token_file, + ) + + def test_workload_exchange_provider_does_not_override_explicit_service(self, monkeypatch, tmp_path): + subject_token_file = tmp_path / "workload-token" + subject_token_file.write_text("subject-token\n", encoding="utf-8") + monkeypatch.setenv(WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, str(subject_token_file)) + monkeypatch.delenv("NMP_PRINCIPAL", raising=False) + + with patch("nemo_platform_plugin.client.oidc_factory.resolve_workload_exchange_provider") as resolve_provider: + client = get_nemo_client(as_service="jobs") + + assert client._auth is None + resolve_provider.assert_not_called() + # --------------------------------------------------------------------------- # Security: token repr and endpoint validation diff --git a/packages/nmp_common/src/nmp/common/auth/__init__.py b/packages/nmp_common/src/nmp/common/auth/__init__.py index f90a1da3da..d2e53e99bc 100644 --- a/packages/nmp_common/src/nmp/common/auth/__init__.py +++ b/packages/nmp_common/src/nmp/common/auth/__init__.py @@ -24,6 +24,25 @@ from .models import NMP_PRINCIPAL_ENVVAR, AuthContext, Principal from .permissions import ALL_WORKSPACES, compute_accessible_workspaces from .tasks import principal_from_env +from .workload_delegations import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, + JWT_WORKLOAD_SUBJECT_TOKEN_TYPE, + OPAQUE_DOCKER_PROOF_PREFIX, + WORKLOAD_DELEGATION_ENTITY_TYPE, + InvalidWorkloadProofTokenError, + ParsedOpaqueDockerProofToken, + WorkloadDelegationConflictError, + WorkloadDelegationEntity, + WorkloadDelegationError, + WorkloadDelegationStore, + WorkloadDelegationValidationError, + create_opaque_docker_proof_token, + docker_delegation_name, + parse_opaque_docker_proof_token, + reference_delegation_name, + subject_token_type_for_exchange, + verify_opaque_docker_proof_token_hash, +) # Testing utilities are NOT exported here to avoid importing dev dependencies (respx) # at runtime. Import directly from nmp.common.auth.testing when needed in tests. @@ -41,14 +60,31 @@ "ACCESS_KEY_JWKS_PATH", "ACCESS_KEY_TOKEN_TYPE", "AccessKeyIssuerService", + "DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE", + "JWT_WORKLOAD_SUBJECT_TOKEN_TYPE", + "OPAQUE_DOCKER_PROOF_PREFIX", + "WORKLOAD_DELEGATION_ENTITY_TYPE", + "InvalidWorkloadProofTokenError", "NMP_PRINCIPAL_ENVVAR", + "ParsedOpaqueDockerProofToken", "Principal", + "WorkloadDelegationConflictError", + "WorkloadDelegationEntity", + "WorkloadDelegationError", + "WorkloadDelegationStore", + "WorkloadDelegationValidationError", "auth_as_service", "auth_client_context", "build_service_principal_headers", + "create_opaque_docker_proof_token", + "docker_delegation_name", "principal_from_env", + "parse_opaque_docker_proof_token", + "reference_delegation_name", + "subject_token_type_for_exchange", "compute_accessible_workspaces", "get_auth_client", "get_principal_auth_headers", "validate_access_key_token", + "verify_opaque_docker_proof_token_hash", ] diff --git a/packages/nmp_common/src/nmp/common/auth/jwks.py b/packages/nmp_common/src/nmp/common/auth/jwks.py index b47c592f40..44e8eda346 100644 --- a/packages/nmp_common/src/nmp/common/auth/jwks.py +++ b/packages/nmp_common/src/nmp/common/auth/jwks.py @@ -65,6 +65,10 @@ async def get_signing_key_from_jwt(self, token: str) -> Any: refreshed_jwks = await self._refresh_jwks_for_unknown_kid() return signing_jwk_from_jwks(token, refreshed_jwks) + async def get_jwks(self) -> dict[str, Any]: + jwks, _ = await self._fetch_jwks() + return jwks + def clear_cache(self) -> None: self._jwks = None self._jwks_cache_time = 0.0 diff --git a/packages/nmp_common/src/nmp/common/auth/jwt.py b/packages/nmp_common/src/nmp/common/auth/jwt.py index 4f13b95cd7..0fcb9e2d0c 100644 --- a/packages/nmp_common/src/nmp/common/auth/jwt.py +++ b/packages/nmp_common/src/nmp/common/auth/jwt.py @@ -6,7 +6,7 @@ import logging import time from dataclasses import dataclass -from typing import Optional +from typing import Any, Optional import httpx import jwt @@ -27,6 +27,14 @@ _DISCOVERY_CACHE_TTL = 3600 # 1 hour +@dataclass +class ActorClaims: + """Validated RFC 8693 actor claims.""" + + subject: str + groups: list[str] + + @dataclass class TokenClaims: """Validated token claims.""" @@ -36,6 +44,7 @@ class TokenClaims: groups: list[str] scopes: list[str] raw_claims: dict + actor: Optional[ActorClaims] = None class UnsignedJWTRejectedError(Exception): @@ -70,6 +79,21 @@ async def _discover_oidc_config(self) -> dict: self._discovery_cache_time = now return self._discovery_cache + async def jwks_uri(self) -> str: + """Return the configured or discovered OIDC JWKS URI.""" + if self.config.oidc.jwks_uri: + return self.config.oidc.jwks_uri + discovery = await self._discover_oidc_config() + jwks_uri = discovery.get("jwks_uri") + if not isinstance(jwks_uri, str) or not jwks_uri: + raise jwt.InvalidTokenError("OIDC discovery did not include jwks_uri") + return jwks_uri + + async def jwks(self) -> dict[str, Any]: + """Return the cached-or-fetched OIDC JWKS document.""" + jwks_client = await self._get_jwks_client() + return await jwks_client.get_jwks() + async def _get_jwks_client(self) -> AsyncJWKSClient: """Get or create JWKS client for token validation. @@ -81,11 +105,7 @@ async def _get_jwks_client(self) -> AsyncJWKSClient: if self._jwks_client: return self._jwks_client - jwks_uri = self.config.oidc.jwks_uri - if not jwks_uri: - discovery = await self._discover_oidc_config() - jwks_uri = discovery["jwks_uri"] - + jwks_uri = await self.jwks_uri() self._jwks_client = AsyncJWKSClient(jwks_uri, lifespan=_JWKS_CACHE_LIFESPAN) return self._jwks_client @@ -96,17 +116,11 @@ def _extract_token_claims(self, claims: dict) -> Optional[TokenClaims]: logger.warning("Token is missing a valid subject claim") return None - email = claims.get(self.config.oidc.email_claim) + email_value = claims.get(self.config.oidc.email_claim) + email = email_value if isinstance(email_value, str) else None - groups: list[str] = [] - for claim_name in [self.config.oidc.groups_claim, "cognito:groups"]: - if claim_name in claims: - groups_value = claims[claim_name] - if isinstance(groups_value, str): - groups = [g.strip() for g in groups_value.split(",")] - elif isinstance(groups_value, list): - groups = groups_value - break + groups = self._extract_groups_from_claims(claims) + actor = self._extract_actor_claims(claims) scopes: list[str] = [] scope_value = claims.get("scope") or claims.get("scp") @@ -130,6 +144,39 @@ def _extract_token_claims(self, claims: dict) -> Optional[TokenClaims]: groups=groups, scopes=scopes, raw_claims=claims, + actor=actor, + ) + + def _extract_groups_from_claims(self, claims: dict) -> list[str]: + """Extract normalized groups from configured or provider-specific claims.""" + for claim_name in [self.config.oidc.groups_claim, "cognito:groups"]: + if claim_name not in claims: + continue + value = claims[claim_name] + if isinstance(value, str): + return [group.strip() for group in value.split(",") if group.strip()] + if isinstance(value, list): + return [str(group).strip() for group in value if str(group).strip()] + return [] + return [] + + def _extract_actor_claims(self, claims: dict) -> Optional[ActorClaims]: + """Extract RFC 8693 act claims when a valid actor subject is present.""" + actor_claims = claims.get("act") + if not isinstance(actor_claims, dict): + return None + + actor_subject = actor_claims.get("sub") + if not isinstance(actor_subject, str): + return None + + actor_subject = actor_subject.strip() + if not actor_subject: + return None + + return ActorClaims( + subject=actor_subject, + groups=self._extract_groups_from_claims(actor_claims), ) async def validate_token(self, token: str) -> Optional[TokenClaims]: diff --git a/packages/nmp_common/src/nmp/common/auth/token_resolver.py b/packages/nmp_common/src/nmp/common/auth/token_resolver.py index 1dffa0b21b..071f5634a7 100644 --- a/packages/nmp_common/src/nmp/common/auth/token_resolver.py +++ b/packages/nmp_common/src/nmp/common/auth/token_resolver.py @@ -15,6 +15,26 @@ ResolvedTokenKind = Literal["access_key", "oidc_access_token", "workload_access_token", "workload_subject_token"] +def _direct_principal_from_claims(claims: TokenClaims) -> Principal: + return Principal( + id=claims.subject, + email=claims.email, + groups=claims.groups, + ) + + +def _workload_access_principal_from_claims(claims: TokenClaims) -> Principal: + if claims.actor is None: + return _direct_principal_from_claims(claims) + return Principal( + id=claims.actor.subject, + groups=claims.actor.groups, + on_behalf_of=claims.subject, + on_behalf_of_email=claims.email, + on_behalf_of_groups=claims.groups, + ) + + @dataclass(frozen=True) class ResolvedBearerToken: claims: TokenClaims @@ -22,11 +42,9 @@ class ResolvedBearerToken: @property def principal(self) -> Principal: - return Principal( - id=self.claims.subject, - email=self.claims.email, - groups=self.claims.groups, - ) + if self.token_kind == "workload_access_token": + return _workload_access_principal_from_claims(self.claims) + return _direct_principal_from_claims(self.claims) @property def scopes(self) -> list[str]: diff --git a/packages/nmp_common/src/nmp/common/auth/workload_delegations.py b/packages/nmp_common/src/nmp/common/auth/workload_delegations.py new file mode 100644 index 0000000000..97930776c8 --- /dev/null +++ b/packages/nmp_common/src/nmp/common/auth/workload_delegations.py @@ -0,0 +1,349 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Common workload delegation primitives for job OBO token exchange.""" + +from __future__ import annotations + +import base64 +import hashlib +import hmac +import json +import secrets +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import ClassVar + +from nmp.common.entities import SYSTEM_WORKSPACE, EntityBase, EntityClient, EntityConflictError, EntityNotFoundError + +from .models import AuthContext + +WORKLOAD_DELEGATION_ENTITY_TYPE = "workload_delegation" +JWT_WORKLOAD_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" +DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = "urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof" +OPAQUE_DOCKER_PROOF_PREFIX = "nmp_obo_v1" + +_OPAQUE_DOCKER_PROOF_SECRET_BYTES = 32 +_OPAQUE_DOCKER_PROOF_HASH_PREFIX = "v1:sha256:" +_DELEGATION_HASH_LENGTH = 48 + + +class WorkloadDelegationError(Exception): + """Base exception for workload delegation failures.""" + + +class WorkloadDelegationValidationError(WorkloadDelegationError): + """Raised when a delegation record is malformed or unusable.""" + + +class WorkloadDelegationConflictError(WorkloadDelegationError): + """Raised when an active delegation already exists for the same name.""" + + +class InvalidWorkloadProofTokenError(WorkloadDelegationError): + """Raised when a Docker opaque workload proof token is malformed.""" + + +@dataclass(frozen=True) +class ParsedOpaqueDockerProofToken: + """Parsed Docker opaque workload proof token.""" + + delegation_name: str + secret: bytes + + +class WorkloadDelegationEntity(EntityBase): + """Auth-owned delegation record used to mint delegated workload tokens.""" + + __entity_type__: ClassVar[str] = WORKLOAD_DELEGATION_ENTITY_TYPE + + workload_subject: str + workload_audience: str + workload_workspace: str + job_id: str + attempt_id: str + step_id: str + auth_context: AuthContext + bound_reference_name: str | None = None + bound_reference_value: str | None = None + opaque_subject_token_hash: str | None = None + expires_at: datetime + revoked_at: datetime | None = None + + def is_expired(self, *, now: datetime | None = None) -> bool: + """Return whether this delegation is expired at the supplied time.""" + effective_now = _as_aware_utc(now or datetime.now(timezone.utc)) + return _as_aware_utc(self.expires_at) <= effective_now + + def is_active(self, *, now: datetime | None = None) -> bool: + """Return whether this delegation can still be used.""" + return self.revoked_at is None and not self.is_expired(now=now) + + +def docker_delegation_name(*, workload_workspace: str, job_id: str, attempt_id: str, step_id: str) -> str: + """Return the deterministic Docker delegation entity name.""" + return _delegation_hash_name( + "job", + [ + _require_non_empty(workload_workspace, "workload_workspace"), + _require_non_empty(job_id, "job_id"), + _require_non_empty(attempt_id, "attempt_id"), + _require_non_empty(step_id, "step_id"), + ], + ) + + +def reference_delegation_name( + *, + workload_audience: str, + workload_subject: str, + bound_reference_name: str, + bound_reference_value: str, +) -> str: + """Return the deterministic verified-reference delegation entity name.""" + return _delegation_hash_name( + "ref", + [ + _require_non_empty(workload_audience, "workload_audience"), + _require_non_empty(workload_subject, "workload_subject"), + _require_non_empty(bound_reference_name, "bound_reference_name"), + _require_non_empty(bound_reference_value, "bound_reference_value"), + ], + ) + + +def create_opaque_docker_proof_token(delegation_name: str) -> tuple[str, str]: + """Create an opaque Docker workload proof token and its stored secret hash.""" + secret = secrets.token_bytes(_OPAQUE_DOCKER_PROOF_SECRET_BYTES) + token = ".".join( + [ + OPAQUE_DOCKER_PROOF_PREFIX, + _b64url_encode(_require_non_empty(delegation_name, "delegation_name").encode("utf-8")), + _b64url_encode(secret), + ] + ) + return token, _opaque_docker_proof_token_hash(secret) + + +def parse_opaque_docker_proof_token(token: str) -> ParsedOpaqueDockerProofToken: + """Parse a Docker opaque workload proof token envelope.""" + parts = token.split(".") + if len(parts) != 3 or parts[0] != OPAQUE_DOCKER_PROOF_PREFIX: + raise InvalidWorkloadProofTokenError("invalid Docker opaque workload proof token envelope") + + try: + delegation_name = _b64url_decode(parts[1]).decode("utf-8") + secret = _b64url_decode(parts[2]) + except (UnicodeDecodeError, ValueError) as exc: + raise InvalidWorkloadProofTokenError("invalid Docker opaque workload proof token encoding") from exc + + if not delegation_name: + raise InvalidWorkloadProofTokenError("Docker opaque workload proof token is missing a delegation name") + if len(secret) != _OPAQUE_DOCKER_PROOF_SECRET_BYTES: + raise InvalidWorkloadProofTokenError("Docker opaque workload proof token secret has invalid length") + + return ParsedOpaqueDockerProofToken(delegation_name=delegation_name, secret=secret) + + +def verify_opaque_docker_proof_token_hash(secret: bytes, expected_hash: str) -> bool: + """Constant-time check for a parsed Docker opaque proof token secret.""" + if not expected_hash.startswith(_OPAQUE_DOCKER_PROOF_HASH_PREFIX): + return False + return hmac.compare_digest(_opaque_docker_proof_token_hash(secret), expected_hash) + + +def subject_token_type_for_exchange(subject_token: str) -> str: + """Return the RFC 8693 subject_token_type for a workload identity subject token.""" + if subject_token.startswith(f"{OPAQUE_DOCKER_PROOF_PREFIX}."): + return DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE + return JWT_WORKLOAD_SUBJECT_TOKEN_TYPE + + +class WorkloadDelegationStore: + """Identity-based store wrapper for workload delegation records.""" + + def __init__(self, entity_client: EntityClient): + self._entity_client = entity_client + + async def register( + self, + entity: WorkloadDelegationEntity, + *, + expected_db_version: int | None = None, + require_opaque_subject_token_hash: bool = False, + ) -> WorkloadDelegationEntity: + """Create a delegation row, or replace an expired/revoked row with the same name.""" + _validate_delegation(entity, require_opaque_subject_token_hash=require_opaque_subject_token_hash) + + try: + return await self._entity_client.add(entity) + except EntityConflictError as exc: + existing = await self.get(entity.name) + if existing is None: + raise WorkloadDelegationConflictError(f"delegation '{entity.name}' already exists") from exc + if existing.is_active(): + raise WorkloadDelegationConflictError(f"active delegation '{entity.name}' already exists") from exc + if expected_db_version is not None and existing.db_version != expected_db_version: + raise WorkloadDelegationConflictError( + f"delegation '{entity.name}' db_version {existing.db_version} does not match " + f"expected db_version {expected_db_version}" + ) from exc + replacement = _copy_delegation_payload(existing, entity) + return await self._entity_client.update(replacement) + + async def get(self, name: str) -> WorkloadDelegationEntity | None: + """Fetch one delegation row by its deterministic name.""" + try: + return await self._entity_client.get(WorkloadDelegationEntity, name, workspace=SYSTEM_WORKSPACE) + except EntityNotFoundError: + return None + + async def update( + self, + entity: WorkloadDelegationEntity, + *, + expected_db_version: int | None = None, + ) -> WorkloadDelegationEntity: + """Update a delegation row using the db_version from the latest fetched row.""" + _validate_delegation(entity) + existing = await self.get(entity.name) + if existing is None: + raise EntityNotFoundError(f"Entity '{entity.name}' not found in workspace '{SYSTEM_WORKSPACE}'") + if expected_db_version is not None and existing.db_version != expected_db_version: + raise WorkloadDelegationConflictError( + f"delegation '{entity.name}' db_version {existing.db_version} does not match " + f"expected db_version {expected_db_version}" + ) + updated = _copy_delegation_payload(existing, entity) + return await self._entity_client.update(updated) + + async def revoke(self, name: str, *, now: datetime | None = None) -> WorkloadDelegationEntity | None: + """Soft-revoke a delegation row by setting revoked_at.""" + existing = await self.get(name) + if existing is None: + return None + existing.revoked_at = _as_aware_utc(now or datetime.now(timezone.utc)) + return await self._entity_client.update(existing) + + +def _copy_delegation_payload( + target: WorkloadDelegationEntity, + source: WorkloadDelegationEntity, +) -> WorkloadDelegationEntity: + target.name = source.name + target.workspace = SYSTEM_WORKSPACE + target.project = source.project + target.workload_subject = source.workload_subject + target.workload_audience = source.workload_audience + target.workload_workspace = source.workload_workspace + target.job_id = source.job_id + target.attempt_id = source.attempt_id + target.step_id = source.step_id + target.auth_context = source.auth_context + target.bound_reference_name = source.bound_reference_name + target.bound_reference_value = source.bound_reference_value + target.opaque_subject_token_hash = source.opaque_subject_token_hash + target.expires_at = source.expires_at + target.revoked_at = source.revoked_at + return target + + +def _validate_delegation( + entity: WorkloadDelegationEntity, + *, + require_opaque_subject_token_hash: bool = False, +) -> None: + if entity.workspace != SYSTEM_WORKSPACE: + raise WorkloadDelegationValidationError("workload delegations must be stored in the system workspace") + + for field_name in ( + "name", + "workload_subject", + "workload_audience", + "workload_workspace", + "job_id", + "attempt_id", + "step_id", + ): + _require_non_empty(getattr(entity, field_name), field_name) + + if entity.is_expired(): + raise WorkloadDelegationValidationError("workload delegation expires_at must be in the future") + + has_reference_name = bool(entity.bound_reference_name) + has_reference_value = bool(entity.bound_reference_value) + if has_reference_name != has_reference_value: + raise WorkloadDelegationValidationError( + "bound_reference_name and bound_reference_value must be provided together" + ) + + if _is_docker_delegation(entity) and (has_reference_name or has_reference_value): + raise WorkloadDelegationValidationError("Docker workload delegations cannot use bound references") + + expected_name = _expected_delegation_name(entity) + if expected_name is None: + raise WorkloadDelegationValidationError( + "workload delegations must use a Docker or verified-reference lookup key" + ) + if entity.name != expected_name: + raise WorkloadDelegationValidationError("workload delegation name does not match its canonical lookup key") + + if require_opaque_subject_token_hash and not entity.opaque_subject_token_hash: + raise WorkloadDelegationValidationError("opaque Docker workload delegations require a stored token hash") + + +def _expected_delegation_name(entity: WorkloadDelegationEntity) -> str | None: + if _is_docker_delegation(entity): + return docker_delegation_name( + workload_workspace=entity.workload_workspace, + job_id=entity.job_id, + attempt_id=entity.attempt_id, + step_id=entity.step_id, + ) + + if entity.bound_reference_name and entity.bound_reference_value: + return reference_delegation_name( + workload_audience=entity.workload_audience, + workload_subject=entity.workload_subject, + bound_reference_name=entity.bound_reference_name, + bound_reference_value=entity.bound_reference_value, + ) + + return None + + +def _is_docker_delegation(entity: WorkloadDelegationEntity) -> bool: + return entity.name.startswith("job-") and entity.workload_subject == entity.name + + +def _opaque_docker_proof_token_hash(secret: bytes) -> str: + digest = hashlib.sha256(secret).digest() + return _OPAQUE_DOCKER_PROOF_HASH_PREFIX + _b64url_encode(digest) + + +def _require_non_empty(value: str, field_name: str) -> str: + if not isinstance(value, str) or not value: + raise WorkloadDelegationValidationError(f"{field_name} must be a non-empty string") + return value + + +def _delegation_hash_name(prefix: str, values: list[str]) -> str: + canonical = json.dumps(values, separators=(",", ":"), ensure_ascii=False) + return f"{prefix}-{hashlib.sha256(canonical.encode('utf-8')).hexdigest()[:_DELEGATION_HASH_LENGTH]}" + + +def _b64url_encode(value: bytes) -> str: + return base64.urlsafe_b64encode(value).decode("ascii").rstrip("=") + + +def _b64url_decode(value: str) -> bytes: + if not value: + raise ValueError("empty base64url value") + padding = "=" * (-len(value) % 4) + return base64.urlsafe_b64decode((value + padding).encode("ascii")) + + +def _as_aware_utc(value: datetime) -> datetime: + if value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) diff --git a/packages/nmp_common/tests/auth/test_jwt.py b/packages/nmp_common/tests/auth/test_jwt.py index e31f91f014..15455ae20e 100644 --- a/packages/nmp_common/tests/auth/test_jwt.py +++ b/packages/nmp_common/tests/auth/test_jwt.py @@ -12,7 +12,7 @@ from cryptography.hazmat.primitives.asymmetric import rsa from jwt.algorithms import RSAAlgorithm from nmp.common import http_clients -from nmp.common.auth.jwt import JWTValidator, TokenClaims, UnsignedJWTRejectedError +from nmp.common.auth.jwt import ActorClaims, JWTValidator, TokenClaims, UnsignedJWTRejectedError from nmp.common.config import AuthConfig from nmp.common.config.base import OIDCConfig @@ -79,6 +79,24 @@ def test_token_claims_with_none_email(self): assert claims.subject == "user123" assert claims.email is None + def test_token_claims_with_actor(self): + """Test TokenClaims with an RFC 8693 actor claim.""" + claims = TokenClaims( + subject="creator@example.com", + email="creator@example.com", + groups=["workspace-editors"], + scopes=[], + raw_claims={"sub": "creator@example.com"}, + actor=ActorClaims( + subject="system:serviceaccount:nemo-runs:job-runner", + groups=["system:serviceaccounts"], + ), + ) + + assert claims.actor is not None + assert claims.actor.subject == "system:serviceaccount:nemo-runs:job-runner" + assert claims.actor.groups == ["system:serviceaccounts"] + class TestOIDCConfigClaimDefaults: """Tests for issuer-based claim defaults.""" @@ -386,6 +404,92 @@ async def test_validate_token_success(self, jwt_validator): assert result.groups == ["admin", "users"] assert result.scopes == ["openid", "profile", "email"] + @pytest.mark.asyncio + async def test_validate_token_with_actor_claim(self, jwt_validator): + """Test token validation with an RFC 8693 act claim.""" + valid_claims = { + "sub": "creator@example.com", + "email": "creator@example.com", + "groups": "workspace-editors", + "act": { + "sub": "system:serviceaccount:nemo-runs:job-runner", + "groups": "system:serviceaccounts", + }, + "exp": int(time.time()) + 3600, + "iat": int(time.time()), + "aud": "test-audience", + "iss": "https://sso.example.com", + } + + with patch.object(jwt_validator, "_get_jwks_client") as mock_get_jwks: + mock_jwks = MagicMock() + mock_signing_key = MagicMock() + mock_signing_key.key = "test-key" + mock_jwks.get_signing_key_from_jwt = AsyncMock(return_value=mock_signing_key) + mock_get_jwks.return_value = mock_jwks + + with patch("jwt.decode", return_value=valid_claims): + result = await jwt_validator.validate_token("valid.token.here") + + assert result is not None + assert result.subject == "creator@example.com" + assert result.email == "creator@example.com" + assert result.groups == ["workspace-editors"] + assert result.actor == ActorClaims( + subject="system:serviceaccount:nemo-runs:job-runner", + groups=["system:serviceaccounts"], + ) + + @pytest.mark.asyncio + async def test_validate_token_ignores_actor_without_subject(self, jwt_validator): + """Test act is ignored unless act.sub is a non-empty string.""" + valid_claims = { + "sub": "creator@example.com", + "act": {"groups": "system:serviceaccounts"}, + "exp": int(time.time()) + 3600, + "iat": int(time.time()), + "aud": "test-audience", + "iss": "https://sso.example.com", + } + + with patch.object(jwt_validator, "_get_jwks_client") as mock_get_jwks: + mock_jwks = MagicMock() + mock_signing_key = MagicMock() + mock_signing_key.key = "test-key" + mock_jwks.get_signing_key_from_jwt = AsyncMock(return_value=mock_signing_key) + mock_get_jwks.return_value = mock_jwks + + with patch("jwt.decode", return_value=valid_claims): + result = await jwt_validator.validate_token("valid.token.here") + + assert result is not None + assert result.actor is None + + @pytest.mark.asyncio + async def test_validate_token_ignores_actor_with_whitespace_subject(self, jwt_validator): + """Test act is ignored when act.sub normalizes to an empty string.""" + valid_claims = { + "sub": "creator@example.com", + "act": {"sub": " ", "groups": "system:serviceaccounts"}, + "exp": int(time.time()) + 3600, + "iat": int(time.time()), + "aud": "test-audience", + "iss": "https://sso.example.com", + } + + with patch.object(jwt_validator, "_get_jwks_client") as mock_get_jwks: + mock_jwks = MagicMock() + mock_signing_key = MagicMock() + mock_signing_key.key = "test-key" + mock_jwks.get_signing_key_from_jwt = AsyncMock(return_value=mock_signing_key) + mock_get_jwks.return_value = mock_jwks + + with patch("jwt.decode", return_value=valid_claims): + result = await jwt_validator.validate_token("valid.token.here") + + assert result is not None + assert result.actor is None + @pytest.mark.asyncio async def test_validate_token_with_string_groups(self, jwt_validator): """Test token validation with comma-separated groups string.""" @@ -551,6 +655,50 @@ async def test_validate_token_uses_configured_jwks_uri(self, auth_config): lifespan=_JWKS_CACHE_LIFESPAN, ) + @pytest.mark.asyncio + async def test_jwks_uri_returns_configured_uri(self, auth_config): + """Configured JWKS URI is exposed without discovery.""" + auth_config.oidc.jwks_uri = "https://custom.example.com/jwks" + validator = JWTValidator(auth_config) + + with patch.object(validator, "_discover_oidc_config", new=AsyncMock()) as discover: + assert await validator.jwks_uri() == "https://custom.example.com/jwks" + + discover.assert_not_awaited() + + @pytest.mark.asyncio + async def test_jwks_uri_returns_discovered_uri(self, auth_config): + """Discovery JWKS URI is exposed when no explicit URI is configured.""" + auth_config.oidc.jwks_uri = None + validator = JWTValidator(auth_config) + + with patch.object( + validator, + "_discover_oidc_config", + new=AsyncMock(return_value={"jwks_uri": "https://sso.example.com/discovered-jwks"}), + ) as discover: + assert await validator.jwks_uri() == "https://sso.example.com/discovered-jwks" + + discover.assert_awaited_once() + + @pytest.mark.asyncio + async def test_jwks_uses_async_jwks_client_cache(self, auth_config): + """JWKS access reuses the validator's async JWKS client.""" + auth_config.oidc.jwks_uri = "https://custom.example.com/jwks" + validator = JWTValidator(auth_config) + jwks = {"keys": [{"kty": "RSA", "kid": "idp-key", "n": "modulus", "e": "AQAB"}]} + + with patch("nmp.common.auth.jwt.AsyncJWKSClient") as jwks_client_class: + jwks_client = MagicMock() + jwks_client.get_jwks = AsyncMock(return_value=jwks) + jwks_client_class.return_value = jwks_client + + assert await validator.jwks() == jwks + assert await validator.jwks() == jwks + + jwks_client_class.assert_called_once() + assert jwks_client.get_jwks.await_count == 2 + @pytest.mark.asyncio async def test_validate_token_fetches_jwks_with_async_client(self, auth_config, monkeypatch): """OIDC JWKS lookup must not use sync PyJWKClient in async validation.""" diff --git a/packages/nmp_common/tests/auth/test_middleware.py b/packages/nmp_common/tests/auth/test_middleware.py index 6a83f4cc81..5fe86c6eb2 100644 --- a/packages/nmp_common/tests/auth/test_middleware.py +++ b/packages/nmp_common/tests/auth/test_middleware.py @@ -12,7 +12,7 @@ from fastapi.testclient import TestClient from nmp.common.auth.client import AuthClient from nmp.common.auth.dependencies import get_auth_client -from nmp.common.auth.jwt import TokenClaims, UnsignedJWTRejectedError +from nmp.common.auth.jwt import ActorClaims, TokenClaims, UnsignedJWTRejectedError from nmp.common.auth.middleware import BYPASS_PREFIXES, HEALTH_ENDPOINTS, PUBLIC_GET_PATHS, AuthorizationMiddleware from nmp.common.auth.models import Principal from nmp.common.auth.token_resolver import ResolvedBearerToken @@ -556,10 +556,17 @@ def test_bearer_token_sets_auth_client_context_for_service_handler(self, auth_co @app.get("/whoami") async def whoami(auth_client: AuthClient = Depends(get_auth_client)): principal = auth_client.principal + effective_principal = principal.effective_principal return { "principal": principal.id, "email": principal.email, "groups": principal.groups, + "on_behalf_of": principal.on_behalf_of, + "on_behalf_of_email": principal.on_behalf_of_email, + "on_behalf_of_groups": principal.on_behalf_of_groups, + "effective_principal": effective_principal.id, + "effective_email": effective_principal.email, + "effective_groups": effective_principal.groups, } Configuration.set_override(auth_config_enabled) @@ -587,6 +594,130 @@ async def whoami(auth_client: AuthClient = Depends(get_auth_client)): "principal": "alice@example.com", "email": "alice@example.com", "groups": ["team-ml", "team-ai"], + "on_behalf_of": None, + "on_behalf_of_email": None, + "on_behalf_of_groups": None, + "effective_principal": "alice@example.com", + "effective_email": "alice@example.com", + "effective_groups": ["team-ml", "team-ai"], + } + resolver.assert_awaited_once() + mock_authorize.assert_called_once() + assert mock_authorize.call_args.kwargs["scopes"] == ["models:read"] + + def test_bearer_token_with_actor_sets_delegated_auth_client_context(self, auth_config_enabled): + app = FastAPI() + + @app.get("/whoami") + async def whoami(auth_client: AuthClient = Depends(get_auth_client)): + principal = auth_client.principal + effective_principal = principal.effective_principal + return { + "principal": principal.id, + "email": principal.email, + "groups": principal.groups, + "on_behalf_of": principal.on_behalf_of, + "on_behalf_of_email": principal.on_behalf_of_email, + "on_behalf_of_groups": principal.on_behalf_of_groups, + "effective_principal": effective_principal.id, + "effective_email": effective_principal.email, + "effective_groups": effective_principal.groups, + } + + Configuration.set_override(auth_config_enabled) + app.add_middleware(AuthorizationMiddleware, service_name="test-service") + client = TestClient(app, raise_server_exceptions=False) + claims = TokenClaims( + subject="creator@example.com", + email="creator@example.com", + groups=["workspace-editors"], + scopes=["models:read"], + raw_claims={}, + actor=ActorClaims( + subject="system:serviceaccount:nemo-runs:job-runner", + groups=["system:serviceaccounts"], + ), + ) + resolved = ResolvedBearerToken(claims=claims, token_kind="workload_access_token") + + with patch( + "nmp.common.auth.middleware.resolve_bearer_token", + new=AsyncMock(return_value=resolved), + ) as resolver: + with patch.object(AuthClient, "authorize_request", autospec=True) as mock_authorize: + mock_authorize.return_value = MagicMock(allowed=True) + response = client.get("/whoami", headers={"Authorization": "Bearer workload-token"}) + + assert response.status_code == 200 + assert response.json() == { + "principal": "system:serviceaccount:nemo-runs:job-runner", + "email": None, + "groups": ["system:serviceaccounts"], + "on_behalf_of": "creator@example.com", + "on_behalf_of_email": "creator@example.com", + "on_behalf_of_groups": ["workspace-editors"], + "effective_principal": "creator@example.com", + "effective_email": "creator@example.com", + "effective_groups": ["workspace-editors"], + } + resolver.assert_awaited_once() + mock_authorize.assert_called_once() + assert mock_authorize.call_args.kwargs["scopes"] == ["models:read"] + + def test_oidc_bearer_token_with_actor_uses_direct_auth_client_context(self, auth_config_enabled): + app = FastAPI() + + @app.get("/whoami") + async def whoami(auth_client: AuthClient = Depends(get_auth_client)): + principal = auth_client.principal + effective_principal = principal.effective_principal + return { + "principal": principal.id, + "email": principal.email, + "groups": principal.groups, + "on_behalf_of": principal.on_behalf_of, + "on_behalf_of_email": principal.on_behalf_of_email, + "on_behalf_of_groups": principal.on_behalf_of_groups, + "effective_principal": effective_principal.id, + "effective_email": effective_principal.email, + "effective_groups": effective_principal.groups, + } + + Configuration.set_override(auth_config_enabled) + app.add_middleware(AuthorizationMiddleware, service_name="test-service") + client = TestClient(app, raise_server_exceptions=False) + claims = TokenClaims( + subject="creator@example.com", + email="creator@example.com", + groups=["workspace-editors"], + scopes=["models:read"], + raw_claims={"act": {"sub": "system:serviceaccount:nemo-runs:job-runner"}}, + actor=ActorClaims( + subject="system:serviceaccount:nemo-runs:job-runner", + groups=["system:serviceaccounts"], + ), + ) + resolved = ResolvedBearerToken(claims=claims, token_kind="oidc_access_token") + + with patch( + "nmp.common.auth.middleware.resolve_bearer_token", + new=AsyncMock(return_value=resolved), + ) as resolver: + with patch.object(AuthClient, "authorize_request", autospec=True) as mock_authorize: + mock_authorize.return_value = MagicMock(allowed=True) + response = client.get("/whoami", headers={"Authorization": "Bearer oidc-token"}) + + assert response.status_code == 200 + assert response.json() == { + "principal": "creator@example.com", + "email": "creator@example.com", + "groups": ["workspace-editors"], + "on_behalf_of": None, + "on_behalf_of_email": None, + "on_behalf_of_groups": None, + "effective_principal": "creator@example.com", + "effective_email": "creator@example.com", + "effective_groups": ["workspace-editors"], } resolver.assert_awaited_once() mock_authorize.assert_called_once() diff --git a/packages/nmp_common/tests/auth/test_token_resolver.py b/packages/nmp_common/tests/auth/test_token_resolver.py index b712c15db5..baa9aab65e 100644 --- a/packages/nmp_common/tests/auth/test_token_resolver.py +++ b/packages/nmp_common/tests/auth/test_token_resolver.py @@ -4,8 +4,8 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from nmp.common.auth.jwt import TokenClaims -from nmp.common.auth.token_resolver import ResolvedBearerToken, resolve_bearer_token +from nmp.common.auth.jwt import ActorClaims, TokenClaims +from nmp.common.auth.token_resolver import ResolvedBearerToken, ResolvedTokenKind, resolve_bearer_token from nmp.common.config import AuthConfig from nmp.common.config.base import AccessKeyConfig, OIDCConfig @@ -20,6 +20,62 @@ def _claims(subject: str = "alice@example.com") -> TokenClaims: ) +def _claims_with_actor() -> TokenClaims: + return TokenClaims( + subject="user:alice", + email="alice@example.test", + groups=["researchers"], + scopes=["models.read"], + raw_claims={ + "sub": "user:alice", + "email": "alice@example.test", + "groups": "researchers", + "act": {"sub": "service:jobs", "groups": "system:serviceaccounts"}, + }, + actor=ActorClaims(subject="service:jobs", groups=["system:serviceaccounts"]), + ) + + +def test_resolved_token_maps_actor_claims_to_delegated_principal_headers() -> None: + claims = _claims_with_actor() + resolved = ResolvedBearerToken(claims=claims, token_kind="workload_access_token") + + assert resolved.principal.id == "service:jobs" + assert resolved.principal.groups == ["system:serviceaccounts"] + assert resolved.principal.on_behalf_of == "user:alice" + assert resolved.principal.on_behalf_of_email == "alice@example.test" + assert resolved.principal.on_behalf_of_groups == ["researchers"] + assert resolved.principal.effective_principal.id == "user:alice" + assert resolved.principal_headers() == { + "X-NMP-Principal-Id": "service:jobs", + "X-NMP-Principal-Groups": "system:serviceaccounts", + "X-NMP-Principal-On-Behalf-Of": "user:alice", + "X-NMP-Principal-On-Behalf-Of-Groups": "researchers", + "X-NMP-Principal-On-Behalf-Of-Email": "alice@example.test", + "X-NMP-Scopes": "models.read", + } + + +@pytest.mark.parametrize("token_kind", ["oidc_access_token", "access_key", "workload_subject_token"]) +def test_non_workload_access_tokens_ignore_actor_claims(token_kind: ResolvedTokenKind) -> None: + resolved = ResolvedBearerToken(claims=_claims_with_actor(), token_kind=token_kind) + + principal = resolved.principal + + assert principal.id == "user:alice" + assert principal.email == "alice@example.test" + assert principal.groups == ["researchers"] + assert principal.on_behalf_of is None + assert principal.on_behalf_of_email is None + assert principal.on_behalf_of_groups is None + assert resolved.principal_headers() == { + "X-NMP-Principal-Id": "user:alice", + "X-NMP-Principal-Email": "alice@example.test", + "X-NMP-Principal-Groups": "researchers", + "X-NMP-Scopes": "models.read", + } + + @pytest.mark.asyncio async def test_resolver_skips_access_key_validator_when_access_keys_are_disabled() -> None: config = AuthConfig( diff --git a/packages/nmp_common/tests/auth/test_workload_delegations.py b/packages/nmp_common/tests/auth/test_workload_delegations.py new file mode 100644 index 0000000000..9be97c3f11 --- /dev/null +++ b/packages/nmp_common/tests/auth/test_workload_delegations.py @@ -0,0 +1,357 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import hashlib +import inspect +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock + +import pytest +from nmp.common.auth import AuthContext, Principal +from nmp.common.auth.workload_delegations import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, + OPAQUE_DOCKER_PROOF_PREFIX, + InvalidWorkloadProofTokenError, + WorkloadDelegationConflictError, + WorkloadDelegationEntity, + WorkloadDelegationStore, + WorkloadDelegationValidationError, + create_opaque_docker_proof_token, + docker_delegation_name, + parse_opaque_docker_proof_token, + reference_delegation_name, + verify_opaque_docker_proof_token_hash, +) +from nmp.common.entities import SYSTEM_WORKSPACE, EntityClient, EntityConflictError, EntityNotFoundError + + +def _expires_at() -> datetime: + return datetime.now(timezone.utc) + timedelta(hours=1) + + +def _entity(**overrides) -> WorkloadDelegationEntity: + values = { + "name": docker_delegation_name( + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + ), + "workspace": SYSTEM_WORKSPACE, + "workload_subject": docker_delegation_name( + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + ), + "workload_audience": "nemo-platform", + "workload_workspace": "default", + "job_id": "job-123", + "attempt_id": "attempt-1", + "step_id": "step-a", + "auth_context": AuthContext.from_principal(Principal(id="creator@example.com")), + "expires_at": _expires_at(), + } + values.update(overrides) + return WorkloadDelegationEntity(**values) + + +def _reference_entity(**overrides) -> WorkloadDelegationEntity: + values = { + "name": reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:ns:runner", + bound_reference_name="authentication.kubernetes.io/pod-uid", + bound_reference_value="pod-uid-123", + ), + "workload_subject": "system:serviceaccount:ns:runner", + "bound_reference_name": "authentication.kubernetes.io/pod-uid", + "bound_reference_value": "pod-uid-123", + } + values.update(overrides) + return _entity(**values) + + +def _entity_client() -> AsyncMock: + return AsyncMock(spec=EntityClient) + + +def test_docker_delegation_name_is_deterministic() -> None: + expected_input = '["default","job-123","attempt-1","step-a"]'.encode() + expected = "job-" + hashlib.sha256(expected_input).hexdigest()[:48] + + name = docker_delegation_name( + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + ) + + assert name == expected + assert ( + docker_delegation_name( + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + ) + == name + ) + assert set(inspect.signature(docker_delegation_name).parameters) == { + "workload_workspace", + "job_id", + "attempt_id", + "step_id", + } + + +def test_reference_delegation_name_is_canonical_hash() -> None: + expected_input = ( + '["nemo-platform","system:serviceaccount:ns:runner","authentication.kubernetes.io/pod-uid","pod-uid-123"]' + ).encode() + expected = "ref-" + hashlib.sha256(expected_input).hexdigest()[:48] + + name = reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:ns:runner", + bound_reference_name="authentication.kubernetes.io/pod-uid", + bound_reference_value="pod-uid-123", + ) + + assert name == expected + assert ( + reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:ns:runner", + bound_reference_name="authentication.kubernetes.io/pod-uid", + bound_reference_value="pod-uid-123", + ) + == name + ) + assert set(inspect.signature(reference_delegation_name).parameters) == { + "workload_audience", + "workload_subject", + "bound_reference_name", + "bound_reference_value", + } + + +def test_opaque_docker_proof_token_round_trip() -> None: + token, token_hash = create_opaque_docker_proof_token("job-abc") + parsed = parse_opaque_docker_proof_token(token) + + assert DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE.endswith("docker-opaque-workload-proof") + assert token.startswith(f"{OPAQUE_DOCKER_PROOF_PREFIX}.") + assert len(token.split(".")) == 3 + assert parsed.delegation_name == "job-abc" + assert len(parsed.secret) == 32 + assert verify_opaque_docker_proof_token_hash(parsed.secret, token_hash) + assert not verify_opaque_docker_proof_token_hash(parsed.secret, "v1:sha256:bad") + + +@pytest.mark.parametrize( + "token", + [ + "", + "wrong.job.secret", + "nmp_obo_v1.only-two-parts", + "nmp_obo_v1...secret", + "nmp_obo_v1.am9iOmFiYw.bad-secret", + ], +) +def test_parse_opaque_docker_proof_token_rejects_malformed_tokens(token: str) -> None: + with pytest.raises(InvalidWorkloadProofTokenError): + parse_opaque_docker_proof_token(token) + + +@pytest.mark.asyncio +async def test_register_creates_system_workspace_entity() -> None: + entity_client = _entity_client() + entity = _entity() + entity_client.add.return_value = entity + + saved = await WorkloadDelegationStore(entity_client).register(entity) + + assert saved == entity + entity_client.add.assert_awaited_once_with(entity) + entity_client.get.assert_not_called() + entity_client.list.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_uses_direct_system_workspace_lookup() -> None: + entity_client = _entity_client() + entity = _entity() + entity_client.get.return_value = entity + + result = await WorkloadDelegationStore(entity_client).get(entity.name) + + assert result == entity + entity_client.get.assert_awaited_once_with(WorkloadDelegationEntity, entity.name, workspace=SYSTEM_WORKSPACE) + entity_client.list.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_returns_none_when_missing() -> None: + entity_client = _entity_client() + entity_client.get.side_effect = EntityNotFoundError("missing") + + result = await WorkloadDelegationStore(entity_client).get("job:missing") + + assert result is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "entity", + [ + _entity(workload_audience=""), + _entity(expires_at=datetime.now(timezone.utc) - timedelta(seconds=1)), + _entity(bound_reference_name="authentication.kubernetes.io/pod-uid"), + _entity(bound_reference_value="pod-uid-123"), + _entity( + bound_reference_name="authentication.kubernetes.io/pod-uid", + bound_reference_value="pod-uid-123", + ), + _entity(name="job:wrong"), + _reference_entity(name="ref:wrong"), + _entity(name="workload:runner", workload_subject="system:serviceaccount:ns:runner"), + _entity(workspace="default"), + ], +) +async def test_register_rejects_invalid_entities(entity: WorkloadDelegationEntity) -> None: + with pytest.raises(WorkloadDelegationValidationError): + await WorkloadDelegationStore(_entity_client()).register(entity) + + +@pytest.mark.asyncio +async def test_register_opaque_docker_requires_stored_hash() -> None: + with pytest.raises(WorkloadDelegationValidationError): + await WorkloadDelegationStore(_entity_client()).register( + _entity(), + require_opaque_subject_token_hash=True, + ) + + +@pytest.mark.asyncio +async def test_reference_bound_entity_registers_with_reference_pair() -> None: + entity_client = _entity_client() + entity = _reference_entity() + entity_client.add.return_value = entity + + saved = await WorkloadDelegationStore(entity_client).register(entity) + + assert saved == entity + entity_client.add.assert_awaited_once_with(entity) + + +@pytest.mark.asyncio +async def test_update_uses_fetched_db_version() -> None: + entity_client = _entity_client() + fetched = _entity(workload_audience="old-audience") + fetched._db_version = 7 + replacement = _entity(workload_audience="new-audience") + entity_client.get.return_value = fetched + entity_client.update.side_effect = lambda entity: entity + + updated = await WorkloadDelegationStore(entity_client).update(replacement) + + entity_client.get.assert_awaited_once_with(WorkloadDelegationEntity, replacement.name, workspace=SYSTEM_WORKSPACE) + entity_client.update.assert_awaited_once() + update_arg = entity_client.update.call_args.args[0] + assert update_arg.db_version == 7 + assert update_arg.workload_audience == "new-audience" + assert updated.db_version == 7 + + +@pytest.mark.asyncio +async def test_update_rejects_stale_expected_db_version() -> None: + entity_client = _entity_client() + fetched = _entity(workload_audience="old-audience") + fetched._db_version = 7 + replacement = _entity(workload_audience="new-audience") + entity_client.get.return_value = fetched + + with pytest.raises(WorkloadDelegationConflictError): + await WorkloadDelegationStore(entity_client).update(replacement, expected_db_version=6) + + entity_client.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_revoke_uses_fetched_db_version() -> None: + entity_client = _entity_client() + fetched = _entity() + fetched._db_version = 9 + entity_client.get.return_value = fetched + entity_client.update.side_effect = lambda entity: entity + now = datetime.now(timezone.utc) + + revoked = await WorkloadDelegationStore(entity_client).revoke(fetched.name, now=now) + + assert revoked is not None + entity_client.get.assert_awaited_once_with(WorkloadDelegationEntity, fetched.name, workspace=SYSTEM_WORKSPACE) + entity_client.update.assert_awaited_once() + update_arg = entity_client.update.call_args.args[0] + assert update_arg.db_version == 9 + assert update_arg.revoked_at == now + + +@pytest.mark.asyncio +async def test_revoke_returns_none_when_missing() -> None: + entity_client = _entity_client() + entity_client.get.side_effect = EntityNotFoundError("missing") + + result = await WorkloadDelegationStore(entity_client).revoke("job:missing") + + assert result is None + entity_client.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_register_conflict_on_active_row_raises_domain_conflict() -> None: + entity_client = _entity_client() + existing = _entity() + entity_client.add.side_effect = EntityConflictError("already exists") + entity_client.get.return_value = existing + + with pytest.raises(WorkloadDelegationConflictError): + await WorkloadDelegationStore(entity_client).register(_entity()) + + entity_client.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_register_conflict_replaces_expired_row_with_fetched_db_version() -> None: + entity_client = _entity_client() + expired = _entity(expires_at=datetime.now(timezone.utc) - timedelta(hours=1)) + expired._db_version = 12 + replacement = _entity() + entity_client.add.side_effect = EntityConflictError("already exists") + entity_client.get.return_value = expired + entity_client.update.side_effect = lambda entity: entity + + saved = await WorkloadDelegationStore(entity_client).register(replacement) + + entity_client.update.assert_awaited_once() + update_arg = entity_client.update.call_args.args[0] + assert update_arg.db_version == 12 + assert update_arg.expires_at == replacement.expires_at + assert saved.db_version == 12 + + +@pytest.mark.asyncio +async def test_register_conflict_rejects_stale_expected_db_version_for_replacement() -> None: + entity_client = _entity_client() + expired = _entity(expires_at=datetime.now(timezone.utc) - timedelta(hours=1)) + expired._db_version = 12 + replacement = _entity() + entity_client.add.side_effect = EntityConflictError("already exists") + entity_client.get.return_value = expired + + with pytest.raises(WorkloadDelegationConflictError): + await WorkloadDelegationStore(entity_client).register(replacement, expected_db_version=11) + + entity_client.update.assert_not_called() diff --git a/sdk/python/nemo-platform/.nmpcontext/openapi.yaml b/sdk/python/nemo-platform/.nmpcontext/openapi.yaml index 33a0d3c022..d503c2233b 100644 --- a/sdk/python/nemo-platform/.nmpcontext/openapi.yaml +++ b/sdk/python/nemo-platform/.nmpcontext/openapi.yaml @@ -126,6 +126,7 @@ paths: type: string enum: - urn:ietf:params:oauth:token-type:jwt + - urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof description: Token type identifier for the subject token. requested_token_type: type: string @@ -136,6 +137,18 @@ paths: audience: type: string description: Requested audience for the issued access token. + resource: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. + actor_token_type: + not: {} + type: string + description: Unsupported for the jobs workload token exchange profile. scope: type: string description: Space-separated scopes requested for the issued access @@ -10443,9 +10456,7 @@ components: - $ref: '#/components/schemas/DockerJobNetworkConfig' description: Docker networking configuration workload_identity: - allOf: - - $ref: '#/components/schemas/DockerWorkloadIdentityConfig' - description: Docker workload identity subject-token issuer configuration. + $ref: '#/components/schemas/DockerWorkloadIdentityConfig' type: object title: DockerJobExecutionProfileConfig description: Configuration for Docker Job execution profile. @@ -10516,59 +10527,11 @@ components: - mount_path title: DockerVolumeMount DockerWorkloadIdentityConfig: - properties: - enabled: - title: Enabled - description: Enable Docker workload identity token-file injection. Defaults - to auth.oidc.workload_token_exchange_enabled. - type: boolean - token_endpoint: - title: Token Endpoint - description: OAuth token endpoint used by the Docker demo issuer. Defaults - to auth.oidc.token_endpoint. - type: string - client_id: - title: Client Id - description: OAuth client ID used by the Docker demo issuer. Defaults to - auth.oidc.workload_client_id or auth.oidc.client_id. - type: string - client_secret: - format: password - title: Client Secret - description: OAuth client secret for the Docker demo issuer. - writeOnly: true - type: string - username: - title: Username - description: Username for the Docker demo issuer password grant. - type: string - password_env_var: - type: string - title: Password Env Var - description: Controller environment variable that contains the Docker demo - issuer password grant shared secret. - default: AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD - scope: - title: Scope - description: OAuth scope for the Docker demo issuer. - type: string - subject_token_ttl_seconds: - type: integer - minimum: 1.0 - title: Subject Token Ttl Seconds - description: Fallback subject-token lifetime when the Docker demo issuer - response omits expires_in. - refresh_margin_seconds: - type: integer - minimum: 0.0 - title: Refresh Margin Seconds - description: Seconds before subject-token expiry when the Docker refresher - issues a replacement token. - default: 60 + properties: {} additionalProperties: false type: object title: DockerWorkloadIdentityConfig - description: Docker-only subject token issuer configuration for workload identity. + description: Docker workload identity configuration. E2EJobExecutionProfile: properties: provider: @@ -19544,7 +19507,7 @@ components: type: string title: Error description: OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, - invalid_request, invalid_grant, invalid_scope, or invalid_target. + invalid_request, invalid_scope, or invalid_target. error_description: title: Error Description description: Human-readable ASCII text providing additional information diff --git a/sdk/python/nemo-platform/src/nemo_platform/auth/workload_exchange.py b/sdk/python/nemo-platform/src/nemo_platform/auth/workload_exchange.py index c54423fc01..b648da1c2d 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/auth/workload_exchange.py +++ b/sdk/python/nemo-platform/src/nemo_platform/auth/workload_exchange.py @@ -16,7 +16,14 @@ from urllib.parse import urlparse import httpx -from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +from nemo_platform_plugin.client.constants import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE as _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, +) +from nemo_platform_plugin.client.constants import ( + JWT_WORKLOAD_SUBJECT_TOKEN_TYPE, + WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, + subject_token_type_for_exchange, +) from nemo_platform.auth.token_provider import DEFAULT_REFRESH_MARGIN_SECONDS, TokenSet from nemo_platform.client.tls import client_verify_from_env @@ -24,7 +31,8 @@ logger = logging.getLogger(__name__) TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" -JWT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" +JWT_TOKEN_TYPE = JWT_WORKLOAD_SUBJECT_TOKEN_TYPE +DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE = _DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE ACCESS_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token" @@ -87,7 +95,7 @@ def token_exchange_grant( "grant_type": TOKEN_EXCHANGE_GRANT_TYPE, "client_id": client_id, "subject_token": subject_token, - "subject_token_type": JWT_TOKEN_TYPE, + "subject_token_type": subject_token_type_for_exchange(subject_token), "requested_token_type": ACCESS_TOKEN_TYPE, } if audience: diff --git a/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/jobs/__init__.py b/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/jobs/__init__.py index 9a6fcc0b41..c0ac03827b 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/jobs/__init__.py +++ b/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/jobs/__init__.py @@ -400,7 +400,12 @@ def list_execution_profiles_jobs( columns: OutputColumnsOption = None, stream: StreamOutputOption = False, ) -> None: - """Get all currently configured execution profiles.""" + """Get all currently configured execution profiles. + + Returns the capability-filtered merge from jobs config. In local standalone the + controller may prune the shared list further after registry boot; in split + topologies the API advertises its own merge result (not controller process + memory).""" state: CLIContext = ctx.obj output_format = state.get_output_format(output_format) validate_stream_output_format(output_format, stream) diff --git a/sdk/python/nemo-platform/src/nemo_platform/resources/jobs/jobs.py b/sdk/python/nemo-platform/src/nemo_platform/resources/jobs/jobs.py index 1c8bf65dbe..3fb1c4e410 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/resources/jobs/jobs.py +++ b/sdk/python/nemo-platform/src/nemo_platform/resources/jobs/jobs.py @@ -459,7 +459,14 @@ def list_execution_profiles( extra_body: Body | None = None, timeout: float | httpx.Timeout | None | NotGiven = not_given, ) -> JobListExecutionProfilesResponse: - """Get all currently configured execution profiles.""" + """ + Get all currently configured execution profiles. + + Returns the capability-filtered merge from jobs config. In local standalone the + controller may prune the shared list further after registry boot; in split + topologies the API advertises its own merge result (not controller process + memory). + """ return self._get( "/apis/jobs/v2/execution-profiles", options=make_request_options( @@ -971,7 +978,14 @@ async def list_execution_profiles( extra_body: Body | None = None, timeout: float | httpx.Timeout | None | NotGiven = not_given, ) -> JobListExecutionProfilesResponse: - """Get all currently configured execution profiles.""" + """ + Get all currently configured execution profiles. + + Returns the capability-filtered merge from jobs config. In local standalone the + controller may prune the shared list further after registry boot; in split + topologies the API advertises its own merge result (not controller process + memory). + """ return await self._get( "/apis/jobs/v2/execution-profiles", options=make_request_options( diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_job_execution_profile_config.py b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_job_execution_profile_config.py index b1b31cad62..59b08338c6 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_job_execution_profile_config.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_job_execution_profile_config.py @@ -61,4 +61,4 @@ class DockerJobExecutionProfileConfig(BaseModel): ttl_seconds_before_active: Optional[int] = None workload_identity: Optional[DockerWorkloadIdentityConfig] = None - """Docker-only subject token issuer configuration for workload identity.""" + """Docker workload identity configuration.""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_workload_identity_config.py b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_workload_identity_config.py index d4910a5138..284bf50aa8 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_workload_identity_config.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/docker_workload_identity_config.py @@ -15,54 +15,10 @@ # File generated from our OpenAPI spec by Stainless. See CONTRIBUTING.md for details. -from typing import Optional - from ..._models import BaseModel __all__ = ["DockerWorkloadIdentityConfig"] class DockerWorkloadIdentityConfig(BaseModel): - """Docker-only subject token issuer configuration for workload identity.""" - - client_id: Optional[str] = None - """OAuth client ID used by the Docker demo issuer. - - Defaults to auth.oidc.workload_client_id or auth.oidc.client_id. - """ - - enabled: Optional[bool] = None - """Enable Docker workload identity token-file injection. - - Defaults to auth.oidc.workload_token_exchange_enabled. - """ - - password_env_var: Optional[str] = None - """ - Controller environment variable that contains the Docker demo issuer password - grant shared secret. - """ - - refresh_margin_seconds: Optional[int] = None - """ - Seconds before subject-token expiry when the Docker refresher issues a - replacement token. - """ - - scope: Optional[str] = None - """OAuth scope for the Docker demo issuer.""" - - subject_token_ttl_seconds: Optional[int] = None - """ - Fallback subject-token lifetime when the Docker demo issuer response omits - expires_in. - """ - - token_endpoint: Optional[str] = None - """OAuth token endpoint used by the Docker demo issuer. - - Defaults to auth.oidc.token_endpoint. - """ - - username: Optional[str] = None - """Username for the Docker demo issuer password grant.""" + """Docker workload identity configuration.""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/shared/workload_token_exchange_error_response.py b/sdk/python/nemo-platform/src/nemo_platform/types/shared/workload_token_exchange_error_response.py index e02b5d084b..05ba154ede 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/shared/workload_token_exchange_error_response.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/shared/workload_token_exchange_error_response.py @@ -28,7 +28,7 @@ class WorkloadTokenExchangeErrorResponse(BaseModel): error: str """ OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, - invalid_request, invalid_grant, invalid_scope, or invalid_target. + invalid_request, invalid_scope, or invalid_target. """ error_description: Optional[str] = None diff --git a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/auth/test_workload_exchange.py b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/auth/test_workload_exchange.py index 311c43173b..5684677222 100644 --- a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/auth/test_workload_exchange.py +++ b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/auth/test_workload_exchange.py @@ -3,14 +3,19 @@ """Tests for RFC 8693 workload identity token exchange.""" +import ast import json import time from base64 import urlsafe_b64encode +from pathlib import Path from unittest.mock import MagicMock, patch +import httpx +import nemo_platform.auth.workload_exchange as workload_exchange_module import pytest from nemo_platform.auth.workload_exchange import ( ACCESS_TOKEN_TYPE, + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, JWT_TOKEN_TYPE, TOKEN_EXCHANGE_GRANT_TYPE, WorkloadTokenExchangeError, @@ -22,6 +27,19 @@ from nemo_platform_plugin.client.constants import WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR +def test_workload_exchange_module_has_no_nmp_common_dependency(): + assert workload_exchange_module.__file__ is not None + tree = ast.parse(Path(workload_exchange_module.__file__).read_text(encoding="utf-8")) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + imported = [alias.name for alias in node.names] + elif isinstance(node, ast.ImportFrom): + imported = [node.module or ""] + else: + continue + assert not any(name == "nmp.common" or name.startswith("nmp.common.") for name in imported) + + def _make_jwt(claims: dict) -> str: header = {"alg": "RS256", "typ": "JWT"} h = urlsafe_b64encode(json.dumps(header).encode()).rstrip(b"=").decode() @@ -74,6 +92,19 @@ def test_token_exchange_grant_sends_rfc8693_request(mock_post): } +@patch("nemo_platform.auth.workload_exchange.httpx.post") +def test_token_exchange_grant_uses_docker_opaque_subject_token_type(mock_post): + mock_post.return_value = httpx.Response(200, json={"access_token": "exchanged-token", "expires_in": 300}) + + token_exchange_grant( + token_endpoint="https://idp.example.com/token", + client_id="nemo-platform-workload", + subject_token="nmp_obo_v1.delegation.secret", + ) + + assert mock_post.call_args.kwargs["data"]["subject_token_type"] == DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE + + @patch("nemo_platform.auth.workload_exchange.httpx.post") def test_token_exchange_grant_rejects_http_non_loopback_endpoint_before_sending_subject_token(mock_post): with pytest.raises(ValueError, match="must use HTTPS"): diff --git a/services/core/auth/src/nmp/core/auth/api/v2/authenticate.py b/services/core/auth/src/nmp/core/auth/api/v2/authenticate.py index 97fad73e93..4f7f821160 100644 --- a/services/core/auth/src/nmp/core/auth/api/v2/authenticate.py +++ b/services/core/auth/src/nmp/core/auth/api/v2/authenticate.py @@ -6,10 +6,11 @@ import logging from typing import Any +import httpx import jwt from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from nmp.common.auth.bearer import MalformedBearerTokenError, parse_bearer_authorization_header -from nmp.common.auth.jwt import TokenClaims +from nmp.common.auth.jwt import ActorClaims, JWTValidator, TokenClaims from nmp.common.auth.token_resolver import ResolvedBearerToken, ResolvedTokenKind, resolve_bearer_token from nmp.common.config import AuthConfig, get_auth_config from nmp.core.auth.api.v2.workload_token_exchange import ( @@ -17,6 +18,7 @@ _allowed_audiences, _workload_token_issuer, get_workload_token_exchange_service, + workload_jwks_url, ) from pydantic import BaseModel, Field @@ -53,6 +55,10 @@ class AuthenticateResponse(BaseModel): } +class TokenIssuerBoundaryMisconfigurationError(Exception): + """Raised when IdP and NeMo-issued token validation paths overlap.""" + + def _bearer_token_from_request(request: Request) -> str: try: token = parse_bearer_authorization_header(request.headers.get("authorization")) @@ -80,6 +86,93 @@ def _scopes_from_claims(claims: dict[str, object]) -> list[str]: return [] +def _normalize_url(value: str) -> str: + return value.rstrip("/") + + +def _normalize_jwk_public_material(jwk: dict[str, object]) -> tuple[tuple[str, str], ...] | None: + kty = jwk.get("kty") + if kty == "RSA": + keys = ("kty", "n", "e") + elif kty == "EC": + keys = ("kty", "crv", "x", "y") + elif kty == "OKP": + keys = ("kty", "crv", "x") + else: + return None + + material: list[tuple[str, str]] = [] + for key in keys: + value = jwk.get(key) + if not isinstance(value, str) or not value: + return None + material.append((key, value)) + return tuple(material) + + +async def _validate_token_issuer_boundary( + config: AuthConfig, + request: Request, + workload_token_exchange_service: WorkloadTokenExchangeService, + jwt_validator: JWTValidator, +) -> None: + if not (config.oidc.enabled and config.oidc.workload_token_exchange_enabled): + return + + workload_issuer = _normalize_url(_workload_token_issuer(config, request)) + idp_issuers = {_normalize_url(issuer) for issuer in [config.oidc.issuer, *config.oidc.additional_issuers] if issuer} + if workload_issuer in idp_issuers: + raise TokenIssuerBoundaryMisconfigurationError("OIDC issuer must not match NeMo workload token issuer") + + try: + idp_jwks_uri = _normalize_url(await jwt_validator.jwks_uri()) + except (httpx.HTTPError, jwt.PyJWTError): + return + + nemo_jwks_uri = _normalize_url(workload_jwks_url(request)) + if idp_jwks_uri == nemo_jwks_uri: + raise TokenIssuerBoundaryMisconfigurationError("OIDC JWKS URI must not match NeMo workload JWKS URI") + + try: + signing_key = await workload_token_exchange_service.public_jwk_async(config) + except (OSError, RuntimeError, ValueError): + return + + nemo_material = _normalize_jwk_public_material(signing_key) + if nemo_material is None: + return + + try: + idp_jwks = await jwt_validator.jwks() + except (httpx.HTTPError, jwt.PyJWTError): + return + + for jwk in idp_jwks.get("keys", []): + if isinstance(jwk, dict) and _normalize_jwk_public_material(jwk) == nemo_material: + raise TokenIssuerBoundaryMisconfigurationError( + "OIDC JWKS must not contain the NeMo workload token signing key" + ) + + +def _actor_from_claims(claims: dict[str, object]) -> ActorClaims | None: + actor_claims = claims.get("act") + if not isinstance(actor_claims, dict): + return None + + actor_subject = actor_claims.get("sub") + if not isinstance(actor_subject, str): + return None + + actor_subject = actor_subject.strip() + if not actor_subject: + return None + + return ActorClaims( + subject=actor_subject, + groups=_groups_from_claim(actor_claims.get("groups", [])), + ) + + def _stamp_principal_headers(response: Response, resolved: ResolvedBearerToken) -> None: for header_name, header_value in resolved.principal_headers().items(): response.headers[header_name] = header_value @@ -138,6 +231,7 @@ async def _validate_workload_access_token( groups=_groups_from_claim(claims.get("groups", [])), scopes=_scopes_from_claims(claims), raw_claims=claims, + actor=_actor_from_claims(claims), ) except jwt.PyJWTError: return None @@ -190,6 +284,7 @@ async def _authenticate_bearer_token( ) -> AuthenticateResponse: token = _bearer_token_from_request(request) config = get_auth_config() + jwt_validator = JWTValidator(config) async def resolve_workload_access(candidate: str) -> ResolvedBearerToken | None: return await _resolve_workload_access_token(config, request, workload_token_exchange_service, candidate) @@ -197,9 +292,19 @@ async def resolve_workload_access(candidate: str) -> ResolvedBearerToken | None: async def resolve_workload_subject(candidate: str) -> ResolvedBearerToken | None: return await _resolve_workload_subject_token(config, workload_token_exchange_service, candidate) + try: + await _validate_token_issuer_boundary(config, request, workload_token_exchange_service, jwt_validator) + except TokenIssuerBoundaryMisconfigurationError as exc: + logger.error("Authentication token issuer boundary is misconfigured: %s", exc) + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Authentication token issuers are misconfigured", + ) from exc + resolved = await resolve_bearer_token( config, token, + jwt_validator=jwt_validator, extra_resolvers=[resolve_workload_access, resolve_workload_subject], ) if resolved is None: diff --git a/services/core/auth/src/nmp/core/auth/api/v2/workload_token_exchange.py b/services/core/auth/src/nmp/core/auth/api/v2/workload_token_exchange.py index fa42f8246a..d2335159f5 100644 --- a/services/core/auth/src/nmp/core/auth/api/v2/workload_token_exchange.py +++ b/services/core/auth/src/nmp/core/auth/api/v2/workload_token_exchange.py @@ -5,11 +5,14 @@ from __future__ import annotations +import asyncio import json import logging +import math import time from collections.abc import Awaitable, Callable from dataclasses import dataclass +from datetime import datetime, timezone from pathlib import Path from typing import Any @@ -19,10 +22,23 @@ from fastapi.responses import JSONResponse from nmp.common.auth.access_keys import public_jwk_from_private_key_pem_async from nmp.common.auth.signing_keys import RSASigningKey, RSASigningKeyCache +from nmp.common.auth.workload_delegations import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, + InvalidWorkloadProofTokenError, + WorkloadDelegationEntity, + WorkloadDelegationStore, + parse_opaque_docker_proof_token, + reference_delegation_name, + verify_opaque_docker_proof_token_hash, +) from nmp.common.config import AuthConfig, get_auth_config, get_platform_config +from nmp.common.entities import EntityClient +from nmp.common.service.dependencies import get_entity_client +from opentelemetry import metrics from pydantic import BaseModel, ConfigDict, Field logger = logging.getLogger(__name__) +meter = metrics.get_meter(__name__) TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" JWT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt" @@ -32,6 +48,16 @@ WORKLOAD_JWKS_PATH = "/apis/auth/jwks" DEFAULT_WORKLOAD_AUDIENCE = "nemo-platform" DEFAULT_WORKLOAD_SCOPE = "openid email groups" +KUBERNETES_POD_UID_REFERENCE_NAME = "authentication.kubernetes.io/pod-uid" +_BOUND_REFERENCE_NAME_CLAIM = "_nmp_bound_reference_name" +_BOUND_REFERENCE_VALUE_CLAIM = "_nmp_bound_reference_value" +_BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM = "_nmp_bound_reference_trusted_source" +_BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES = "kubernetes" + +_delegation_lookup_retry_exhausted_total = meter.create_counter( + name="nmp.auth.workload_delegation.lookup_retry_exhausted.total", + description="Number of workload delegation lookups that exhausted their retry budget", +) _TOKEN_EXCHANGE_FORM_REQUEST_BODY: dict[str, Any] = { "required": True, @@ -57,7 +83,7 @@ "subject_token_type": { "type": "string", "description": "Token type identifier for the subject token.", - "enum": [JWT_TOKEN_TYPE], + "enum": [JWT_TOKEN_TYPE, DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE], }, "requested_token_type": { "type": "string", @@ -69,6 +95,21 @@ "type": "string", "description": "Requested audience for the issued access token.", }, + "resource": { + "type": "string", + "description": "Unsupported for the jobs workload token exchange profile.", + "not": {}, + }, + "actor_token": { + "type": "string", + "description": "Unsupported for the jobs workload token exchange profile.", + "not": {}, + }, + "actor_token_type": { + "type": "string", + "description": "Unsupported for the jobs workload token exchange profile.", + "not": {}, + }, "scope": { "type": "string", "description": "Space-separated scopes requested for the issued access token.", @@ -94,6 +135,24 @@ class _SubjectTokenDecoder: decode: Callable[[], Awaitable[dict[str, Any]]] +@dataclass(frozen=True) +class VerifiedWorkloadReference: + name: str + value: str + + +@dataclass(frozen=True) +class VerifiedSubjectToken: + subject: str + groups: list[str] + email: str | None = None + delegation_name: str | None = None + bound_reference: VerifiedWorkloadReference | None = None + is_docker_proof: bool = False + is_opaque_docker_proof: bool = False + opaque_secret: bytes | None = None + + @dataclass(frozen=True) class _SigningKeyLoadRequest: kid: str @@ -102,6 +161,10 @@ class _SigningKeyLoadRequest: invalid_private_key_message: str +class _InvalidGrantError(Exception): + """Raised when a validated subject token is not authorized for delegation.""" + + class WorkloadTokenExchangeResponse(BaseModel): """RFC 8693 token exchange response for workload identity access tokens.""" @@ -118,7 +181,7 @@ class WorkloadTokenExchangeErrorResponse(BaseModel): error: str = Field( description=( "OAuth 2.0 or RFC 8693 token exchange error code, such as invalid_client, " - "invalid_request, invalid_grant, invalid_scope, or invalid_target." + "invalid_request, invalid_scope, or invalid_target." ), ) error_description: str | None = Field( @@ -234,9 +297,19 @@ def _validate_subject_jwks(jwks: dict[str, Any]) -> None: class WorkloadTokenExchangeService: """Stateful helpers for workload token exchange endpoints.""" - def __init__(self, signing_key_cache: RSASigningKeyCache | None = None) -> None: + def __init__( + self, + signing_key_cache: RSASigningKeyCache | None = None, + *, + delegation_lookup_retry_timeout_seconds: float = 5.0, + delegation_lookup_retry_interval_seconds: float = 0.1, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + ) -> None: self._signing_key_cache = signing_key_cache or RSASigningKeyCache() self._subject_jwks_cache: dict[str, _SubjectJWKSCacheEntry] = {} + self._delegation_lookup_retry_timeout_seconds = delegation_lookup_retry_timeout_seconds + self._delegation_lookup_retry_interval_seconds = delegation_lookup_retry_interval_seconds + self._sleep = sleep def workload_signing_key(self, config: AuthConfig) -> RSASigningKey: load_request = _workload_signing_key_load_request(config) @@ -327,7 +400,7 @@ async def decode_jwt_subject_token(self, config: AuthConfig, subject_token: str) issuer = claims.get("iss") if issuer not in _allowed_subject_issuers(config): raise jwt.InvalidIssuerError(f"unexpected subject token issuer: {issuer!r}") - return claims + return _strip_private_bound_reference_claims(claims) async def decode_subject_token(self, config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: errors: list[str] = [] @@ -345,6 +418,37 @@ async def decode_subject_token(self, config: AuthConfig, subject_token: str, aud errors.append(f"{decoder.name}: {message}") raise jwt.InvalidTokenError("; ".join(errors)) + async def get_delegation_with_retry( + self, + store: WorkloadDelegationStore, + name: str, + *, + retry: bool, + ) -> WorkloadDelegationEntity | None: + """Fetch a delegation row by name with a bounded retry for Kubernetes startup races.""" + if not retry: + return await store.get(name) + + retry_attempts = ( + max( + 0, + math.ceil( + self._delegation_lookup_retry_timeout_seconds / self._delegation_lookup_retry_interval_seconds + ), + ) + if self._delegation_lookup_retry_timeout_seconds > 0 and self._delegation_lookup_retry_interval_seconds > 0 + else 0 + ) + while True: + entity = await store.get(name) + if entity is not None: + return entity + if retry_attempts <= 0: + _delegation_lookup_retry_exhausted_total.add(1) + return None + retry_attempts -= 1 + await self._sleep(self._delegation_lookup_retry_interval_seconds) + def get_workload_token_exchange_service(request: Request) -> WorkloadTokenExchangeService: service = getattr(request.app.state, _WORKLOAD_TOKEN_EXCHANGE_SERVICE_STATE_KEY, None) @@ -415,6 +519,26 @@ def _validated_audience(config: AuthConfig, requested_audience: Any) -> str: return audience +def _epoch_seconds(value: datetime) -> int: + if value.tzinfo is None: + value = value.replace(tzinfo=timezone.utc) + return int(value.timestamp()) + + +def _split_scope(scope: Any) -> list[str]: + if not isinstance(scope, str): + return [] + return [part for part in scope.split() if part] + + +def _granted_workload_scope(config: AuthConfig, requested_scope: Any) -> str | None: + allowed_scopes = _split_scope(config.oidc.workload_scope or DEFAULT_WORKLOAD_SCOPE) + requested_scopes = _split_scope(requested_scope) or allowed_scopes + allowed = set(allowed_scopes) + granted = [scope for scope in requested_scopes if scope in allowed] + return " ".join(granted) or None + + def _groups_claim_for_gateway_header(groups: Any) -> str | None: if isinstance(groups, str): return groups @@ -427,6 +551,14 @@ def _allowed_subject_issuers(config: AuthConfig) -> set[str]: return {issuer for issuer in config.oidc.workload_subject_issuers if issuer} +def _strip_private_bound_reference_claims(claims: dict[str, Any]) -> dict[str, Any]: + sanitized = dict(claims) + sanitized.pop(_BOUND_REFERENCE_NAME_CLAIM, None) + sanitized.pop(_BOUND_REFERENCE_VALUE_CLAIM, None) + sanitized.pop(_BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM, None) + return sanitized + + def _kubernetes_reviewer_credentials() -> tuple[str, str]: service_account_dir = Path("/var/run/secrets/kubernetes.io/serviceaccount") reviewer_token = (service_account_dir / "token").read_text(encoding="utf-8").strip() @@ -489,10 +621,157 @@ async def _decode_kubernetes_subject_token( if not subject: raise jwt.InvalidTokenError("Kubernetes TokenReview response did not include a username") - return { + claims: dict[str, Any] = { "sub": subject, "groups": user.get("groups", []), } + extra = user.get("extra") + if isinstance(extra, dict) and KUBERNETES_POD_UID_REFERENCE_NAME in extra: + pod_uids = extra.get(KUBERNETES_POD_UID_REFERENCE_NAME) + if not isinstance(pod_uids, list) or len(pod_uids) != 1 or not isinstance(pod_uids[0], str) or not pod_uids[0]: + raise _InvalidGrantError("Kubernetes TokenReview did not include exactly one pod UID reference") + claims[_BOUND_REFERENCE_NAME_CLAIM] = KUBERNETES_POD_UID_REFERENCE_NAME + claims[_BOUND_REFERENCE_VALUE_CLAIM] = pod_uids[0] + claims[_BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM] = _BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES + + return claims + + +def _subject_groups_from_claims(claims: dict[str, Any]) -> list[str]: + groups = claims.get("groups", []) + if isinstance(groups, str): + return [group.strip() for group in groups.split(",") if group.strip()] + if isinstance(groups, list): + return [str(group).strip() for group in groups if str(group).strip()] + return [] + + +async def _normalize_verified_subject_token( + config: AuthConfig, + workload_token_exchange_service: WorkloadTokenExchangeService, + *, + subject_token: str, + subject_token_type: str, +) -> VerifiedSubjectToken: + if subject_token_type == DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE: + try: + parsed = parse_opaque_docker_proof_token(subject_token) + except InvalidWorkloadProofTokenError as exc: + raise _InvalidGrantError("Docker opaque workload proof token is malformed") from exc + return VerifiedSubjectToken( + subject=parsed.delegation_name, + groups=[], + delegation_name=parsed.delegation_name, + is_docker_proof=True, + is_opaque_docker_proof=True, + opaque_secret=parsed.secret, + ) + + claims = await workload_token_exchange_service.decode_subject_token( + config, + subject_token, + _workload_subject_audience(config), + ) + subject = claims.get("sub") + if not isinstance(subject, str) or not subject: + raise jwt.InvalidTokenError("Subject token did not include a subject") + + bound_reference = None + bound_reference_name = claims.get(_BOUND_REFERENCE_NAME_CLAIM) + bound_reference_value = claims.get(_BOUND_REFERENCE_VALUE_CLAIM) + trusted_reference_source = claims.get(_BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM) + if ( + trusted_reference_source == _BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES + and isinstance(bound_reference_name, str) + and isinstance(bound_reference_value, str) + ): + bound_reference = VerifiedWorkloadReference(name=bound_reference_name, value=bound_reference_value) + + return VerifiedSubjectToken( + subject=subject, + groups=_subject_groups_from_claims(claims), + email=claims.get("email") if isinstance(claims.get("email"), str) else None, + bound_reference=bound_reference, + ) + + +def _delegation_lookup_name(verified_subject: VerifiedSubjectToken, audience: str) -> str | None: + if verified_subject.is_docker_proof: + return verified_subject.delegation_name + if verified_subject.bound_reference is None: + return None + return reference_delegation_name( + workload_audience=audience, + workload_subject=verified_subject.subject, + bound_reference_name=verified_subject.bound_reference.name, + bound_reference_value=verified_subject.bound_reference.value, + ) + + +def _validate_delegation_for_exchange( + delegation: WorkloadDelegationEntity, + verified_subject: VerifiedSubjectToken, + *, + audience: str, +) -> None: + if not delegation.is_active(now=datetime.now(timezone.utc)): + raise _InvalidGrantError("Workload delegation is expired or revoked") + if delegation.workload_audience != audience: + raise _InvalidGrantError("Workload delegation audience does not match") + if delegation.workload_subject != verified_subject.subject: + raise _InvalidGrantError("Workload delegation subject does not match") + + stored_reference = None + if delegation.bound_reference_name or delegation.bound_reference_value: + if not delegation.bound_reference_name or not delegation.bound_reference_value: + raise _InvalidGrantError("Workload delegation bound reference is incomplete") + stored_reference = VerifiedWorkloadReference( + name=delegation.bound_reference_name, + value=delegation.bound_reference_value, + ) + + if stored_reference != verified_subject.bound_reference: + raise _InvalidGrantError("Workload delegation bound reference does not match") + if verified_subject.is_docker_proof and stored_reference is not None: + raise _InvalidGrantError("Docker workload delegation cannot use a bound reference") + + if verified_subject.is_opaque_docker_proof: + if verified_subject.opaque_secret is None or not delegation.opaque_subject_token_hash: + raise _InvalidGrantError("Docker opaque workload delegation is missing proof-token state") + if not verify_opaque_docker_proof_token_hash( + verified_subject.opaque_secret, + delegation.opaque_subject_token_hash, + ): + raise _InvalidGrantError("Docker opaque workload proof-token hash does not match") + elif delegation.opaque_subject_token_hash: + raise _InvalidGrantError("Workload delegation requires an opaque proof token") + + +def _build_workload_only_claims(verified_subject: VerifiedSubjectToken) -> dict[str, Any]: + claims: dict[str, Any] = {"sub": verified_subject.subject} + if verified_subject.email: + claims["email"] = verified_subject.email + if verified_subject.groups: + claims["groups"] = ",".join(verified_subject.groups) + return claims + + +def _build_delegated_claims( + delegation: WorkloadDelegationEntity, + verified_subject: VerifiedSubjectToken, +) -> dict[str, Any]: + delegated_principal = delegation.auth_context.to_principal().effective_principal + claims: dict[str, Any] = {"sub": delegated_principal.id} + if delegated_principal.email: + claims["email"] = delegated_principal.email + if delegated_principal.groups: + claims["groups"] = ",".join(delegated_principal.groups) + + actor_claims: dict[str, Any] = {"sub": verified_subject.subject} + if verified_subject.groups: + actor_claims["groups"] = ",".join(verified_subject.groups) + claims["act"] = actor_claims + return claims @router.get( @@ -528,6 +807,7 @@ async def jwks( async def token_exchange( request: Request, workload_token_exchange_service: WorkloadTokenExchangeService = Depends(get_workload_token_exchange_service), + entity_client: EntityClient | None = Depends(get_entity_client), ) -> WorkloadTokenExchangeResponse | JSONResponse: """Exchange an RFC 8693 workload identity subject token for a NeMo access token.""" config = get_auth_config() @@ -541,51 +821,91 @@ async def token_exchange( return _oauth_error(400, "unsupported_grant_type", "Only RFC 8693 token exchange is supported") if form.get("client_id") != client_id: return _oauth_error(401, "invalid_client", "Unknown workload token exchange client") - if form.get("subject_token_type") != JWT_TOKEN_TYPE: - return _oauth_error(400, "invalid_request", "subject_token_type must be a JWT token type") + subject_token_type = str(form.get("subject_token_type") or "") + if subject_token_type not in {JWT_TOKEN_TYPE, DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE}: + return _oauth_error(400, "invalid_request", "subject_token_type must be a supported workload token type") if form.get("requested_token_type", ACCESS_TOKEN_TYPE) != ACCESS_TOKEN_TYPE: return _oauth_error(400, "invalid_request", "requested_token_type must be access_token") + if "actor_token" in form or "actor_token_type" in form: + return _oauth_error( + 400, + "invalid_request", + "actor_token is not supported for this workload token exchange profile", + ) + + if "resource" in form: + return _oauth_error(400, "invalid_target", "resource is not supported for workload token exchange") + + audience_values = [value for value in form.getlist("audience") if str(value or "").strip()] + if len(audience_values) > 1: + return _oauth_error(400, "invalid_target", "Only one audience is supported for workload token exchange") + requested_audience = audience_values[0] if audience_values else None subject_token = form.get("subject_token") if not subject_token: return _oauth_error(400, "invalid_request", "subject_token is required") try: - audience = _validated_audience(config, form.get("audience")) + audience = _validated_audience(config, requested_audience) except jwt.InvalidAudienceError as exc: logger.info("Requested token audience validation failed: %s", exc) return _oauth_error(400, "invalid_target", "Requested audience is not allowed") try: - subject_claims = await workload_token_exchange_service.decode_subject_token( - config, str(subject_token), _workload_subject_audience(config) + verified_subject = await _normalize_verified_subject_token( + config, + workload_token_exchange_service, + subject_token=str(subject_token), + subject_token_type=subject_token_type, ) - subject = subject_claims.get("sub") - if not subject: - raise jwt.InvalidTokenError("Subject token did not include a subject") except jwt.InvalidTokenError as exc: logger.info("Subject token validation failed: %s", exc) # RFC 8693 clients only need a stable invalid_request response. Avoid # returning decoder details that may include infrastructure internals. return _oauth_error(400, "invalid_request", "Could not validate subject token") + except _InvalidGrantError as exc: + logger.info("Subject token was not eligible for workload delegation: %s", exc) + return _oauth_error(400, "invalid_request", "Subject token is not authorized for workload delegation") + + try: + lookup_name = _delegation_lookup_name(verified_subject, audience) + delegation = None + if lookup_name is not None: + if entity_client is None: + raise _InvalidGrantError("Entity client is unavailable for workload delegation lookup") + delegation_store = WorkloadDelegationStore(entity_client) + delegation = await workload_token_exchange_service.get_delegation_with_retry( + delegation_store, + lookup_name, + retry=verified_subject.bound_reference is not None, + ) + if delegation is None: + raise _InvalidGrantError("Workload delegation is not ready") + _validate_delegation_for_exchange(delegation, verified_subject, audience=audience) + except _InvalidGrantError as exc: + logger.info("Workload delegation lookup failed: %s", exc) + return _oauth_error(400, "invalid_request", "Subject token is not authorized for workload delegation") now = int(time.time()) - scope = str(form.get("scope") or config.oidc.workload_scope or DEFAULT_WORKLOAD_SCOPE) + scope = _granted_workload_scope(config, form.get("scope")) + subject_claims = ( + _build_delegated_claims(delegation, verified_subject) + if delegation is not None + else _build_workload_only_claims(verified_subject) + ) + configured_exp = now + config.oidc.workload_token_ttl_seconds + token_exp = min(configured_exp, _epoch_seconds(delegation.expires_at)) if delegation is not None else configured_exp + expires_in = max(0, token_exp - now) exchanged_claims: dict[str, Any] = { "iss": _workload_token_issuer(config, request), - "sub": subject, + **subject_claims, "aud": audience, "iat": now, "nbf": now, - "exp": now + config.oidc.workload_token_ttl_seconds, - "scope": scope, + "exp": token_exp, } - if "email" in subject_claims: - exchanged_claims["email"] = subject_claims["email"] - if "groups" in subject_claims: - groups_claim = _groups_claim_for_gateway_header(subject_claims["groups"]) - if groups_claim: - exchanged_claims["groups"] = groups_claim + if scope: + exchanged_claims["scope"] = scope signing_key = await workload_token_exchange_service.workload_signing_key_async(config) access_token = jwt.encode( @@ -598,6 +918,6 @@ async def token_exchange( access_token=access_token, issued_token_type=ACCESS_TOKEN_TYPE, token_type="Bearer", - expires_in=config.oidc.workload_token_ttl_seconds, + expires_in=expires_in, scope=scope, ) diff --git a/services/core/auth/tests/integration/test_gateway_header_spoofing.py b/services/core/auth/tests/integration/test_gateway_header_spoofing.py index f77992b5ca..bda5167711 100644 --- a/services/core/auth/tests/integration/test_gateway_header_spoofing.py +++ b/services/core/auth/tests/integration/test_gateway_header_spoofing.py @@ -14,6 +14,7 @@ "x-nmp-principal-on-behalf-of", "x-nmp-principal-on-behalf-of-email", "x-nmp-principal-on-behalf-of-groups", + "x-nmp-scopes", } diff --git a/services/core/auth/tests/test_authenticate.py b/services/core/auth/tests/test_authenticate.py index 493307ac74..20c105a463 100644 --- a/services/core/auth/tests/test_authenticate.py +++ b/services/core/auth/tests/test_authenticate.py @@ -13,7 +13,7 @@ from cryptography.hazmat.primitives.asymmetric import rsa from fastapi import FastAPI from fastapi.testclient import TestClient -from nmp.common.auth.jwt import TokenClaims +from nmp.common.auth.jwt import ActorClaims, JWTValidator, TokenClaims from nmp.common.auth.token_resolver import ResolvedBearerToken from nmp.common.config import AuthConfig from nmp.common.config.base import AccessKeyConfig, OIDCConfig, TokenSigningConfig @@ -47,6 +47,34 @@ def _test_client( yield TestClient(app) +def _auth_config_with_workload_exchange( + tmp_path, + *, + oidc_issuer: str = "https://sso.example.com", + additional_issuers: list[str] | None = None, + oidc_jwks_uri: str | None = "https://sso.example.com/jwks", + token_issuer: str | None = "http://testserver/apis/auth", +) -> AuthConfig: + private_key_file = tmp_path / "private.pem" + private_key_file.write_bytes(_private_key_pem()) + return AuthConfig( + enabled=True, + token_signing=TokenSigningConfig( + issuer=token_issuer, + key_id="test-workload", + private_key_file=str(private_key_file), + ), + oidc=OIDCConfig( + enabled=True, + issuer=oidc_issuer, + additional_issuers=additional_issuers or [], + jwks_uri=oidc_jwks_uri, + workload_token_exchange_enabled=True, + workload_audience="nemo-platform", + ), + ) + + def test_authenticate_access_key_returns_principal_headers(tmp_path): config = AuthConfig( enabled=True, @@ -131,6 +159,55 @@ def test_authenticate_callout_accepts_original_request_methods(tmp_path): resolver.assert_awaited_once() +def test_authenticate_oidc_access_token_with_actor_returns_direct_principal_headers(tmp_path): + config = AuthConfig( + enabled=True, + token_signing=TokenSigningConfig(private_key_file=str(tmp_path / "private.pem")), + oidc=OIDCConfig(enabled=True, issuer="https://sso.example.com", client_id="nemo-platform-cli"), + ) + (tmp_path / "private.pem").write_bytes(_private_key_pem()) + claims = TokenClaims( + subject="user:alice", + email="alice@example.test", + groups=["researchers"], + scopes=["models.read"], + raw_claims={ + "sub": "user:alice", + "email": "alice@example.test", + "groups": "researchers", + "act": {"sub": "service:jobs", "groups": "system:serviceaccounts"}, + }, + actor=ActorClaims(subject="service:jobs", groups=["system:serviceaccounts"]), + ) + resolved = ResolvedBearerToken(claims=claims, token_kind="oidc_access_token") + + with ( + _test_client(config) as client, + patch( + "nmp.core.auth.api.v2.authenticate.resolve_bearer_token", + new=AsyncMock(return_value=resolved), + ), + ): + response = client.get("/authenticate", headers={"Authorization": "Bearer oidc.token"}) + + assert response.status_code == 200 + assert response.json() == { + "principal": "user:alice", + "email": "alice@example.test", + "groups": ["researchers"], + "scopes": ["models.read"], + "jti": None, + "token_kind": "oidc_access_token", + } + assert response.headers["X-NMP-Principal-Id"] == "user:alice" + assert response.headers["X-NMP-Principal-Email"] == "alice@example.test" + assert response.headers["X-NMP-Principal-Groups"] == "researchers" + assert "X-NMP-Principal-On-Behalf-Of" not in response.headers + assert "X-NMP-Principal-On-Behalf-Of-Email" not in response.headers + assert "X-NMP-Principal-On-Behalf-Of-Groups" not in response.headers + assert response.headers["X-NMP-Scopes"] == "models.read" + + def test_authenticate_rejects_unresolved_bearer_token(tmp_path): config = AuthConfig(enabled=True, token_signing=TokenSigningConfig(private_key_file=str(tmp_path / "private.pem"))) (tmp_path / "private.pem").write_bytes(_private_key_pem()) @@ -197,12 +274,64 @@ def test_authenticate_workload_access_token_returns_principal_headers(tmp_path): assert response.headers["X-NMP-Scopes"] == "openid email groups" +def test_authenticate_delegated_workload_access_token_returns_obo_principal_headers(tmp_path): + private_key_file = tmp_path / "private.pem" + private_key_file.write_bytes(_private_key_pem()) + config = AuthConfig( + enabled=True, + token_signing=TokenSigningConfig( + issuer="http://testserver/apis/auth", + key_id="test-workload", + private_key_file=str(private_key_file), + ), + oidc=OIDCConfig( + workload_token_exchange_enabled=True, + workload_audience="nemo-platform", + ), + ) + signing_key = WorkloadTokenExchangeService().workload_signing_key(config) + now = datetime.now(tz=UTC) + token = jwt.encode( + { + "iss": "http://testserver/apis/auth", + "sub": "submitter@example.com", + "email": "submitter@example.com", + "aud": "nemo-platform", + "iat": now, + "nbf": now, + "exp": now + timedelta(minutes=5), + "scope": "openid email groups", + "groups": "workspace-editors", + "act": { + "sub": "system:serviceaccount:nemo:job", + "groups": ["system:serviceaccounts", "nemo-jobs"], + }, + }, + signing_key.private_key, + algorithm="RS256", + headers={"kid": signing_key.kid}, + ) + with _test_client(config) as client: + response = client.get( + "/authenticate", + headers={"Authorization": f"Bearer {token}"}, + ) + + assert response.status_code == 200 + assert response.json()["token_kind"] == "workload_access_token" + assert response.headers["X-NMP-Principal-Id"] == "system:serviceaccount:nemo:job" + assert response.headers["X-NMP-Principal-Groups"] == "system:serviceaccounts,nemo-jobs" + assert response.headers["X-NMP-Principal-On-Behalf-Of"] == "submitter@example.com" + assert response.headers["X-NMP-Principal-On-Behalf-Of-Email"] == "submitter@example.com" + assert response.headers["X-NMP-Principal-On-Behalf-Of-Groups"] == "workspace-editors" + assert response.headers["X-NMP-Scopes"] == "openid email groups" + + def test_authenticate_workload_subject_token_uses_resolver_callback(tmp_path): config = AuthConfig( enabled=True, token_signing=TokenSigningConfig(private_key_file=str(tmp_path / "private.pem")), oidc=OIDCConfig( - enabled=True, issuer="https://sso.example.com/application/o/nemo-cli/", client_id="nemo-platform-cli", workload_token_exchange_enabled=True, @@ -290,6 +419,119 @@ def test_authenticate_workload_access_token_surfaces_signing_key_misconfiguratio assert "Failed to load workload access token signing key" in caplog.text +def test_authenticate_fails_closed_when_oidc_issuer_matches_workload_token_issuer(tmp_path): + config = _auth_config_with_workload_exchange( + tmp_path, + oidc_issuer="http://testserver/apis/auth", + token_issuer="http://testserver/apis/auth", + ) + + with _test_client(config) as client: + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 500 + assert response.json()["detail"] == "Authentication token issuers are misconfigured" + + +def test_authenticate_fails_closed_when_additional_oidc_issuer_matches_workload_token_issuer(tmp_path): + config = _auth_config_with_workload_exchange( + tmp_path, + additional_issuers=["http://testserver/apis/auth"], + token_issuer="http://testserver/apis/auth", + ) + + with _test_client(config) as client: + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 500 + assert response.json()["detail"] == "Authentication token issuers are misconfigured" + + +def test_authenticate_fails_closed_when_oidc_jwks_uri_matches_workload_jwks_uri(tmp_path): + config = _auth_config_with_workload_exchange( + tmp_path, + oidc_jwks_uri="http://testserver/apis/auth/jwks", + ) + + with _test_client(config) as client: + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 500 + assert response.json()["detail"] == "Authentication token issuers are misconfigured" + + +def test_authenticate_fails_closed_when_discovered_oidc_jwks_uri_matches_workload_jwks_uri(tmp_path): + config = _auth_config_with_workload_exchange( + tmp_path, + oidc_jwks_uri=None, + ) + + with ( + _test_client(config) as client, + patch.object( + JWTValidator, + "_discover_oidc_config", + new=AsyncMock(return_value={"jwks_uri": "http://testserver/apis/auth/jwks"}), + ), + ): + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 500 + assert response.json()["detail"] == "Authentication token issuers are misconfigured" + + +def test_authenticate_fails_closed_when_oidc_jwks_contains_workload_public_key_material(tmp_path): + config = _auth_config_with_workload_exchange(tmp_path) + service = WorkloadTokenExchangeService() + mirrored_jwk = {**service.public_jwk(config), "kid": "idp-key"} + + with ( + _test_client(config, workload_token_exchange_service=service) as client, + patch.object(JWTValidator, "jwks", new=AsyncMock(return_value={"keys": [mirrored_jwk]})), + ): + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 500 + assert response.json()["detail"] == "Authentication token issuers are misconfigured" + + +def test_authenticate_allows_same_kid_when_oidc_jwks_key_material_differs(tmp_path): + config = _auth_config_with_workload_exchange(tmp_path) + other_private_key_file = tmp_path / "other-private.pem" + other_private_key_file.write_bytes(_private_key_pem()) + other_config = AuthConfig( + enabled=True, + token_signing=TokenSigningConfig( + issuer="https://other.example.test", + key_id="test-workload", + private_key_file=str(other_private_key_file), + ), + ) + service = WorkloadTokenExchangeService() + idp_jwk = service.public_jwk(other_config) + claims = TokenClaims( + subject="user:alice", + email="alice@example.test", + groups=["researchers"], + scopes=[], + raw_claims={}, + ) + resolved = ResolvedBearerToken(claims=claims, token_kind="oidc_access_token") + + with ( + _test_client(config, workload_token_exchange_service=service) as client, + patch.object(JWTValidator, "jwks", new=AsyncMock(return_value={"keys": [idp_jwk]})), + patch( + "nmp.core.auth.api.v2.authenticate.resolve_bearer_token", + new=AsyncMock(return_value=resolved), + ), + ): + response = client.get("/authenticate", headers={"Authorization": "Bearer token"}) + + assert response.status_code == 200 + assert response.headers["X-NMP-Principal-Id"] == "user:alice" + + def test_authenticate_openapi_documents_error_responses(tmp_path): config = AuthConfig(enabled=True, token_signing=TokenSigningConfig(private_key_file=str(tmp_path / "private.pem"))) (tmp_path / "private.pem").write_bytes(_private_key_pem()) diff --git a/services/core/auth/tests/test_workload_token_exchange.py b/services/core/auth/tests/test_workload_token_exchange.py index 93d2663809..d2153e72bd 100644 --- a/services/core/auth/tests/test_workload_token_exchange.py +++ b/services/core/auth/tests/test_workload_token_exchange.py @@ -3,8 +3,10 @@ import asyncio import json +from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any +from urllib.parse import urlencode import nmp.common.auth.signing_keys as signing_keys_mod import pytest @@ -13,9 +15,18 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from jwt.algorithms import RSAAlgorithm +from nmp.common.auth import AuthContext, Principal from nmp.common.auth.signing_keys import RSASigningKeyCache +from nmp.common.auth.workload_delegations import ( + DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE, + WorkloadDelegationEntity, + create_opaque_docker_proof_token, + docker_delegation_name, + reference_delegation_name, +) from nmp.common.config import AuthConfig, Configuration from nmp.common.config.base import AccessKeyConfig, OIDCConfig, TokenSigningConfig +from nmp.common.entities import SYSTEM_WORKSPACE, EntityNotFoundError from nmp.core.auth.api.v2 import workload_token_exchange as exchange from pydantic import ValidationError @@ -31,6 +42,32 @@ def exchange_service() -> exchange.WorkloadTokenExchangeService: return exchange.WorkloadTokenExchangeService() +class _FakeEntityClient: + def __init__(self) -> None: + self.entities: dict[str, WorkloadDelegationEntity] = {} + self.get_calls: list[str] = [] + + async def get( + self, + entity_type: type[WorkloadDelegationEntity], + name: str, + *, + workspace: str, + ) -> WorkloadDelegationEntity: + self.get_calls.append(name) + assert entity_type is WorkloadDelegationEntity + assert workspace == SYSTEM_WORKSPACE + try: + return self.entities[name] + except KeyError as exc: + raise EntityNotFoundError(f"Entity '{name}' not found") from exc + + +@pytest.fixture +def entity_client() -> _FakeEntityClient: + return _FakeEntityClient() + + @pytest.fixture def workload_signing_key() -> rsa.RSAPrivateKey: return rsa.generate_private_key(public_exponent=65537, key_size=2048) @@ -50,6 +87,89 @@ def _openapi(client: TestClient) -> dict[str, Any]: return app.openapi() +def _exchange_form( + subject_token: str, + *, + subject_token_type: str = exchange.JWT_TOKEN_TYPE, + audience: str | None = None, + scope: str | None = None, +) -> dict[str, str]: + data = { + "grant_type": exchange.TOKEN_EXCHANGE_GRANT_TYPE, + "client_id": "nemo-platform-workload", + "subject_token": subject_token, + "subject_token_type": subject_token_type, + } + if audience is not None: + data["audience"] = audience + if scope is not None: + data["scope"] = scope + return data + + +def _post_form_items(client: TestClient, items: list[tuple[str, str]]) -> Any: + return client.post( + "/token", + content=urlencode(items), + headers={"content-type": "application/x-www-form-urlencoded"}, + ) + + +def _decode_access_token( + token: str, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + *, + audience: str = "nemo-platform", +) -> dict[str, Any]: + signing_key = exchange_service.workload_signing_key(exchange_config) + return exchange.jwt.decode(token, signing_key.public_key, algorithms=["RS256"], audience=audience) + + +def _docker_delegation_name() -> str: + return docker_delegation_name( + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + ) + + +def _delegation_entity(**overrides: Any) -> WorkloadDelegationEntity: + delegation_name = _docker_delegation_name() + entity = WorkloadDelegationEntity( + name=delegation_name, + workspace=SYSTEM_WORKSPACE, + workload_subject=delegation_name, + workload_audience="nemo-platform", + workload_workspace="default", + job_id="job-123", + attempt_id="attempt-1", + step_id="step-a", + auth_context=AuthContext.from_principal( + Principal(id="creator@example.com", email="creator@example.com", groups=["workspace-editors"]) + ), + expires_at=datetime.now(timezone.utc) + timedelta(minutes=30), + ) + return entity.model_copy(update=overrides) + + +def _reference_delegation_entity(**overrides: Any) -> WorkloadDelegationEntity: + values = { + "name": reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:nemo-runs:job-runner", + bound_reference_name=exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value="pod-uid-123", + ), + "workload_subject": "system:serviceaccount:nemo-runs:job-runner", + "bound_reference_name": exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + "bound_reference_value": "pod-uid-123", + } + values.update(overrides) + return _delegation_entity(**values) + + @pytest.fixture def exchange_config(workload_signing_key: rsa.RSAPrivateKey, tmp_path) -> AuthConfig: private_key_file = tmp_path / "workload-token-private-key.pem" @@ -73,10 +193,15 @@ def exchange_config(workload_signing_key: rsa.RSAPrivateKey, tmp_path) -> AuthCo @pytest.fixture -def client(exchange_config: AuthConfig, exchange_service: exchange.WorkloadTokenExchangeService) -> TestClient: +def client( + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, +) -> TestClient: Configuration.set_override(exchange_config) app = FastAPI() app.dependency_overrides[exchange.get_workload_token_exchange_service] = lambda: exchange_service + app.dependency_overrides[exchange.get_entity_client] = lambda: entity_client app.include_router(exchange.router) return TestClient(app, raise_server_exceptions=False) @@ -227,8 +352,18 @@ def test_token_exchange_openapi_documents_form_request_and_token_response(client "subject_token_type", "requested_token_type", "audience", + "resource", + "actor_token", + "actor_token_type", "scope", } <= set(request_schema["properties"]) + assert DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE in request_schema["properties"]["subject_token_type"]["enum"] + for unsupported_property in ("resource", "actor_token", "actor_token_type"): + assert request_schema["properties"][unsupported_property] == { + "type": "string", + "description": "Unsupported for the jobs workload token exchange profile.", + "not": {}, + } assert set(operation["responses"]) == {"200", "400", "401"} assert operation["responses"]["200"]["description"] == "Successful Response" assert operation["responses"]["200"]["content"]["application/json"]["schema"] == { @@ -245,7 +380,7 @@ def test_token_exchange_openapi_documents_form_request_and_token_response(client error_schema = openapi["components"]["schemas"]["WorkloadTokenExchangeErrorResponse"] assert error_schema["properties"]["error"]["type"] == "string" error_description = error_schema["properties"]["error"]["description"] - for error_code in ("invalid_client", "invalid_request", "invalid_grant", "invalid_scope", "invalid_target"): + for error_code in ("invalid_client", "invalid_request", "invalid_scope", "invalid_target"): assert error_code in error_description response_schema = openapi["components"]["schemas"]["WorkloadTokenExchangeResponse"] assert { @@ -318,6 +453,180 @@ async def decode_subject_token(config: AuthConfig, subject_token: str, audience: } +@pytest.mark.parametrize( + "actor_fields", + [ + {"actor_token": "actor-token"}, + {"actor_token_type": exchange.JWT_TOKEN_TYPE}, + {"actor_token": "actor-token", "actor_token_type": exchange.JWT_TOKEN_TYPE}, + ], +) +def test_token_exchange_rejects_actor_token_parameters( + client: TestClient, + actor_fields: dict[str, str], +) -> None: + response = client.post( + "/token", + data={**_exchange_form("subject-token"), **actor_fields}, + ) + + assert response.status_code == 400 + assert response.json() == { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + } + + +@pytest.mark.parametrize( + ("field", "values", "expected_response"), + [ + ( + "actor_token", + [""], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "actor_token", + [" "], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "actor_token", + ["", " "], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "actor_token_type", + [""], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "actor_token_type", + [" "], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "actor_token_type", + ["", " "], + { + "error": "invalid_request", + "error_description": "actor_token is not supported for this workload token exchange profile", + }, + ), + ( + "resource", + [""], + { + "error": "invalid_target", + "error_description": "resource is not supported for workload token exchange", + }, + ), + ( + "resource", + [" "], + { + "error": "invalid_target", + "error_description": "resource is not supported for workload token exchange", + }, + ), + ( + "resource", + ["", " "], + { + "error": "invalid_target", + "error_description": "resource is not supported for workload token exchange", + }, + ), + ], +) +def test_token_exchange_rejects_unsupported_fields_by_presence( + client: TestClient, + field: str, + values: list[str], + expected_response: dict[str, str], +) -> None: + response = _post_form_items( + client, + [ + *_exchange_form("subject-token").items(), + *((field, value) for value in values), + ], + ) + + assert response.status_code == 400 + assert response.json() == expected_response + + +def test_token_exchange_rejects_resource_target_before_subject_validation(client: TestClient) -> None: + response = client.post( + "/token", + data={**_exchange_form("subject-token"), "resource": "https://nmp.example.com/apis/secrets"}, + ) + + assert response.status_code == 400 + assert response.json() == { + "error": "invalid_target", + "error_description": "resource is not supported for workload token exchange", + } + + +def test_token_exchange_rejects_multiple_audiences_before_subject_validation(client: TestClient) -> None: + response = _post_form_items( + client, + [ + *_exchange_form("subject-token").items(), + ("audience", "nemo-platform"), + ("audience", "extra-audience"), + ], + ) + + assert response.status_code == 400 + assert response.json() == { + "error": "invalid_target", + "error_description": "Only one audience is supported for workload token exchange", + } + + +def test_token_exchange_accepts_single_allowed_audience( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange_config.oidc.workload_allowed_audiences.append("extra-audience") + + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return {"sub": "workload-subject"} + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = client.post("/token", data=_exchange_form("subject-token", audience="extra-audience")) + + assert response.status_code == 200 + claims = _decode_access_token( + response.json()["access_token"], + exchange_config, + exchange_service, + audience="extra-audience", + ) + assert claims["aud"] == "extra-audience" + + def test_token_exchange_mints_access_token_signed_by_configured_key( client: TestClient, exchange_config: AuthConfig, @@ -347,7 +656,9 @@ async def decode_subject_token(config: AuthConfig, subject_token: str, audience: ) assert response.status_code == 200 - access_token = response.json()["access_token"] + response_body = response.json() + assert response_body["expires_in"] == exchange_config.oidc.workload_token_ttl_seconds + access_token = response_body["access_token"] signing_key = exchange_service.workload_signing_key(exchange_config) assert exchange.jwt.get_unverified_header(access_token)["kid"] == signing_key.kid claims = exchange.jwt.decode(access_token, signing_key.public_key, algorithms=["RS256"], audience="nemo-platform") @@ -357,6 +668,318 @@ async def decode_subject_token(config: AuthConfig, subject_token: str, audience: assert claims["groups"] == "svc-group,system:serviceaccounts" +def test_token_exchange_intersects_requested_scope_with_configured_workload_scope( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return {"sub": "workload-subject"} + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = client.post( + "/token", + data=_exchange_form("subject-token", scope="openid admin:write email"), + ) + + assert response.status_code == 200 + response_body = response.json() + assert response_body["scope"] == "openid email" + claims = _decode_access_token(response_body["access_token"], exchange_config, exchange_service) + assert claims["scope"] == "openid email" + + +def test_token_exchange_omits_scope_when_requested_scope_is_not_allowed( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return {"sub": "workload-subject"} + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = client.post( + "/token", + data=_exchange_form("subject-token", scope="admin:write"), + ) + + assert response.status_code == 200 + response_body = response.json() + assert response_body["scope"] is None + claims = _decode_access_token(response_body["access_token"], exchange_config, exchange_service) + assert "scope" not in claims + + +def test_opaque_docker_proof_token_with_missing_row_returns_invalid_request(client: TestClient) -> None: + subject_token, _token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + + response = client.post( + "/token", + data=_exchange_form(subject_token, subject_token_type=DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE), + ) + + assert response.status_code == 400 + assert response.json()["error"] == "invalid_request" + + +def test_opaque_docker_proof_token_mints_delegated_access_token( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, +) -> None: + subject_token, token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + delegation = _delegation_entity(opaque_subject_token_hash=token_hash) + entity_client.entities[delegation.name] = delegation + + response = client.post( + "/token", + data=_exchange_form(subject_token, subject_token_type=DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE), + ) + + assert response.status_code == 200 + claims = _decode_access_token(response.json()["access_token"], exchange_config, exchange_service) + assert claims["sub"] == "creator@example.com" + assert claims["act"] == {"sub": delegation.name} + + +def test_opaque_docker_proof_token_requires_private_subject_token_type( + client: TestClient, + entity_client: _FakeEntityClient, +) -> None: + subject_token, token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + delegation = _delegation_entity(opaque_subject_token_hash=token_hash) + entity_client.entities[delegation.name] = delegation + + response = client.post("/token", data=_exchange_form(subject_token, subject_token_type=exchange.JWT_TOKEN_TYPE)) + + assert response.status_code == 400 + assert response.json()["error"] == "invalid_request" + assert entity_client.get_calls == [] + + +def test_delegated_access_token_expiry_is_capped_by_delegation_expiry( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, +) -> None: + exchange_config.oidc.workload_token_ttl_seconds = 600 + subject_token, token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + delegation_expires_at = datetime.now(timezone.utc) + timedelta(seconds=120) + delegation = _delegation_entity( + opaque_subject_token_hash=token_hash, + expires_at=delegation_expires_at, + ) + entity_client.entities[delegation.name] = delegation + + response = client.post( + "/token", + data=_exchange_form(subject_token, subject_token_type=DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE), + ) + + assert response.status_code == 200 + response_body = response.json() + assert 0 < response_body["expires_in"] <= 120 + claims = _decode_access_token(response_body["access_token"], exchange_config, exchange_service) + assert claims["exp"] <= int(delegation_expires_at.timestamp()) + + +def test_delegated_access_token_expiry_keeps_configured_ttl_when_delegation_expires_later( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, +) -> None: + exchange_config.oidc.workload_token_ttl_seconds = 120 + subject_token, token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + delegation = _delegation_entity( + opaque_subject_token_hash=token_hash, + expires_at=datetime.now(timezone.utc) + timedelta(seconds=600), + ) + entity_client.entities[delegation.name] = delegation + + response = client.post( + "/token", + data=_exchange_form(subject_token, subject_token_type=DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE), + ) + + assert response.status_code == 200 + response_body = response.json() + assert 0 < response_body["expires_in"] <= 120 + claims = _decode_access_token(response_body["access_token"], exchange_config, exchange_service) + assert claims["exp"] - claims["iat"] <= 120 + + +@pytest.mark.parametrize( + "delegation_overrides", + [ + {"opaque_subject_token_hash": "v1:sha256:wrong"}, + {"expires_at": datetime.now(timezone.utc) - timedelta(seconds=1)}, + {"revoked_at": datetime.now(timezone.utc)}, + { + "bound_reference_name": exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + "bound_reference_value": "pod-uid-123", + }, + ], +) +def test_opaque_docker_proof_token_rejects_ineligible_rows( + client: TestClient, + entity_client: _FakeEntityClient, + delegation_overrides: dict[str, Any], +) -> None: + subject_token, token_hash = create_opaque_docker_proof_token(_docker_delegation_name()) + entity_overrides = {"opaque_subject_token_hash": token_hash, **delegation_overrides} + delegation = _delegation_entity(**entity_overrides) + entity_client.entities[delegation.name] = delegation + + response = client.post( + "/token", + data=_exchange_form(subject_token, subject_token_type=DOCKER_OPAQUE_WORKLOAD_PROOF_TOKEN_TYPE), + ) + + assert response.status_code == 400 + assert response.json()["error"] == "invalid_request" + + +def test_kubernetes_tokenreview_without_pod_uid_keeps_workload_only_exchange( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return { + "sub": "system:serviceaccount:nemo-runs:job-runner", + "groups": ["system:serviceaccounts"], + } + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = client.post("/token", data=_exchange_form("kubernetes-token")) + + assert response.status_code == 200 + claims = _decode_access_token(response.json()["access_token"], exchange_config, exchange_service) + assert entity_client.get_calls == [] + assert claims["sub"] == "system:serviceaccount:nemo-runs:job-runner" + assert claims["groups"] == "system:serviceaccounts" + assert "act" not in claims + + +def test_kubernetes_reference_row_mints_delegated_access_token( + client: TestClient, + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + entity_client: _FakeEntityClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + delegation = _reference_delegation_entity() + entity_client.entities[delegation.name] = delegation + + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return { + "sub": "system:serviceaccount:nemo-runs:job-runner", + "groups": ["system:serviceaccounts"], + exchange._BOUND_REFERENCE_NAME_CLAIM: exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + exchange._BOUND_REFERENCE_VALUE_CLAIM: "pod-uid-123", + exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM: exchange._BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES, + } + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = client.post("/token", data=_exchange_form("kubernetes-token")) + + assert response.status_code == 200 + claims = _decode_access_token(response.json()["access_token"], exchange_config, exchange_service) + assert entity_client.get_calls == [delegation.name] + assert claims["sub"] == "creator@example.com" + assert claims["act"] == { + "sub": "system:serviceaccount:nemo-runs:job-runner", + "groups": "system:serviceaccounts", + } + + +def test_kubernetes_reference_lookup_retries_missing_row_without_waiting( + exchange_config: AuthConfig, + entity_client: _FakeEntityClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + sleeps: list[float] = [] + + async def sleep(seconds: float) -> None: + sleeps.append(seconds) + + exchange_service = exchange.WorkloadTokenExchangeService( + delegation_lookup_retry_timeout_seconds=0.3, + delegation_lookup_retry_interval_seconds=0.1, + sleep=sleep, + ) + app = FastAPI() + Configuration.set_override(exchange_config) + app.dependency_overrides[exchange.get_workload_token_exchange_service] = lambda: exchange_service + app.dependency_overrides[exchange.get_entity_client] = lambda: entity_client + app.include_router(exchange.router) + retry_client = TestClient(app, raise_server_exceptions=False) + + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return { + "sub": "system:serviceaccount:nemo-runs:job-runner", + exchange._BOUND_REFERENCE_NAME_CLAIM: exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + exchange._BOUND_REFERENCE_VALUE_CLAIM: "pod-uid-123", + exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM: exchange._BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES, + } + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = retry_client.post("/token", data=_exchange_form("kubernetes-token")) + + assert response.status_code == 400 + assert response.json()["error"] == "invalid_request" + assert len(entity_client.get_calls) == 4 + assert sleeps == [0.1, 0.1, 0.1] + + +def test_kubernetes_reference_lookup_default_retry_budget_uses_five_seconds( + exchange_config: AuthConfig, + entity_client: _FakeEntityClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + sleeps: list[float] = [] + + async def sleep(seconds: float) -> None: + sleeps.append(seconds) + + exchange_service = exchange.WorkloadTokenExchangeService(sleep=sleep) + app = FastAPI() + Configuration.set_override(exchange_config) + app.dependency_overrides[exchange.get_workload_token_exchange_service] = lambda: exchange_service + app.dependency_overrides[exchange.get_entity_client] = lambda: entity_client + app.include_router(exchange.router) + retry_client = TestClient(app, raise_server_exceptions=False) + + async def decode_subject_token(config: AuthConfig, subject_token: str, audience: str) -> dict[str, Any]: + return { + "sub": "system:serviceaccount:nemo-runs:job-runner", + exchange._BOUND_REFERENCE_NAME_CLAIM: exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + exchange._BOUND_REFERENCE_VALUE_CLAIM: "pod-uid-123", + exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM: exchange._BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES, + } + + monkeypatch.setattr(exchange_service, "decode_subject_token", decode_subject_token) + + response = retry_client.post("/token", data=_exchange_form("kubernetes-token")) + + assert response.status_code == 400 + assert response.json()["error"] == "invalid_request" + assert len(entity_client.get_calls) == 51 + assert sleeps == [0.1] * 50 + + def test_validated_audience_accepts_configured_allowlist(exchange_config: AuthConfig) -> None: exchange_config.oidc.workload_allowed_audiences.append("extra-audience") @@ -541,16 +1164,20 @@ def _signed_subject_token( issuer: str | None = None, private_key: rsa.RSAPrivateKey | None = None, key_id: str | None = None, + extra_claims: dict[str, Any] | None = None, ) -> str: token_issuer = issuer or config.oidc.workload_subject_issuers[0] signing_key = exchange_service.workload_signing_key(config) + claims = { + "iss": token_issuer, + "sub": "authentik-user", + "aud": audience, + "exp": int(exchange.time.time()) + 300, + } + if extra_claims: + claims.update(extra_claims) return exchange.jwt.encode( - { - "iss": token_issuer, - "sub": "authentik-user", - "aud": audience, - "exp": int(exchange.time.time()) + 300, - }, + claims, private_key or signing_key.private_key, algorithm="RS256", headers={"kid": key_id or signing_key.kid}, @@ -608,6 +1235,30 @@ def test_jwt_subject_token_decoder_fetches_configured_jwks( } +def test_jwt_subject_token_decoder_strips_untrusted_bound_reference_claims( + exchange_config: AuthConfig, + exchange_service: exchange.WorkloadTokenExchangeService, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _mock_subject_jwks_client(exchange_config, exchange_service, monkeypatch) + subject_token = _signed_subject_token( + exchange_config, + exchange_service, + audience="nemo-platform-workload", + extra_claims={ + exchange._BOUND_REFERENCE_NAME_CLAIM: exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + exchange._BOUND_REFERENCE_VALUE_CLAIM: "pod-uid-123", + exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM: exchange._BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES, + }, + ) + + claims = asyncio.run(exchange_service.decode_jwt_subject_token(exchange_config, subject_token)) + + assert exchange._BOUND_REFERENCE_NAME_CLAIM not in claims + assert exchange._BOUND_REFERENCE_VALUE_CLAIM not in claims + assert exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM not in claims + + def test_jwt_subject_token_decoder_caches_configured_jwks( exchange_config: AuthConfig, exchange_service: exchange.WorkloadTokenExchangeService, @@ -854,3 +1505,102 @@ async def post(self, url: str, **kwargs) -> _FakeResponse: "Content-Type": "application/json", }, } + + +def test_kubernetes_subject_token_decoder_preserves_single_pod_uid_reference( + exchange_config: AuthConfig, + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange_config.oidc.workload_kubernetes_token_review_enabled = True + monkeypatch.setenv("KUBERNETES_SERVICE_HOST", "kubernetes.default.svc") + monkeypatch.setenv("KUBERNETES_SERVICE_PORT", "443") + monkeypatch.setattr(exchange, "_kubernetes_reviewer_credentials", lambda: ("reviewer-token", "/tmp/ca.crt")) + + class FakeAsyncClient: + def __init__(self, *, timeout: float, verify: str) -> None: + return None + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback) -> None: + return None + + async def post(self, url: str, **kwargs) -> _FakeResponse: + return _FakeResponse( + { + "status": { + "authenticated": True, + "user": { + "username": "system:serviceaccount:nemo-runs:job-runner", + "groups": ["system:serviceaccounts"], + "extra": { + exchange.KUBERNETES_POD_UID_REFERENCE_NAME: ["pod-uid-123"], + }, + }, + } + } + ) + + monkeypatch.setattr(exchange.httpx, "AsyncClient", FakeAsyncClient) + + claims = asyncio.run( + exchange._decode_kubernetes_subject_token(exchange_config, "kubernetes-subject-token", "nemo-platform") + ) + + assert claims == { + "sub": "system:serviceaccount:nemo-runs:job-runner", + "groups": ["system:serviceaccounts"], + exchange._BOUND_REFERENCE_NAME_CLAIM: exchange.KUBERNETES_POD_UID_REFERENCE_NAME, + exchange._BOUND_REFERENCE_VALUE_CLAIM: "pod-uid-123", + exchange._BOUND_REFERENCE_TRUSTED_SOURCE_CLAIM: exchange._BOUND_REFERENCE_TRUSTED_SOURCE_KUBERNETES, + } + + +@pytest.mark.parametrize( + "pod_uids", + [ + [], + ["pod-uid-1", "pod-uid-2"], + ], +) +def test_kubernetes_subject_token_decoder_rejects_ambiguous_pod_uid_reference( + exchange_config: AuthConfig, + monkeypatch: pytest.MonkeyPatch, + pod_uids: list[str], +) -> None: + exchange_config.oidc.workload_kubernetes_token_review_enabled = True + monkeypatch.setenv("KUBERNETES_SERVICE_HOST", "kubernetes.default.svc") + monkeypatch.setattr(exchange, "_kubernetes_reviewer_credentials", lambda: ("reviewer-token", "/tmp/ca.crt")) + + class FakeAsyncClient: + def __init__(self, *, timeout: float, verify: str) -> None: + return None + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback) -> None: + return None + + async def post(self, url: str, **kwargs) -> _FakeResponse: + return _FakeResponse( + { + "status": { + "authenticated": True, + "user": { + "username": "system:serviceaccount:nemo-runs:job-runner", + "extra": { + exchange.KUBERNETES_POD_UID_REFERENCE_NAME: pod_uids, + }, + }, + } + } + ) + + monkeypatch.setattr(exchange.httpx, "AsyncClient", FakeAsyncClient) + + with pytest.raises(exchange._InvalidGrantError): + asyncio.run( + exchange._decode_kubernetes_subject_token(exchange_config, "kubernetes-subject-token", "nemo-platform") + ) diff --git a/services/core/jobs/jobs-launcher/cmd/workload_auth.go b/services/core/jobs/jobs-launcher/cmd/workload_auth.go index 15123d85b1..fe18df463a 100644 --- a/services/core/jobs/jobs-launcher/cmd/workload_auth.go +++ b/services/core/jobs/jobs-launcher/cmd/workload_auth.go @@ -17,15 +17,17 @@ import ( ) const ( - nmpBaseURLEnv = "NMP_BASE_URL" - workloadIdentityTokenFileEnv = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE" - tokenExchangeGrantType = "urn:ietf:params:oauth:grant-type:token-exchange" - jwtTokenType = "urn:ietf:params:oauth:token-type:jwt" - accessTokenType = "urn:ietf:params:oauth:token-type:access_token" - workloadAuthRequestTimeoutSeconds = 30 - maxAuthResponseBodyBytes = 64 * 1024 - workloadAuthRefreshMarginFraction = 5 - workloadAuthMaxRefreshMargin = time.Minute + nmpBaseURLEnv = "NMP_BASE_URL" + workloadIdentityTokenFileEnv = "NMP_WORKLOAD_IDENTITY_TOKEN_FILE" + tokenExchangeGrantType = "urn:ietf:params:oauth:grant-type:token-exchange" + jwtTokenType = "urn:ietf:params:oauth:token-type:jwt" + accessTokenType = "urn:ietf:params:oauth:token-type:access_token" + dockerOpaqueWorkloadProofTokenType = "urn:nvidia:nemo:params:oauth:token-type:docker-opaque-workload-proof" + dockerOpaqueWorkloadProofPrefix = "nmp_obo_v1." + workloadAuthRequestTimeoutSeconds = 30 + maxAuthResponseBodyBytes = 64 * 1024 + workloadAuthRefreshMarginFraction = 5 + workloadAuthMaxRefreshMargin = time.Minute ) type authDiscoveryResponse struct { @@ -215,6 +217,13 @@ func readSubjectToken(path string) (string, error) { return token, nil } +func subjectTokenTypeForExchange(subjectToken string) string { + if strings.HasPrefix(subjectToken, dockerOpaqueWorkloadProofPrefix) { + return dockerOpaqueWorkloadProofTokenType + } + return jwtTokenType +} + func exchangeWorkloadToken( ctx context.Context, tokenEndpoint string, @@ -227,7 +236,7 @@ func exchangeWorkloadToken( form.Set("grant_type", tokenExchangeGrantType) form.Set("client_id", clientID) form.Set("subject_token", subjectToken) - form.Set("subject_token_type", jwtTokenType) + form.Set("subject_token_type", subjectTokenTypeForExchange(subjectToken)) form.Set("requested_token_type", accessTokenType) if discovery.WorkloadAudience != "" { form.Set("audience", discovery.WorkloadAudience) diff --git a/services/core/jobs/jobs-launcher/cmd/workload_auth_test.go b/services/core/jobs/jobs-launcher/cmd/workload_auth_test.go index 354aa1a66a..4f28849914 100644 --- a/services/core/jobs/jobs-launcher/cmd/workload_auth_test.go +++ b/services/core/jobs/jobs-launcher/cmd/workload_auth_test.go @@ -51,6 +51,15 @@ func assertLogNotContains(t *testing.T, logOutput string, unexpected string) { } } +func TestSubjectTokenTypeForExchange(t *testing.T) { + if got := subjectTokenTypeForExchange("subject-token"); got != jwtTokenType { + t.Fatalf("expected JWT subject token type, got %q", got) + } + if got := subjectTokenTypeForExchange("nmp_obo_v1.delegation.secret"); got != dockerOpaqueWorkloadProofTokenType { + t.Fatalf("expected Docker opaque subject token type, got %q", got) + } +} + func TestGetOTLPLogWorkloadAuthHeadersReturnsAuthorizationWithoutMutatingEnv(t *testing.T) { subjectTokenPath := filepath.Join(t.TempDir(), "subject.jwt") if err := os.WriteFile(subjectTokenPath, []byte("subject-token\n"), 0o600); err != nil { diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/backends/docker.py b/services/core/jobs/src/nmp/core/jobs/controllers/backends/docker.py index 43dba2f13b..c4e373249d 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/backends/docker.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/backends/docker.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import asyncio import datetime import hashlib import io @@ -13,14 +14,17 @@ import uuid from abc import abstractmethod from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass -from typing import Any, Generic, Self, TypeVar +from dataclasses import dataclass, field +from typing import Any, Generic, TypeVar import docker.types from docker.errors import APIError, ImageNotFound, NotFound from docker.models.containers import Container from docker.types import LogConfig, Mount +from nemo_platform import AsyncNeMoPlatform from nemo_platform_plugin.capabilities import CapabilityUnavailableError, probe_docker +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.entities.client import AsyncEntitiesClient from nemo_platform_plugin.jobs.execution_profiles import ( DockerJobExecutionProfile as PluginDockerJobExecutionProfile, ) @@ -41,9 +45,16 @@ PlatformJobStepWithContext, PlatformJobTaskUpdate, ) -from nmp.common.auth import AuthContext +from nmp.common.auth import ( + AuthContext, + WorkloadDelegationEntity, + WorkloadDelegationStore, + create_opaque_docker_proof_token, + docker_delegation_name, +) from nmp.common.config import get_platform_config, nmp_user_data_dir from nmp.common.docker.gpu_pool import GPUAllocationError +from nmp.common.entities import SYSTEM_WORKSPACE, EntityClient from nmp.common.jobs.constants import ( CONFIG_TASK_STORAGE_PATH_ENVVAR, DEFAULT_CONFIG_STORAGE_PATH, @@ -112,15 +123,12 @@ SchedulingDeferred, ) from nmp.core.jobs.controllers.backends.workload_tokens import ( - DEFAULT_SUBJECT_TOKEN_REFRESH_MARGIN_SECONDS, - DEFAULT_SUBJECT_TOKEN_TTL_SECONDS, - OAuthPasswordGrantSubjectTokenIssuer, - SubjectTokenIssuer, - SubjectTokenRefreshLoop, build_token_archive, + create_authenticated_async_nmp_sdk, + workload_delegation_expires_at, ) from opentelemetry import trace -from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic import BaseModel, ConfigDict, Field import docker @@ -150,7 +158,6 @@ def k8s_shm_quantity_to_docker(quantity: str) -> str: # Default is 30 seconds which matches the Kubernetes default grace period for pod termination. DOCKER_STOP_TIMEOUT = int(os.getenv("NEMO_JOBS_DEFAULT_DOCKER_STOP_TIMEOUT", "30")) NMP_JOBS_DOCKER_OWNER_ID_ENVVAR = "NMP_JOBS_DOCKER_OWNER_ID" -DOCKER_WORKLOAD_IDENTITY_PASSWORD_ENV_VAR_DEFAULT = "AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD" DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL = "nmp.nvidia.com/workload_identity_token_file" DOCKER_WORKLOAD_IDENTITY_VOLUME_LABEL = "nmp.nvidia.com/workload_identity_volume" @@ -178,6 +185,13 @@ class DockerTimestampParseResult: is_zero: bool +@dataclass(frozen=True, slots=True) +class DockerWorkloadProofToken: + token: str = field(repr=False) + expires_at: datetime.datetime + opaque_subject_token_hash: str | None = None + + # Server-side override: the default network name comes from the # ``NEMO_JOBS_DEFAULT_DOCKER_NETWORK`` env var (used by quickstart and e2e). # No docstring on purpose — a docstring would surface as the schema @@ -189,75 +203,10 @@ class DockerJobNetworkConfig(PluginDockerJobNetworkConfig): class DockerWorkloadIdentityConfig(BaseModel): - """Docker-only subject token issuer configuration for workload identity.""" + """Docker workload identity configuration.""" model_config = ConfigDict(extra="forbid") - enabled: bool | None = Field( - default=None, - description="Enable Docker workload identity token-file injection. Defaults to auth.oidc.workload_token_exchange_enabled.", - ) - token_endpoint: str | None = Field( - default=None, - description="OAuth token endpoint used by the Docker demo issuer. Defaults to auth.oidc.token_endpoint.", - ) - client_id: str | None = Field( - default=None, - description="OAuth client ID used by the Docker demo issuer. Defaults to auth.oidc.workload_client_id or auth.oidc.client_id.", - ) - client_secret: str | None = Field( - default=None, - description="OAuth client secret for the Docker demo issuer.", - json_schema_extra={"format": "password", "writeOnly": True}, - ) - username: str | None = Field(default=None, description="Username for the Docker demo issuer password grant.") - password_env_var: str = Field( - default=DOCKER_WORKLOAD_IDENTITY_PASSWORD_ENV_VAR_DEFAULT, - description="Controller environment variable that contains the Docker demo issuer password grant shared secret.", - ) - scope: str | None = Field(default=None, description="OAuth scope for the Docker demo issuer.") - subject_token_ttl_seconds: int = Field( - default_factory=lambda: int( - os.environ.get("NMP_WORKLOAD_IDENTITY_TOKEN_TTL_SECONDS", DEFAULT_SUBJECT_TOKEN_TTL_SECONDS) - ), - ge=1, - validate_default=True, - description="Fallback subject-token lifetime when the Docker demo issuer response omits expires_in.", - ) - refresh_margin_seconds: int = Field( - default=DEFAULT_SUBJECT_TOKEN_REFRESH_MARGIN_SECONDS, - ge=0, - validate_default=True, - description="Seconds before subject-token expiry when the Docker refresher issues a replacement token.", - ) - - @field_validator("password_env_var") - @classmethod - def validate_password_env_var(cls, value: str) -> str: - value = value.strip() - if not value: - raise ValueError("password_env_var must name a non-empty environment variable") - if value[0].isdigit() or not all(char == "_" or char.isalnum() for char in value): - raise ValueError("password_env_var must be a valid environment variable name") - return value - - @model_validator(mode="after") - def validate_refresh_margin_before_expiry(self) -> Self: - if self.refresh_margin_seconds >= self.subject_token_ttl_seconds: - raise ValueError("refresh_margin_seconds must be less than subject_token_ttl_seconds") - return self - - -def _resolve_docker_workload_identity_scope(workload_config: DockerWorkloadIdentityConfig, oidc_config: Any) -> str: - return workload_config.scope or getattr(oidc_config, "workload_scope", None) or "openid email groups" - - -def _resolve_docker_workload_identity_password(workload_config: DockerWorkloadIdentityConfig) -> str | None: - password = os.environ.get(workload_config.password_env_var) - if password: - return password - return None - class DockerJobExecutionProfileConfig(PluginDockerJobExecutionProfileConfig): """Configuration for Docker Job execution profile.""" @@ -268,7 +217,7 @@ class DockerJobExecutionProfileConfig(PluginDockerJobExecutionProfileConfig): ) workload_identity: DockerWorkloadIdentityConfig = Field( default_factory=DockerWorkloadIdentityConfig, - description="Docker workload identity subject-token issuer configuration.", + description="Docker workload identity configuration.", ) @@ -289,7 +238,6 @@ def init(self) -> None: self._jobs_controller_instance_id = _resolve_jobs_controller_instance_id() self._container_start_admission = threading.BoundedSemaphore(DOCKER_CONTAINER_START_WORKERS) self._container_run_threadpool = ThreadPoolExecutor(max_workers=DOCKER_CONTAINER_START_WORKERS) - self._workload_identity_refreshers: dict[str, SubjectTokenRefreshLoop] = {} # Short probe first — avoid docker.from_env(timeout=180) hanging when the # daemon is down. CapabilityUnavailableError is soft-skipped by the registry. probe = probe_docker() @@ -308,62 +256,212 @@ def init(self) -> None: def shutdown(self) -> None: self._container_run_threadpool.shutdown(wait=True) - self._stop_all_workload_identity_refreshers() self._client.close() def _is_workload_identity_enabled(self) -> bool: - configured = self._execution_profile_config.workload_identity.enabled - if configured is not None: - return configured return is_workload_identity_token_exchange_enabled() - def _create_docker_subject_token_issuer(self) -> SubjectTokenIssuer: - workload_config = self._execution_profile_config.workload_identity + def _should_enable_workload_identity_for_step(self, step: PlatformJobStepWithContext) -> bool: + if not self._is_workload_identity_enabled(): + return False + if step.auth_context is None: + logger.debug( + "Docker workload identity is enabled, but the job step has no auth_context", + extra={"workspace": step.workspace, "job": step.job, "step": step.name}, + ) + return False + return True + + @staticmethod + def _workload_delegation_audience() -> str: try: from nmp.common.config import get_auth_config oidc_config = get_auth_config().oidc except Exception: - oidc_config = None + logger.warning( + "Could not resolve auth config for Docker workload delegation audience; using the default audience", + exc_info=True, + ) + return "nemo-platform" + return ( + getattr(oidc_config, "workload_audience", None) or getattr(oidc_config, "audience", None) or "nemo-platform" + ) - token_endpoint = workload_config.token_endpoint or getattr(oidc_config, "token_endpoint", None) - client_id = ( - workload_config.client_id - or getattr(oidc_config, "workload_client_id", None) - or getattr(oidc_config, "client_id", None) + def _create_async_nmp_sdk(self) -> AsyncNeMoPlatform: + return create_authenticated_async_nmp_sdk( + self._nmp_sdk, + missing_headers_message="Docker workload delegation requires authenticated SDK headers", ) - password = _resolve_docker_workload_identity_password(workload_config) - scope = _resolve_docker_workload_identity_scope(workload_config, oidc_config) - missing = [ - name - for name, value in { - "token_endpoint": token_endpoint, - "client_id": client_id, - "username": workload_config.username, - workload_config.password_env_var: password, - }.items() - if not value - ] - if missing: - raise JobStorageError( - "Docker workload identity token exchange is enabled, but the Docker profile is missing issuer " - f"configuration: {', '.join(missing)}" + + async def _register_workload_delegation_async(self, entity: WorkloadDelegationEntity) -> None: + async_sdk = self._create_async_nmp_sdk() + try: + entity_client = EntityClient(client_from_platform(async_sdk, AsyncEntitiesClient)).as_service( + "jobs", internal=True + ) + store = WorkloadDelegationStore(entity_client) + await store.register( + entity, + require_opaque_subject_token_hash=entity.opaque_subject_token_hash is not None, + ) + finally: + await async_sdk.close() + + async def _revoke_workload_delegation_async(self, delegation_name: str) -> None: + async_sdk = self._create_async_nmp_sdk() + try: + entity_client = EntityClient(client_from_platform(async_sdk, AsyncEntitiesClient)).as_service( + "jobs", internal=True + ) + store = WorkloadDelegationStore(entity_client) + await store.revoke(delegation_name) + finally: + await async_sdk.close() + + def _register_workload_delegation(self, entity: WorkloadDelegationEntity) -> None: + try: + asyncio.get_running_loop() + except RuntimeError: + asyncio.run(self._register_workload_delegation_async(entity)) + return + raise JobStorageError("Docker workload delegation registration cannot run from an active asyncio event loop") + + def _revoke_workload_delegation(self, delegation_name: str) -> None: + try: + asyncio.get_running_loop() + except RuntimeError: + asyncio.run(self._revoke_workload_delegation_async(delegation_name)) + return + raise JobStorageError("Docker workload delegation revocation cannot run from an active asyncio event loop") + + def _provision_docker_workload_proof_token(self, delegation_name: str) -> DockerWorkloadProofToken: + token, token_hash = create_opaque_docker_proof_token(delegation_name) + expires_at = workload_delegation_expires_at( + ttl_seconds_active=self._execution_profile_config.ttl_seconds_active + ) + return DockerWorkloadProofToken( + token=token, + expires_at=expires_at, + opaque_subject_token_hash=token_hash, + ) + + def _build_docker_workload_delegation( + self, + *, + step: PlatformJobStepWithContext, + delegation_name: str, + proof_token: DockerWorkloadProofToken, + ) -> WorkloadDelegationEntity: + if step.auth_context is None: + raise JobStorageError("Docker workload identity requires a job auth_context for on-behalf-of delegation") + + auth_context = AuthContext.model_validate(step.auth_context.model_dump(mode="python", exclude_none=True)) + return WorkloadDelegationEntity( + name=delegation_name, + workspace=SYSTEM_WORKSPACE, + workload_subject=delegation_name, + workload_audience=self._workload_delegation_audience(), + workload_workspace=step.workspace, + job_id=step.job, + attempt_id=step.attempt_id, + step_id=step.id, + auth_context=auth_context, + opaque_subject_token_hash=proof_token.opaque_subject_token_hash, + expires_at=proof_token.expires_at, + ) + + def _prepare_workload_identity_for_step( + self, + *, + step: PlatformJobStepWithContext, + workload_identity_volume_name: str, + ) -> str: + delegation_name = docker_delegation_name( + workload_workspace=step.workspace, + job_id=step.job, + attempt_id=step.attempt_id, + step_id=step.id, + ) + proof_token = self._provision_docker_workload_proof_token(delegation_name) + delegation = self._build_docker_workload_delegation( + step=step, + delegation_name=delegation_name, + proof_token=proof_token, + ) + registered = False + try: + self._register_workload_delegation(delegation) + registered = True + self._write_workload_identity_subject_token(workload_identity_volume_name, proof_token.token) + except Exception: + if registered: + try: + self._revoke_workload_delegation(delegation_name) + except Exception: + logger.exception( + "Failed to revoke Docker workload delegation after token provisioning failure", + extra={"delegation_name": delegation_name}, + ) + raise + return delegation_name + + def _revoke_workload_delegation_after_failed_start(self, delegation_name: str | None) -> None: + if not delegation_name: + return + try: + self._revoke_workload_delegation(delegation_name) + except Exception: + logger.exception( + "Failed to revoke Docker workload delegation after container scheduling failure", + extra={"delegation_name": delegation_name}, ) - assert token_endpoint is not None - assert client_id is not None - assert workload_config.username is not None - assert password is not None - return OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint=token_endpoint, - client_id=client_id, - client_secret=workload_config.client_secret, - username=workload_config.username, - password=password, - scope=scope, - default_expires_in_seconds=workload_config.subject_token_ttl_seconds, + def _workload_delegation_name_for_step(self, step: PlatformJobStepWithContext) -> str: + return docker_delegation_name( + workload_workspace=step.workspace, + job_id=step.job, + attempt_id=step.attempt_id, + step_id=step.id, + ) + + @staticmethod + def _workload_delegation_name_from_container(container: Container) -> str | None: + labels = getattr(container, "labels", None) or {} + if DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL not in labels: + return None + + workload_workspace = labels.get(JOB_WORKSPACE_ID_LABEL) + job_id = labels.get(JOB_ID_LABEL) + attempt_id = labels.get(JOB_ATTEMPT_ID_LABEL) + step_id = labels.get(JOB_STEP_ID_LABEL) + if not isinstance(workload_workspace, str) or not workload_workspace: + return None + if not isinstance(job_id, str) or not job_id: + return None + if not isinstance(attempt_id, str) or not attempt_id: + return None + if not isinstance(step_id, str) or not step_id: + return None + + return docker_delegation_name( + workload_workspace=workload_workspace, + job_id=job_id, + attempt_id=attempt_id, + step_id=step_id, ) + def _revoke_workload_delegation_after_terminal(self, delegation_name: str | None) -> None: + if not delegation_name: + return + try: + self._revoke_workload_delegation(delegation_name) + except Exception: + logger.exception( + "Failed to revoke Docker workload delegation after terminal status", + extra={"delegation_name": delegation_name}, + ) + def _write_workload_identity_subject_token(self, volume_name: str, token: str) -> None: storage_config = self._execution_profile_config.storage permissions_image = ( @@ -416,109 +514,6 @@ def _write_workload_identity_subject_token(self, volume_name: str, token: str) - except Exception: logger.debug("Failed to remove workload identity token writer container", exc_info=True) - def _build_workload_identity_refresher(self, volume_name: str) -> SubjectTokenRefreshLoop: - issuer = self._create_docker_subject_token_issuer() - return SubjectTokenRefreshLoop( - issuer=issuer, - write_token=lambda token: self._write_workload_identity_subject_token(volume_name, token), - refresh_margin_seconds=self._execution_profile_config.workload_identity.refresh_margin_seconds, - ) - - def _start_workload_identity_refresher( - self, container_name: str, refresher: SubjectTokenRefreshLoop | None - ) -> None: - if refresher is None: - return - old_refresher = self._workload_identity_refreshers.get(container_name) - if old_refresher is not None: - old_refresher.stop() - self._workload_identity_refreshers.pop(container_name, None) - refresher.start() - self._workload_identity_refreshers[container_name] = refresher - - def _stop_workload_identity_refresher(self, container_name: str) -> None: - refresher = self._workload_identity_refreshers.get(container_name) - if refresher is not None: - refresher.stop() - self._workload_identity_refreshers.pop(container_name, None) - - def _stop_all_workload_identity_refreshers(self) -> None: - for container_name, refresher in list(self._workload_identity_refreshers.items()): - try: - refresher.stop() - except Exception: - logger.warning( - "Failed to stop workload identity subject token refresher for Docker container %s", - container_name, - exc_info=True, - ) - finally: - self._workload_identity_refreshers.pop(container_name, None) - - @staticmethod - def _get_workload_identity_volume_from_mounts(container: Container) -> str | None: - attrs = getattr(container, "attrs", None) or {} - mounts = attrs.get("Mounts", []) or [] - for mount in mounts: - if mount.get("Type") != "volume" or mount.get("Destination") != WORKLOAD_IDENTITY_VOLUME_PATH: - continue - volume_name = mount.get("Name") or mount.get("Source") - if isinstance(volume_name, str) and volume_name: - return volume_name - return None - - def _get_workload_identity_volume_for_container(self, container: Container) -> str | None: - labels = getattr(container, "labels", None) or {} - token_file = labels.get(DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL) - if token_file is not None and token_file != WORKLOAD_IDENTITY_TOKEN_FILE_PATH: - logger.warning( - "Skipping workload identity refresher restore for container with unexpected token file label", - extra={ - "container_name": getattr(container, "name", None), - "token_file": token_file, - }, - ) - return None - - volume_name = labels.get(DOCKER_WORKLOAD_IDENTITY_VOLUME_LABEL) - if volume_name: - return volume_name - return self._get_workload_identity_volume_from_mounts(container) - - def _restore_workload_identity_refresher_for_container(self, container: Container) -> None: - container_name = getattr(container, "name", None) - if not isinstance(container_name, str) or not container_name: - return - if container_name in self._workload_identity_refreshers: - return - if container.status != "running": - return - if not self._is_container_owned_by_this_controller(container): - return - - labels = getattr(container, "labels", None) or {} - if labels.get(JOB_TYPE_LABEL) != JOB_TYPE_JOB: - return - - volume_name = self._get_workload_identity_volume_for_container(container) - if volume_name is None: - return - - try: - self._start_workload_identity_refresher( - container_name, - self._build_workload_identity_refresher(volume_name), - ) - logger.info( - "Restored Docker workload identity refresher for running job container", - extra={"container_name": container_name, "volume_name": volume_name}, - ) - except Exception: - logger.exception( - "Failed to restore Docker workload identity refresher for running job container", - extra={"container_name": container_name, "volume_name": volume_name}, - ) - @staticmethod def get_label_from_container(container: Container, label: str) -> str: return container.labels[label] @@ -966,12 +961,12 @@ def schedule_single_container( NEMO_JOB_SECRETS_ENVVAR: self.get_secrets_environment_variable_for_injection(step), } ) - workload_identity_enabled = self._is_workload_identity_enabled() + workload_identity_enabled = self._should_enable_workload_identity_for_step(step) if workload_identity_enabled: env[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] = WORKLOAD_IDENTITY_TOKEN_FILE_PATH # Set auth context env var for job containers to make authenticated API calls - if step.auth_context: + if step.auth_context and not workload_identity_enabled: sdk_auth_context = step.auth_context auth_context = AuthContext.model_validate(sdk_auth_context.model_dump(mode="python", exclude_none=True)) principal = auth_context.to_principal() @@ -1021,6 +1016,7 @@ def schedule_single_container( # The admission slot is owned by this method until submit succeeds. # After that, run_container releases it when the start worker exits. + container_args = None try: container_args = self._prepare_container_args_for_start( executor_config=executor_config, @@ -1035,6 +1031,9 @@ def schedule_single_container( submitted_to_threadpool_at = time.monotonic() self._container_run_threadpool.submit(self.run_container, step, container_args, submitted_to_threadpool_at) except Exception: + if container_args is not None: + self._revoke_workload_delegation_after_failed_start(container_args.get("_nmp_workload_delegation_name")) + self.cleanup_task_storage_volumes(step.workspace, step.job, task_id) self._container_start_admission.release() raise logger.debug( @@ -1066,7 +1065,7 @@ def _prepare_container_args_for_start( job_volume_name = storage_config.volume_name if storage_config is not None else "" task_volume_name = self.task_storage_volume_name(workspace=step.workspace, job=step.job, task=task_id) config_volume_name = self.task_config_volume_name(workspace=step.workspace, job=step.job, task=task_id) - workload_identity_enabled = self._is_workload_identity_enabled() + workload_identity_enabled = self._should_enable_workload_identity_for_step(step) workload_identity_volume_name = ( self.task_workload_identity_volume_name(workspace=step.workspace, job=step.job, task=task_id) if workload_identity_enabled @@ -1087,14 +1086,7 @@ def _prepare_container_args_for_start( step_config_json=step_config_json, workload_identity_volume_name=workload_identity_volume_name, ) - workload_identity_refresher = None - if workload_identity_volume_name is not None: - try: - workload_identity_refresher = self._build_workload_identity_refresher(workload_identity_volume_name) - workload_identity_refresher.refresh_once() - except Exception: - self.cleanup_task_storage_volumes(step.workspace, step.job, task_id) - raise + workload_delegation_name = None logger.debug( "Docker job storage ensured", extra={ @@ -1153,11 +1145,24 @@ def _prepare_container_args_for_start( additional_volume_mounts=additional_volume_mounts, ), } - if workload_identity_refresher is not None: - container_args["_nmp_workload_identity_refresher"] = workload_identity_refresher + if workload_identity_volume_name is not None: + try: + workload_delegation_name = self._prepare_workload_identity_for_step( + step=step, + workload_identity_volume_name=workload_identity_volume_name, + ) + container_args["_nmp_workload_delegation_name"] = workload_delegation_name + except Exception: + self.cleanup_task_storage_volumes(step.workspace, step.job, task_id) + raise container_args["network"] = self._execution_profile_config.networking.job_container_network - return self.configure_container(container_args, executor_config) + try: + return self.configure_container(container_args, executor_config) + except Exception: + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) + self.cleanup_task_storage_volumes(step.workspace, step.job, task_id) + raise def cancel_scheduling(self, step: PlatformJobStepWithContext) -> bool: """Check if the job step is cancelling or pausing, and update status accordingly.""" @@ -1226,10 +1231,12 @@ def run_container( if submitted_to_threadpool_at is not None: log_extra["queue_delay_seconds"] = time.monotonic() - submitted_to_threadpool_at logger.debug("Docker run_container worker started", extra=log_extra) + workload_delegation_name = container_args.get("_nmp_workload_delegation_name") try: self._run_container_in_thread(step, container_args) except FailedToScheduleError as e: logger.exception("Failed to schedule container for job step") + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) status = PlatformJobStatus.ERROR try: self._jobs.update_job_step_status( @@ -1242,13 +1249,14 @@ def run_container( logger.exception("Failed to persist scheduling error for job step") except Exception: logger.exception("Unexpected error while scheduling container for job step") + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) finally: self._container_start_admission.release() def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_args: dict): status_details = {} status = PlatformJobStatus.PENDING - workload_identity_refresher = container_args.pop("_nmp_workload_identity_refresher", None) + workload_delegation_name = container_args.pop("_nmp_workload_delegation_name", None) # If a request to pause or cancel came in while we were waiting for scheduling loop, # cancel scheduling the container @@ -1263,6 +1271,7 @@ def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_a "duration_seconds": time.monotonic() - cancel_check_started_at, }, ) + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) return logger.debug( "Docker pre-create cancellation check completed", @@ -1364,6 +1373,7 @@ def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_a # cancel scheduling the container logger.debug("Checking for cancellation or pausing before creating container after image pull") if self.cancel_scheduling(step): + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) return # Send the status update to indicate we are starting the container @@ -1441,6 +1451,7 @@ def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_a "duration_seconds": time.monotonic() - pre_start_cancel_check_started_at, }, ) + self._revoke_workload_delegation_after_failed_start(workload_delegation_name) return logger.debug( "Docker pre-start cancellation check completed", @@ -1475,7 +1486,6 @@ def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_a max_attempts = 3 attempts = 0 start_started_at = time.monotonic() - self._start_workload_identity_refresher(container.name, workload_identity_refresher) while not started and attempts < max_attempts: attempts += 1 try: @@ -1503,7 +1513,6 @@ def _run_container_in_thread(self, step: PlatformJobStepWithContext, container_a }, ) except Exception as e: - self._stop_workload_identity_refresher(container.name) raise FailedToScheduleError( f"Failed to start container {container.name} for job step", error_details={"message": f"Failed to start container: {e}"}, @@ -1759,14 +1768,13 @@ def docker_state_debug_fields(self, container: Container) -> dict[str, Any]: } def create_step_update(self, step: PlatformJobStepWithContext, container: Container) -> JobUpdate: - if step.status in (PlatformJobStatus.ACTIVE, PlatformJobStatus.PENDING): - self._restore_workload_identity_refresher_for_container(container) - status, status_details, error_stack = self.map_docker_container_status_to_platform_status(step, container) task_id = self.get_label_from_container(container, JOB_TASK_ID_LABEL) error_details = {} if status == PlatformJobStatus.ERROR: error_details["message"] = status_details.get("message", "Job encountered an error") + if status in PlatformJobStatus.terminals() and self._should_enable_workload_identity_for_step(step): + self._revoke_workload_delegation_after_terminal(self._workload_delegation_name_for_step(step)) logger.debug( "Docker container status mapped to platform status", @@ -2016,8 +2024,9 @@ def cleanup_single_container(self, container: Container) -> None: job = self.get_label_from_container(container, JOB_ID_LABEL) task = self.get_label_from_container(container, JOB_TASK_ID_LABEL) exit_code = container.attrs.get("State", {}).get("ExitCode", 0) + delegation_name = self._workload_delegation_name_from_container(container) - self._stop_workload_identity_refresher(container.name) + self._revoke_workload_delegation_after_terminal(delegation_name) self.cleanup_container(container) logger.debug( "Cleaned up container", diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/common.py b/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/common.py index 5c5683745b..575efd9330 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/common.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/common.py @@ -978,7 +978,7 @@ def create_pod_template_spec( """ platform_config = get_platform_config() - workload_identity_enabled = is_workload_identity_token_exchange_enabled() + workload_identity_enabled = is_workload_identity_token_exchange_enabled() and step.auth_context is not None workload_identity_token_audience = config.workload_identity_token_audience or get_workload_identity_token_audience() # Profile-level env vars first (e.g. HOME=/tmp); system, step, and shared env override these @@ -1015,7 +1015,7 @@ def create_pod_template_spec( env.append(client.V1EnvVar(name=WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR, value=WORKLOAD_IDENTITY_TOKEN_FILE_PATH)) # Set auth context env var for job containers to make authenticated API calls - if step.auth_context: + if step.auth_context and not workload_identity_enabled: sdk_auth_context = step.auth_context auth_context = AuthContext.model_validate(sdk_auth_context.model_dump(mode="python", exclude_none=True)) principal = auth_context.to_principal() diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/kubernetes_job.py b/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/kubernetes_job.py index 684755d4ec..a68d3646b1 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/kubernetes_job.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/backends/kubernetes/kubernetes_job.py @@ -1,20 +1,34 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import asyncio import logging from typing import Any, Generic, Literal, TypeVar from kubernetes import client from kubernetes.client.models import V1Job, V1JobStatus from kubernetes.client.rest import ApiException +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.client.adapter import client_from_platform +from nemo_platform_plugin.entities.client import AsyncEntitiesClient from nemo_platform_plugin.jobs.types import PlatformJobStepWithContext, PlatformJobTaskUpdate +from nmp.common.auth import ( + AuthContext, + WorkloadDelegationConflictError, + WorkloadDelegationEntity, + WorkloadDelegationStore, + reference_delegation_name, +) +from nmp.common.entities import SYSTEM_WORKSPACE, EntityClient from nmp.common.jobs.schemas import PlatformJobStatus from nmp.core.jobs.app.constants import ( + JOB_ATTEMPT_ID_LABEL, JOB_EXECUTION_BACKEND_LABEL, JOB_EXECUTION_PROFILE_LABEL, JOB_ID_LABEL, JOB_MANAGED_BY_JOBS_CONTROLLER, JOB_MANAGED_BY_LABEL, + JOB_STEP_ID_LABEL, JOB_STEP_NAME_LABEL, JOB_TYPE_JOB, JOB_TYPE_LABEL, @@ -30,7 +44,13 @@ GPUExecutionProvider, ) from nmp.core.jobs.app.schemas import BaseExecutionProfile -from nmp.core.jobs.controllers.backends.base import JobBackend, JobUpdate, staleness_error_message +from nmp.core.jobs.controllers.backends.base import ( + JobBackend, + JobUpdate, + is_workload_identity_token_exchange_enabled, + staleness_error_message, +) +from nmp.core.jobs.controllers.backends.exceptions import JobStorageError from nmp.core.jobs.controllers.backends.kubernetes.common import ( BaseKubernetesExecutionProfileConfig, aggregate_pod_statuses_for_job_step, @@ -43,13 +63,19 @@ delete_configmap, get_namespace_from_environment, list_pod_status, + list_pods_by_labels, load_kubernetes_config, name_for_step, update_all_tasks, ) +from nmp.core.jobs.controllers.backends.workload_tokens import ( + create_authenticated_async_nmp_sdk, + workload_delegation_expires_at, +) from pydantic import Field logger = logging.getLogger(__name__) +KUBERNETES_POD_UID_REFERENCE_NAME = "authentication.kubernetes.io/pod-uid" ProviderT = TypeVar("ProviderT", bound=ExecutionProviderT) @@ -87,12 +113,237 @@ def init(self) -> None: self._batch_v1 = client.BatchV1Api() self._core_v1 = client.CoreV1Api() self.namespace = self._execution_profile_config.namespace or get_namespace_from_environment() + self._workload_delegations_by_job: dict[str, set[str]] = {} def shutdown(self): self._batch_v1.api_client.close() self._core_v1.api_client.close() return + def _should_manage_workload_delegation_for_step(self, step: PlatformJobStepWithContext) -> bool: + if not is_workload_identity_token_exchange_enabled(): + return False + return step.auth_context is not None + + @staticmethod + def _workload_delegation_audience() -> str: + try: + from nmp.common.config import get_auth_config + + oidc_config = get_auth_config().oidc + except Exception: + logger.debug("Could not resolve auth config for Kubernetes workload delegation audience", exc_info=True) + return "nemo-platform" + return ( + getattr(oidc_config, "workload_audience", None) or getattr(oidc_config, "audience", None) or "nemo-platform" + ) + + def _create_async_nmp_sdk(self) -> AsyncNeMoPlatform: + return create_authenticated_async_nmp_sdk( + self._nmp_sdk, + missing_headers_message="Kubernetes workload delegation requires authenticated SDK headers", + ) + + async def _register_workload_delegation_async(self, entity: WorkloadDelegationEntity) -> None: + async_sdk = self._create_async_nmp_sdk() + try: + entity_client = EntityClient(client_from_platform(async_sdk, AsyncEntitiesClient)).as_service( + "jobs", internal=True + ) + store = WorkloadDelegationStore(entity_client) + await store.register(entity) + finally: + await async_sdk.close() + + async def _revoke_workload_delegation_async(self, delegation_name: str) -> None: + async_sdk = self._create_async_nmp_sdk() + try: + entity_client = EntityClient(client_from_platform(async_sdk, AsyncEntitiesClient)).as_service( + "jobs", internal=True + ) + store = WorkloadDelegationStore(entity_client) + await store.revoke(delegation_name) + finally: + await async_sdk.close() + + def _register_workload_delegation(self, entity: WorkloadDelegationEntity) -> None: + try: + asyncio.get_running_loop() + except RuntimeError: + asyncio.run(self._register_workload_delegation_async(entity)) + return + raise JobStorageError( + "Kubernetes workload delegation registration cannot run from an active asyncio event loop" + ) + + def _revoke_workload_delegation(self, delegation_name: str) -> None: + try: + asyncio.get_running_loop() + except RuntimeError: + asyncio.run(self._revoke_workload_delegation_async(delegation_name)) + return + raise JobStorageError("Kubernetes workload delegation revocation cannot run from an active asyncio event loop") + + def _workload_subject_for_job(self, job: V1Job) -> str: + namespace = getattr(job.metadata, "namespace", None) or self.namespace + service_account_name = self._execution_profile_config.service_account_name + if job.spec and job.spec.template and job.spec.template.spec: + service_account_name = job.spec.template.spec.service_account_name or service_account_name + return f"system:serviceaccount:{namespace}:{service_account_name}" + + def _workload_delegation_job_key(self, job: V1Job) -> str | None: + job_name = getattr(job.metadata, "name", None) + if not isinstance(job_name, str) or not job_name: + logger.warning("Skipping Kubernetes workload delegation reconciliation for job with missing name") + return None + namespace = getattr(job.metadata, "namespace", None) or self.namespace + return f"{namespace}/{job_name}" + + def _delegation_name_for_pod(self, job: V1Job, pod: client.V1Pod) -> str | None: + pod_uid = getattr(pod.metadata, "uid", None) + if not pod_uid: + return None + return reference_delegation_name( + workload_audience=self._workload_delegation_audience(), + workload_subject=self._workload_subject_for_job(job), + bound_reference_name=KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value=pod_uid, + ) + + def _workload_delegation_pods_for_job(self, job: V1Job) -> list[client.V1Pod]: + labels = getattr(job.metadata, "labels", None) or {} + selector_labels = { + key: labels[key] + for key in ( + JOB_TYPE_LABEL, + JOB_WORKSPACE_ID_LABEL, + JOB_ID_LABEL, + JOB_ATTEMPT_ID_LABEL, + JOB_STEP_NAME_LABEL, + JOB_STEP_ID_LABEL, + JOB_MANAGED_BY_LABEL, + ) + if key in labels + } + required_labels = { + JOB_TYPE_LABEL, + JOB_WORKSPACE_ID_LABEL, + JOB_ID_LABEL, + JOB_ATTEMPT_ID_LABEL, + JOB_STEP_NAME_LABEL, + JOB_STEP_ID_LABEL, + JOB_MANAGED_BY_LABEL, + } + if set(selector_labels) != required_labels: + logger.warning( + "Skipping Kubernetes workload delegation reconciliation for job with missing labels", + extra={"job_name": getattr(job.metadata, "name", None), "labels": labels}, + ) + return [] + return list_pods_by_labels( + self._core_v1, getattr(job.metadata, "namespace", None) or self.namespace, selector_labels + ) + + def _build_workload_delegation( + self, + *, + step: PlatformJobStepWithContext, + job: V1Job, + pod: client.V1Pod, + delegation_name: str, + ) -> WorkloadDelegationEntity: + if step.auth_context is None: + raise JobStorageError( + "Kubernetes workload identity requires a job auth_context for on-behalf-of delegation" + ) + + pod_uid = getattr(pod.metadata, "uid", None) + if not pod_uid: + raise JobStorageError("Kubernetes workload identity requires a Pod UID for on-behalf-of delegation") + + auth_context = AuthContext.model_validate(step.auth_context.model_dump(mode="python", exclude_none=True)) + expires_at = workload_delegation_expires_at( + ttl_seconds_active=self._execution_profile_config.ttl_seconds_active + ) + return WorkloadDelegationEntity( + name=delegation_name, + workspace=SYSTEM_WORKSPACE, + workload_subject=self._workload_subject_for_job(job), + workload_audience=self._workload_delegation_audience(), + workload_workspace=step.workspace, + job_id=step.job, + attempt_id=step.attempt_id, + step_id=step.id, + auth_context=auth_context, + bound_reference_name=KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value=pod_uid, + expires_at=expires_at, + ) + + def _ensure_workload_delegations_for_job(self, step: PlatformJobStepWithContext, job: V1Job) -> None: + if not self._should_manage_workload_delegation_for_step(step): + return + + job_key = self._workload_delegation_job_key(job) + if job_key is None: + return + registered_delegations = self._workload_delegations_by_job.setdefault(job_key, set()) + + for pod in self._workload_delegation_pods_for_job(job): + delegation_name = self._delegation_name_for_pod(job, pod) + if delegation_name is None or delegation_name in registered_delegations: + continue + delegation = self._build_workload_delegation( + step=step, + job=job, + pod=pod, + delegation_name=delegation_name, + ) + try: + self._register_workload_delegation(delegation) + except WorkloadDelegationConflictError: + logger.debug( + "Kubernetes workload delegation already exists", + extra={"delegation_name": delegation_name}, + ) + except Exception: + logger.exception( + "Failed to register Kubernetes workload delegation; will retry on next sync", + extra={"delegation_name": delegation_name, "job_key": job_key}, + ) + continue + registered_delegations.add(delegation_name) + + def _revoke_workload_delegations_for_job(self, job: V1Job) -> None: + if not is_workload_identity_token_exchange_enabled(): + return + + job_key = self._workload_delegation_job_key(job) + recorded_delegations = self._workload_delegations_by_job.get(job_key, set()) if job_key else set() + delegation_names = set(recorded_delegations) + if not delegation_names: + for pod in self._workload_delegation_pods_for_job(job): + delegation_name = self._delegation_name_for_pod(job, pod) + if delegation_name is not None: + delegation_names.add(delegation_name) + + failed_delegations: set[str] = set() + for delegation_name in sorted(delegation_names): + try: + self._revoke_workload_delegation(delegation_name) + except Exception: + logger.exception( + "Failed to revoke Kubernetes workload delegation; will retry on next sync", + extra={"delegation_name": delegation_name, "job_key": job_key}, + ) + failed_delegations.add(delegation_name) + + if job_key is not None: + if failed_delegations: + self._workload_delegations_by_job[job_key] = failed_delegations + else: + self._workload_delegations_by_job.pop(job_key, None) + def get_job_by_name(self, name: str) -> V1Job | None: try: return self._batch_v1.read_namespaced_job(name=name, namespace=self.namespace) # type: ignore @@ -336,6 +587,11 @@ def create_step_update(self, step: PlatformJobStepWithContext, job: V1Job) -> Jo error_details["message"] = "One or more tasks are in error state" else: error_details["message"] += "; One or more tasks are in error state" + if self._should_manage_workload_delegation_for_step(step): + if status in PlatformJobStatus.terminals(): + self._revoke_workload_delegations_for_job(job) + else: + self._ensure_workload_delegations_for_job(step, job) return JobUpdate(status=status, status_details=status_details, error_details=error_details) def enforce_sync_ttl( @@ -454,6 +710,7 @@ def terminate_job(self, job: V1Job): ) return logger.info("Deleting Kubernetes Job", extra={"job_name": job_name, "namespace": self.namespace}) + self._revoke_workload_delegations_for_job(job) try: self._batch_v1.delete_namespaced_job( name=job_name, diff --git a/services/core/jobs/src/nmp/core/jobs/controllers/backends/workload_tokens.py b/services/core/jobs/src/nmp/core/jobs/controllers/backends/workload_tokens.py index d726ae5fe8..fac68cff7d 100644 --- a/services/core/jobs/src/nmp/core/jobs/controllers/backends/workload_tokens.py +++ b/services/core/jobs/src/nmp/core/jobs/controllers/backends/workload_tokens.py @@ -5,144 +5,53 @@ from __future__ import annotations +import datetime import io -import json -import logging import tarfile -import threading -import time -from dataclasses import dataclass, field -from ipaddress import ip_address -from math import isfinite -from typing import Callable, Protocol -from urllib.parse import urlparse -import httpx - -logger = logging.getLogger(__name__) - -DEFAULT_SUBJECT_TOKEN_TTL_SECONDS = 600 -DEFAULT_SUBJECT_TOKEN_REFRESH_MARGIN_SECONDS = 60 -DEFAULT_SUBJECT_TOKEN_FAILURE_BACKOFF_MAX_SECONDS = 60.0 -DEFAULT_SUBJECT_TOKEN_STOP_TIMEOUT_SECONDS = 5.0 - - -def _is_loopback_host(host: str | None) -> bool: - if host is None: - return False - if host.lower() == "localhost": - return True - try: - return ip_address(host).is_loopback - except ValueError: - return False - - -def _validate_token_endpoint(token_endpoint: str) -> None: - parsed = urlparse(token_endpoint) - scheme = parsed.scheme.lower() - if scheme == "https" and parsed.netloc: - return - if scheme == "http" and _is_loopback_host(parsed.hostname): - return - raise RuntimeError( - "Invalid Docker workload identity token_endpoint configuration: " - "token_endpoint must use https://, except http:// is allowed for loopback hosts only" +from nemo_platform import AsyncNeMoPlatform, NeMoPlatform +from nmp.core.jobs.controllers.backends.exceptions import JobStorageError + +WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS = 300 + + +def create_authenticated_async_nmp_sdk( + nmp_sdk: NeMoPlatform, + *, + missing_headers_message: str, +) -> AsyncNeMoPlatform: + """Create an async SDK preserving authenticated headers from a sync SDK.""" + headers = getattr(nmp_sdk, "_custom_headers", None) + if not headers: + sdk_client = getattr(nmp_sdk, "_client", None) + sdk_client_headers = getattr(sdk_client, "headers", {}) + skip_headers = {"accept", "accept-encoding", "connection", "user-agent", "host"} + headers = {key: value for key, value in sdk_client_headers.items() if key.lower() not in skip_headers} + if not headers: + raise JobStorageError(missing_headers_message) + + async_sdk = AsyncNeMoPlatform( + base_url=str(getattr(nmp_sdk, "base_url")).rstrip("/"), + workspace=getattr(nmp_sdk, "workspace", None), + default_headers=headers or None, + timeout=getattr(nmp_sdk, "timeout", None), + max_retries=getattr(nmp_sdk, "max_retries", 2), ) + router = getattr(nmp_sdk, "_nmp_request_router", None) + if router is not None: + setattr(async_sdk, "_nmp_request_router", router) + return async_sdk -@dataclass(frozen=True) -class SubjectToken: - """Issued workload identity subject token and its refresh deadline metadata.""" - - value: str - expires_at: float - - def seconds_until_refresh(self, margin_seconds: int) -> float: - return max(0.0, self.expires_at - time.time() - margin_seconds) - - -class SubjectTokenIssuer(Protocol): - """Issues a subject token suitable for SDK RFC 8693 token exchange.""" - - def issue(self) -> SubjectToken: ... - - -@dataclass(frozen=True) -class OAuthPasswordGrantSubjectTokenIssuer: - """Demo issuer that obtains a short-lived subject token with OAuth password grant.""" - - token_endpoint: str - client_id: str - username: str - password: str = field(repr=False) - client_secret: str | None = field(default=None, repr=False) - scope: str | None = None - default_expires_in_seconds: int = DEFAULT_SUBJECT_TOKEN_TTL_SECONDS - timeout: float = 30.0 - - def issue(self) -> SubjectToken: - _validate_token_endpoint(self.token_endpoint) - data = { - "grant_type": "password", - "client_id": self.client_id, - "username": self.username, - "password": self.password, - } - if self.client_secret: - data["client_secret"] = self.client_secret - if self.scope: - data["scope"] = self.scope - - response = httpx.post(self.token_endpoint, data=data, timeout=self.timeout) - if response.status_code != 200: - error_data: dict[str, object] = {} - if response.headers.get("content-type", "").startswith("application/json"): - try: - payload = response.json() - except (json.JSONDecodeError, ValueError): - payload = {} - if isinstance(payload, dict): - error_data = payload - error = error_data.get("error", "unknown_error") - description = error_data.get("error_description", response.text) - raise RuntimeError(f"Failed to issue workload subject token: {error} - {description}") - - try: - token_data = response.json() - except (json.JSONDecodeError, ValueError) as exc: - raise RuntimeError( - "Failed to issue workload subject token: invalid_response - " - "Token endpoint response was not a JSON object" - ) from exc - if not isinstance(token_data, dict): - raise RuntimeError( - "Failed to issue workload subject token: invalid_response - " - "Token endpoint response was not a JSON object" - ) - - access_token = token_data.get("access_token") - if not isinstance(access_token, str) or not access_token.strip(): - raise RuntimeError( - "Failed to issue workload subject token: invalid_response - " - "Token endpoint response did not include a non-empty access_token" - ) - - if "expires_in" in token_data: - expires_in = token_data["expires_in"] - if ( - isinstance(expires_in, bool) - or not isinstance(expires_in, int | float) - or not isfinite(expires_in) - or expires_in <= 0 - ): - raise RuntimeError( - "Failed to issue workload subject token: invalid_response - " - "Token endpoint response did not include a positive numeric expires_in" - ) - else: - expires_in = self.default_expires_in_seconds - return SubjectToken(value=access_token, expires_at=time.time() + expires_in) +def workload_delegation_expires_at( + *, + ttl_seconds_active: int, + now: datetime.datetime | None = None, +) -> datetime.datetime: + effective_now = now or datetime.datetime.now(datetime.timezone.utc) + if effective_now.tzinfo is None: + effective_now = effective_now.replace(tzinfo=datetime.timezone.utc) + return effective_now + datetime.timedelta(seconds=ttl_seconds_active + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS) def build_token_archive(token: str, *, name: str = "token.tmp") -> io.BytesIO: @@ -156,74 +65,3 @@ def build_token_archive(token: str, *, name: str = "token.tmp") -> io.BytesIO: tar.addfile(info, io.BytesIO(data)) archive.seek(0) return archive - - -class SubjectTokenRefreshLoop: - """Background refresher for a controller-owned workload subject token file.""" - - def __init__( - self, - *, - issuer: SubjectTokenIssuer, - write_token: Callable[[str], None], - refresh_margin_seconds: int = DEFAULT_SUBJECT_TOKEN_REFRESH_MARGIN_SECONDS, - min_sleep_seconds: float = 1.0, - max_failure_backoff_seconds: float = DEFAULT_SUBJECT_TOKEN_FAILURE_BACKOFF_MAX_SECONDS, - ) -> None: - self._issuer = issuer - self._write_token = write_token - self._refresh_margin_seconds = refresh_margin_seconds - self._min_sleep_seconds = min_sleep_seconds - self._max_failure_backoff_seconds = max(min_sleep_seconds, max_failure_backoff_seconds) - self._stop_timeout_seconds = DEFAULT_SUBJECT_TOKEN_STOP_TIMEOUT_SECONDS - self._stop = threading.Event() - self._thread: threading.Thread | None = None - - def start(self) -> None: - if self._thread is not None and self._thread.is_alive(): - return - self._stop.clear() - self._thread = threading.Thread(target=self._run, name="nmp-workload-token-refresh", daemon=True) - self._thread.start() - - def stop(self) -> None: - self._stop.set() - thread = self._thread - if thread is None: - return - - thread.join(timeout=self._stop_timeout_seconds) - if thread.is_alive(): - raise RuntimeError("Timed out stopping workload identity subject token refresher") - self._thread = None - - def refresh_once(self) -> SubjectToken: - token = self._issuer.issue() - self._write_token(token.value) - return token - - def _refresh_once_for_worker(self) -> SubjectToken | None: - token = self._issuer.issue() - if self._stop.is_set(): - return None - self._write_token(token.value) - return token - - def _run(self) -> None: - token: SubjectToken | None = None - failure_sleep_seconds = self._min_sleep_seconds - while not self._stop.is_set(): - try: - token = self._refresh_once_for_worker() - except Exception: - logger.exception("Failed to refresh workload identity subject token") - if self._stop.wait(failure_sleep_seconds): - return - failure_sleep_seconds = min(self._max_failure_backoff_seconds, failure_sleep_seconds * 2) - continue - if token is None: - return - - failure_sleep_seconds = self._min_sleep_seconds - sleep_seconds = max(self._min_sleep_seconds, token.seconds_until_refresh(self._refresh_margin_seconds)) - self._stop.wait(sleep_seconds) diff --git a/services/core/jobs/tests/controllers/test_docker_backend.py b/services/core/jobs/tests/controllers/test_docker_backend.py index b20a5d4087..d0908b836d 100644 --- a/services/core/jobs/tests/controllers/test_docker_backend.py +++ b/services/core/jobs/tests/controllers/test_docker_backend.py @@ -3,7 +3,6 @@ import datetime import json -import time import uuid from types import SimpleNamespace from typing import Iterator @@ -12,7 +11,14 @@ import pytest from docker.errors import APIError, NotFound from nemo_platform.types.shared import AuthContext as SdkAuthContext -from nmp.common.auth import NMP_PRINCIPAL_ENVVAR, AuthContext, Principal +from nmp.common.auth import ( + NMP_PRINCIPAL_ENVVAR, + AuthContext, + Principal, + docker_delegation_name, + parse_opaque_docker_proof_token, + verify_opaque_docker_proof_token_hash, +) from nmp.common.config import PlatformConfig from nmp.common.docker.gpu_pool import DockerGPUPool from nmp.common.jobs.constants import ( @@ -25,6 +31,7 @@ from nmp.common.jobs.schemas import PlatformJobStatus from nmp.core.jobs.api.v2.jobs.schemas import PlatformJobStepWithContext from nmp.core.jobs.app.constants import ( + JOB_ATTEMPT_ID_LABEL, JOB_CONTROLLER_INSTANCE_ID_LABEL, JOB_EXECUTION_BACKEND_LABEL, JOB_EXECUTION_PROFILE_LABEL, @@ -76,7 +83,7 @@ ResourceAllocationError, SchedulingDeferred, ) -from nmp.core.jobs.controllers.backends.workload_tokens import SubjectToken +from nmp.core.jobs.controllers.backends.workload_tokens import WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS from pydantic import ValidationError from services.core.jobs.tests.controllers.client_mocks import data_response @@ -109,6 +116,16 @@ def assert_created_task_volumes_cleaned_up(docker_client_mock) -> None: assert docker_client_mock.volumes.get.return_value.remove.call_count == 3 +def workload_token_exchange_auth_config(enabled: bool = True) -> SimpleNamespace: + return SimpleNamespace( + oidc=SimpleNamespace( + workload_token_exchange_enabled=enabled, + workload_audience="nemo-platform", + audience=None, + ) + ) + + @pytest.fixture def docker_client_mock(monkeypatch): """Mock docker client for testing.""" @@ -793,93 +810,230 @@ def test_docker_job_rejects_reserved_step_auth_env_vars(docker_job, test_job_ste def test_docker_job_injects_workload_identity_volume_when_token_exchange_enabled( - docker_job, docker_client_mock, test_job_step + docker_job, docker_client_mock, test_job_step_with_auth_context ): - auth_config = SimpleNamespace(oidc=SimpleNamespace(workload_token_exchange_enabled=True)) - - class FakeIssuer: - def issue(self): - return SubjectToken(value="subject-token", expires_at=time.time() + 3600) + auth_config = workload_token_exchange_auth_config() + expected_delegation_name = docker_delegation_name( + workload_workspace=test_job_step_with_auth_context.workspace, + job_id=test_job_step_with_auth_context.job, + attempt_id=test_job_step_with_auth_context.attempt_id, + step_id=test_job_step_with_auth_context.id, + ) with ( patch("nmp.common.config.get_auth_config", return_value=auth_config), - patch.object(docker_job, "_create_docker_subject_token_issuer", return_value=FakeIssuer()), + patch.object(docker_job, "_register_workload_delegation") as register_delegation, + patch.object( + docker_job, + "_write_workload_identity_subject_token", + wraps=docker_job._write_workload_identity_subject_token, + ) as write_workload_token, ): - docker_job.schedule(test_job_step.step_spec.executor, test_job_step) + docker_job.schedule(test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context) docker_job._container_run_threadpool.shutdown(wait=True) docker_job._container_run_threadpool = MagicMock() - try: - job_create_call = next( - call - for call in docker_client_mock.containers.create.call_args_list - if call.kwargs.get("name") == "job-test-job-id-test-step" + register_delegation.assert_called_once() + delegation = register_delegation.call_args.args[0] + assert delegation.name == expected_delegation_name + assert delegation.workload_subject == expected_delegation_name + assert delegation.workload_audience == "nemo-platform" + assert delegation.workload_workspace == test_job_step_with_auth_context.workspace + assert delegation.job_id == test_job_step_with_auth_context.job + assert delegation.attempt_id == test_job_step_with_auth_context.attempt_id + assert delegation.step_id == test_job_step_with_auth_context.id + assert delegation.auth_context.principal_id == "creator@example.com" + assert delegation.opaque_subject_token_hash is not None + write_workload_token.assert_called_once() + written_subject_token = write_workload_token.call_args.args[1] + parsed_subject_token = parse_opaque_docker_proof_token(written_subject_token) + assert parsed_subject_token.delegation_name == expected_delegation_name + assert verify_opaque_docker_proof_token_hash( + parsed_subject_token.secret, + delegation.opaque_subject_token_hash, + ) + + job_create_call = next( + call + for call in docker_client_mock.containers.create.call_args_list + if call.kwargs.get("name") == "job-test-job-id-test-step" + ) + kwargs = job_create_call.kwargs + env = kwargs["environment"] + assert env[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] == WORKLOAD_IDENTITY_TOKEN_FILE_PATH + assert NMP_PRINCIPAL_ENVVAR not in env + + mounts = kwargs["mounts"] + workload_identity_mount = next(m for m in mounts if m["Target"] == WORKLOAD_IDENTITY_VOLUME_PATH) + assert workload_identity_mount["Type"] == "volume" + assert workload_identity_mount["Source"].startswith( + f"task-workload-identity-{test_job_step_with_auth_context.workspace}-{test_job_step_with_auth_context.job}-" + ) + assert workload_identity_mount["ReadOnly"] is True + assert kwargs["labels"][DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL] == WORKLOAD_IDENTITY_TOKEN_FILE_PATH + assert kwargs["labels"][DOCKER_WORKLOAD_IDENTITY_VOLUME_LABEL] == workload_identity_mount["Source"] + workload_token_write_call = next( + call + for call in docker_client_mock.containers.create.call_args_list + if call.kwargs.get("name", "").startswith("workload-token-write-") + ) + assert workload_token_write_call.kwargs["command"] == [ + "sh", + "-c", + "mv /workload-identity-vol/token.tmp /workload-identity-vol/token && chmod 0444 /workload-identity-vol/token", + ] + + +def test_docker_workload_delegation_expiry_covers_active_ttl(docker_job, test_job_step_with_auth_context): + docker_job._execution_profile_config.ttl_seconds_active = 900 + before = datetime.datetime.now(datetime.timezone.utc) + + with ( + patch.object(docker_job, "_register_workload_delegation") as register_delegation, + patch.object(docker_job, "_write_workload_identity_subject_token"), + ): + docker_job._prepare_workload_identity_for_step( + step=test_job_step_with_auth_context, + workload_identity_volume_name="workload-token-volume", ) - kwargs = job_create_call.kwargs - env = kwargs["environment"] - assert env[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] == WORKLOAD_IDENTITY_TOKEN_FILE_PATH - - mounts = kwargs["mounts"] - workload_identity_mount = next(m for m in mounts if m["Target"] == WORKLOAD_IDENTITY_VOLUME_PATH) - assert workload_identity_mount["Type"] == "volume" - assert workload_identity_mount["Source"].startswith( - f"task-workload-identity-{test_job_step.workspace}-{test_job_step.job}-" + + after = datetime.datetime.now(datetime.timezone.utc) + delegation = register_delegation.call_args.args[0] + expected_min = before + datetime.timedelta(seconds=900 + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS) + expected_max = after + datetime.timedelta(seconds=900 + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS) + assert expected_min <= delegation.expires_at <= expected_max + assert delegation.revoked_at is None + + +def test_docker_schedule_cleans_task_volumes_when_workload_identity_proof_token_fails( + docker_job, docker_client_mock, test_job_step_with_auth_context +): + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object( + docker_job, + "_provision_docker_workload_proof_token", + side_effect=JobStorageError("proof failed"), + ), + pytest.raises(JobStorageError, match="proof failed"), + ): + docker_job.schedule_single_container( + test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context ) - assert workload_identity_mount["ReadOnly"] is True - assert kwargs["labels"][DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL] == WORKLOAD_IDENTITY_TOKEN_FILE_PATH - assert kwargs["labels"][DOCKER_WORKLOAD_IDENTITY_VOLUME_LABEL] == workload_identity_mount["Source"] - workload_token_write_call = next( - call - for call in docker_client_mock.containers.create.call_args_list - if call.kwargs.get("name", "").startswith("workload-token-write-") + + assert_created_task_volumes_cleaned_up(docker_client_mock) + + +def test_docker_schedule_cleans_task_volumes_when_workload_identity_registration_fails( + docker_job, docker_client_mock, test_job_step_with_auth_context +): + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object(docker_job, "_register_workload_delegation", side_effect=JobStorageError("register failed")), + patch.object(docker_job, "_write_workload_identity_subject_token") as write_token, + patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation, + pytest.raises(JobStorageError, match="register failed"), + ): + docker_job.schedule_single_container( + test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context ) - assert workload_token_write_call.kwargs["command"] == [ - "sh", - "-c", - "mv /workload-identity-vol/token.tmp /workload-identity-vol/token && chmod 0444 /workload-identity-vol/token", - ] - finally: - for refresher in list(docker_job._workload_identity_refreshers.values()): - refresher.stop() + write_token.assert_not_called() + revoke_delegation.assert_not_called() + assert_created_task_volumes_cleaned_up(docker_client_mock) -def test_docker_schedule_cleans_task_volumes_when_workload_identity_issuer_fails( - docker_job, docker_client_mock, test_job_step + +def test_docker_schedule_cleans_task_volumes_and_revokes_delegation_when_workload_identity_write_fails( + docker_job, docker_client_mock, test_job_step_with_auth_context ): - docker_job._execution_profile_config.workload_identity.enabled = True + expected_delegation_name = docker_delegation_name( + workload_workspace=test_job_step_with_auth_context.workspace, + job_id=test_job_step_with_auth_context.job, + attempt_id=test_job_step_with_auth_context.attempt_id, + step_id=test_job_step_with_auth_context.id, + ) with ( + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object(docker_job, "_register_workload_delegation") as register_delegation, patch.object( docker_job, - "_create_docker_subject_token_issuer", - side_effect=JobStorageError("issuer failed"), + "_write_workload_identity_subject_token", + side_effect=RuntimeError("write failed"), ), - pytest.raises(JobStorageError, match="issuer failed"), + patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation, + pytest.raises(RuntimeError, match="write failed"), ): - docker_job.schedule_single_container(test_job_step.step_spec.executor, test_job_step) + docker_job.schedule_single_container( + test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context + ) + register_delegation.assert_called_once() + revoke_delegation.assert_called_once_with(expected_delegation_name) assert_created_task_volumes_cleaned_up(docker_client_mock) -def test_docker_schedule_cleans_task_volumes_when_initial_workload_identity_refresh_fails( - docker_job, docker_client_mock, test_job_step +def test_docker_schedule_cleans_task_volumes_and_revokes_delegation_when_configure_fails( + docker_job, docker_client_mock, test_job_step_with_auth_context ): - docker_job._execution_profile_config.workload_identity.enabled = True - refresher = MagicMock() - refresher.refresh_once.side_effect = RuntimeError("refresh failed") + expected_delegation_name = docker_delegation_name( + workload_workspace=test_job_step_with_auth_context.workspace, + job_id=test_job_step_with_auth_context.job, + attempt_id=test_job_step_with_auth_context.attempt_id, + step_id=test_job_step_with_auth_context.id, + ) with ( - patch.object(docker_job, "_build_workload_identity_refresher", return_value=refresher), - pytest.raises(RuntimeError, match="refresh failed"), + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object(docker_job, "_register_workload_delegation") as register_delegation, + patch.object(docker_job, "_write_workload_identity_subject_token") as write_token, + patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation, + patch.object(docker_job, "configure_container", side_effect=RuntimeError("configure failed")), + pytest.raises(RuntimeError, match="configure failed"), ): - docker_job.schedule_single_container(test_job_step.step_spec.executor, test_job_step) + docker_job.schedule_single_container( + test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context + ) - refresher.refresh_once.assert_called_once() + register_delegation.assert_called_once() + write_token.assert_called_once() + revoke_delegation.assert_called_once_with(expected_delegation_name) assert_created_task_volumes_cleaned_up(docker_client_mock) -def test_docker_sync_restores_workload_identity_refresher_from_container_labels( +def test_docker_schedule_cleans_task_volumes_and_revokes_delegation_when_submit_fails( + docker_job, docker_client_mock, test_job_step_with_auth_context +): + docker_job._container_run_threadpool = MagicMock() + docker_job._container_run_threadpool.submit.side_effect = RuntimeError("submit failed") + expected_delegation_name = docker_delegation_name( + workload_workspace=test_job_step_with_auth_context.workspace, + job_id=test_job_step_with_auth_context.job, + attempt_id=test_job_step_with_auth_context.attempt_id, + step_id=test_job_step_with_auth_context.id, + ) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object(docker_job, "_register_workload_delegation") as register_delegation, + patch.object(docker_job, "_write_workload_identity_subject_token") as write_token, + patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation, + pytest.raises(RuntimeError, match="submit failed"), + ): + docker_job.schedule_single_container( + test_job_step_with_auth_context.step_spec.executor, test_job_step_with_auth_context + ) + + register_delegation.assert_called_once() + write_token.assert_called_once() + revoke_delegation.assert_called_once_with(expected_delegation_name) + assert_created_task_volumes_cleaned_up(docker_client_mock) + assert docker_job._container_start_admission.acquire(blocking=False) + docker_job._container_start_admission.release() + + +def test_docker_sync_does_not_restore_opaque_workload_token_refresher_from_container_labels( docker_job, docker_client_mock, test_job_step ): volume_name = "task-workload-identity-default-job-test-job-id-task-restored" @@ -902,19 +1056,14 @@ def test_docker_sync_restores_workload_identity_refresher_from_container_labels( docker_client_mock.containers.get.side_effect = None docker_client_mock.containers.get.return_value = container - refresher = MagicMock() test_job_step.status = PlatformJobStatus.ACTIVE - with patch.object(docker_job, "_build_workload_identity_refresher", return_value=refresher) as build_refresher: - update = docker_job.sync(test_job_step) + update = docker_job.sync(test_job_step) assert update.status == PlatformJobStatus.ACTIVE - build_refresher.assert_called_once_with(volume_name) - refresher.start.assert_called_once() - assert docker_job._workload_identity_refreshers[container.name] is refresher -def test_docker_sync_restores_workload_identity_refresher_from_mounted_volume( +def test_docker_sync_does_not_restore_opaque_workload_token_refresher_from_mounted_volume( docker_job, docker_client_mock, test_job_step ): volume_name = "task-workload-identity-default-job-test-job-id-task-mounted" @@ -939,181 +1088,40 @@ def test_docker_sync_restores_workload_identity_refresher_from_mounted_volume( docker_client_mock.containers.get.side_effect = None docker_client_mock.containers.get.return_value = container - refresher = MagicMock() test_job_step.status = PlatformJobStatus.ACTIVE - with patch.object(docker_job, "_build_workload_identity_refresher", return_value=refresher) as build_refresher: - update = docker_job.sync(test_job_step) + update = docker_job.sync(test_job_step) assert update.status == PlatformJobStatus.ACTIVE - build_refresher.assert_called_once_with(volume_name) - refresher.start.assert_called_once() - assert docker_job._workload_identity_refreshers[container.name] is refresher - - -def test_docker_stop_workload_identity_refresher_keeps_refresher_when_stop_fails(docker_job): - refresher = MagicMock() - refresher.stop.side_effect = RuntimeError("Timed out stopping workload identity subject token refresher") - docker_job._workload_identity_refreshers["job-container"] = refresher - - with pytest.raises(RuntimeError, match="Timed out stopping workload identity subject token refresher"): - docker_job._stop_workload_identity_refresher("job-container") - - assert docker_job._workload_identity_refreshers["job-container"] is refresher - - -def test_docker_shutdown_stops_workload_identity_refreshers(docker_job, docker_client_mock): - first_refresher = MagicMock() - second_refresher = MagicMock() - docker_job._workload_identity_refreshers = { - "job-container-one": first_refresher, - "job-container-two": second_refresher, - } - close_calls_before_shutdown = docker_client_mock.close.call_count - - docker_job.shutdown() - - first_refresher.stop.assert_called_once() - second_refresher.stop.assert_called_once() - assert docker_job._workload_identity_refreshers == {} - assert docker_client_mock.close.call_count == close_calls_before_shutdown + 1 -def test_docker_shutdown_continues_stopping_refreshers_when_one_stop_fails(docker_job, docker_client_mock): - failing_refresher = MagicMock() - failing_refresher.stop.side_effect = RuntimeError("Timed out stopping workload identity subject token refresher") - second_refresher = MagicMock() - docker_job._workload_identity_refreshers = { - "job-container-one": failing_refresher, - "job-container-two": second_refresher, - } +def test_docker_shutdown_closes_client(docker_job, docker_client_mock): close_calls_before_shutdown = docker_client_mock.close.call_count docker_job.shutdown() - failing_refresher.stop.assert_called_once() - second_refresher.stop.assert_called_once() - assert docker_job._workload_identity_refreshers == {} assert docker_client_mock.close.call_count == close_calls_before_shutdown + 1 -def test_docker_subject_token_issuer_reads_password_from_configured_env_var(docker_job, monkeypatch): - monkeypatch.setenv("AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", "shared-secret") - auth_config = SimpleNamespace( - oidc=SimpleNamespace( - token_endpoint="http://127.0.0.1:18080/application/o/token/", - workload_client_id="nemo-platform-workload", - client_id="nemo-platform-cli", - workload_scope="openid email groups", - ) - ) - docker_job._execution_profile_config.workload_identity = DockerWorkloadIdentityConfig( - username="svc-nemo", - password_env_var="AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", - ) - - with patch("nmp.common.config.get_auth_config", return_value=auth_config): - issuer = docker_job._create_docker_subject_token_issuer() - - assert issuer.token_endpoint == "http://127.0.0.1:18080/application/o/token/" - assert issuer.client_id == "nemo-platform-workload" - assert issuer.username == "svc-nemo" - assert issuer.password == "shared-secret" - assert issuer.scope == "openid email groups" - - -def test_docker_subject_token_issuer_uses_internal_token_endpoint_override(docker_job, monkeypatch): - monkeypatch.setenv("AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", "shared-secret") - auth_config = SimpleNamespace( - oidc=SimpleNamespace( - token_endpoint="http://127.0.0.1:18080/application/o/token/", - workload_client_id="nemo-platform-workload", - client_id="nemo-platform-cli", - workload_scope="openid email groups", - ) - ) - docker_job._execution_profile_config.workload_identity = DockerWorkloadIdentityConfig( - token_endpoint="https://nemo-gateway:8080/application/o/token/", - username="svc-nemo", - password_env_var="AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", - ) - - with patch("nmp.common.config.get_auth_config", return_value=auth_config): - issuer = docker_job._create_docker_subject_token_issuer() - - assert issuer.token_endpoint == "https://nemo-gateway:8080/application/o/token/" - - -@pytest.mark.parametrize("env_value", [None, ""]) -def test_docker_subject_token_issuer_requires_non_empty_password_env_var(docker_job, monkeypatch, env_value): - if env_value is None: - monkeypatch.delenv("AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", raising=False) - else: - monkeypatch.setenv("AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", env_value) - auth_config = SimpleNamespace( - oidc=SimpleNamespace( - token_endpoint="http://127.0.0.1:18080/application/o/token/", - workload_client_id="nemo-platform-workload", - client_id="nemo-platform-cli", - workload_scope="openid email groups", - ) - ) - docker_job._execution_profile_config.workload_identity = DockerWorkloadIdentityConfig( - username="svc-nemo", - password_env_var="AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", - ) - - with ( - patch("nmp.common.config.get_auth_config", return_value=auth_config), - pytest.raises(JobStorageError, match="AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD"), - ): - docker_job._create_docker_subject_token_issuer() - - -def test_docker_workload_identity_config_rejects_legacy_password_field(): - with pytest.raises(ValidationError) as exc_info: - DockerWorkloadIdentityConfig.model_validate( - { - "username": "svc-nemo", - "password": "inline-secret", - } - ) - - assert "password" in str(exc_info.value) - assert "extra" in str(exc_info.value).lower() - - -def test_docker_workload_identity_client_secret_schema_is_write_only(): - client_secret_schema = DockerWorkloadIdentityConfig.model_json_schema()["properties"]["client_secret"] - - assert client_secret_schema["format"] == "password" - assert client_secret_schema["writeOnly"] is True - - -def test_docker_workload_identity_token_timing_constraints(monkeypatch): +def test_docker_workload_identity_config_rejects_per_executor_fields(): properties = DockerWorkloadIdentityConfig.model_json_schema()["properties"] - assert properties["subject_token_ttl_seconds"]["minimum"] == 1 - assert properties["refresh_margin_seconds"]["minimum"] == 0 - assert ( - DockerWorkloadIdentityConfig(subject_token_ttl_seconds=1, refresh_margin_seconds=0).subject_token_ttl_seconds - == 1 - ) - assert DockerWorkloadIdentityConfig(refresh_margin_seconds=0).refresh_margin_seconds == 0 - with pytest.raises(ValidationError): - DockerWorkloadIdentityConfig.model_validate({"subject_token_ttl_seconds": 0}) - with pytest.raises(ValidationError): - DockerWorkloadIdentityConfig.model_validate({"refresh_margin_seconds": -1}) - for values in ( - {"subject_token_ttl_seconds": 1}, - {"subject_token_ttl_seconds": 30, "refresh_margin_seconds": 30}, - {"subject_token_ttl_seconds": 30, "refresh_margin_seconds": 31}, + assert properties == {} + for legacy_field in ( + "enabled", + "token_endpoint", + "client_id", + "client_secret", + "username", + "password", + "password_env_var", + "scope", + "proof_token_provider", + "subject_token_ttl_seconds", + "refresh_margin_seconds", ): - with pytest.raises(ValidationError, match="refresh_margin_seconds"): - DockerWorkloadIdentityConfig.model_validate(values) - monkeypatch.setenv("NMP_WORKLOAD_IDENTITY_TOKEN_TTL_SECONDS", "0") - with pytest.raises(ValidationError): - DockerWorkloadIdentityConfig() + with pytest.raises(ValidationError, match=legacy_field): + DockerWorkloadIdentityConfig.model_validate({legacy_field: "legacy"}) def test_schedule_docker_gpu(mock_nmp_client, docker_client_mock): @@ -2108,6 +2116,49 @@ def test_failed_schedule_logs_status_update_failure_and_releases_admission(docke docker_job._container_start_admission.release() +def test_failed_schedule_revokes_prepared_workload_delegation(docker_job, test_job_step): + assert docker_job._container_start_admission.acquire(blocking=False) + docker_job._run_container_in_thread = MagicMock( + side_effect=FailedToScheduleError("container failed", error_details={"message": "container failed"}) + ) + + with patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation: + docker_job.run_container( + test_job_step, + {"_nmp_workload_delegation_name": "job:prepared-delegation"}, + ) + + revoke_delegation.assert_called_once_with("job:prepared-delegation") + assert docker_job._container_start_admission.acquire(blocking=False) + docker_job._container_start_admission.release() + + +def test_terminal_step_update_revokes_workload_delegation(docker_job, test_job_step_with_auth_context): + container = MagicMock() + container.name = "job-test-job-id-test-step" + container.labels = {JOB_TASK_ID_LABEL: "task-success"} + docker_job.map_docker_container_status_to_platform_status = MagicMock( + return_value=(PlatformJobStatus.COMPLETED, {}, "") + ) + docker_job.docker_state_debug_fields = MagicMock(return_value={}) + expected_delegation_name = docker_delegation_name( + workload_workspace=test_job_step_with_auth_context.workspace, + job_id=test_job_step_with_auth_context.job, + attempt_id=test_job_step_with_auth_context.attempt_id, + step_id=test_job_step_with_auth_context.id, + ) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_token_exchange_auth_config()), + patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation, + ): + update = docker_job.create_step_update(test_job_step_with_auth_context, container) + + assert update.status == PlatformJobStatus.COMPLETED + revoke_delegation.assert_called_once_with(expected_delegation_name) + docker_job._jobs.update_job_step_task.assert_called_once() + + def test_resuming_step_skips_before_active_ttl_enforcement(docker_job, test_job_step): """RESUMING must not apply ttl_seconds_before_active (pause/resume rebasing).""" ttl_seconds = docker_job._execution_profile_config.ttl_seconds_before_active @@ -2714,6 +2765,41 @@ def test_cleanup_single_container_without_persistent_storage_label(docker_job, d docker_job.cleanup_job_persistent_storage.assert_not_called() +def test_cleanup_single_container_revokes_workload_delegation_from_labels(docker_job, docker_client_mock): + mock_container = MagicMock() + mock_container.name = "test-container-workload-identity" + mock_container.id = "workload-identity-container-id" + mock_container.attrs = { + "State": {"ExitCode": 0}, + } + mock_container.labels = { + JOB_WORKSPACE_ID_LABEL: "default", + JOB_ID_LABEL: "test-job-id", + JOB_ATTEMPT_ID_LABEL: "test-attempt-id", + JOB_STEP_ID_LABEL: "test-step-id", + JOB_STEP_NAME_LABEL: "test-step", + JOB_TASK_ID_LABEL: "task-success", + DOCKER_WORKLOAD_IDENTITY_TOKEN_FILE_LABEL: WORKLOAD_IDENTITY_TOKEN_FILE_PATH, + JOB_MANAGED_BY_LABEL: JOB_MANAGED_BY_JOBS_CONTROLLER, + JOB_CONTROLLER_INSTANCE_ID_LABEL: TEST_JOBS_CONTROLLER_INSTANCE_ID, + } + mock_volume = MagicMock() + docker_client_mock.volumes.get.return_value = mock_volume + expected_delegation_name = docker_delegation_name( + workload_workspace="default", + job_id="test-job-id", + attempt_id="test-attempt-id", + step_id="test-step-id", + ) + + with patch.object(docker_job, "_revoke_workload_delegation") as revoke_delegation: + docker_job.cleanup_single_container(mock_container) + + revoke_delegation.assert_called_once_with(expected_delegation_name) + assert mock_container.remove.call_count == 1 + assert docker_client_mock.volumes.get.call_count == 3 + + def test_cleanup_single_container_step_terminal_but_job_has_more_steps(docker_job, docker_client_mock): """Test that persistent storage is NOT cleaned up when a step is terminal but the job has more steps to run. diff --git a/services/core/jobs/tests/controllers/test_kubernetes_backend.py b/services/core/jobs/tests/controllers/test_kubernetes_backend.py index 2b0dcd11b4..a60b5d16c0 100644 --- a/services/core/jobs/tests/controllers/test_kubernetes_backend.py +++ b/services/core/jobs/tests/controllers/test_kubernetes_backend.py @@ -10,6 +10,11 @@ from jsonschema.exceptions import ValidationError as JsonSchemaValidationError from kubernetes import client from kubernetes.client.rest import ApiException +from nmp.common.auth import ( + NMP_PRINCIPAL_ENVVAR, + WorkloadDelegationConflictError, + reference_delegation_name, +) from nmp.common.config import ImagePullSecret, PlatformConfig from nmp.common.jobs.constants import ( EPHEMERAL_TASK_STORAGE_PATH_ENVVAR, @@ -73,6 +78,10 @@ delete_configmap, name_for_step, ) +from nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job import ( + KUBERNETES_POD_UID_REFERENCE_NAME, +) +from nmp.core.jobs.controllers.backends.workload_tokens import WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS from pydantic import ValidationError DEFAULT_STORAGE = KubernetesJobStorageConfig(pvc_name="job-storage-pvc") @@ -82,6 +91,83 @@ DEFAULT_POD_METADATA = KubernetesObjectMetadata(labels={"foo": "bar"}, annotations={"example.com/annotation": "value"}) +def _kubernetes_job_for_step( + step: PlatformJobStepWithContext, + *, + status: client.V1JobStatus | None = None, + service_account_name: str = "default", +) -> client.V1Job: + labels = common_labels_for_step(step) + labels[JOB_EXECUTION_BACKEND_LABEL] = "kubernetes_job" + labels[JOB_EXECUTION_PROFILE_LABEL] = "default" + return client.V1Job( + metadata=client.V1ObjectMeta( + name=name_for_step(step), + namespace="test-namespace", + labels=labels, + ), + spec=client.V1JobSpec( + template=client.V1PodTemplateSpec( + spec=client.V1PodSpec( + containers=[], + service_account_name=service_account_name, + ) + ) + ), + status=status + or client.V1JobStatus( + active=1, + completion_time=None, + failed=None, + succeeded=None, + ), + ) + + +def _kubernetes_pod(uid: str, *, phase: str = "Running") -> client.V1Pod: + return client.V1Pod( + metadata=client.V1ObjectMeta( + name=f"pod-{uid}", + namespace="test-namespace", + uid=uid, + ), + status=client.V1PodStatus( + phase=phase, + container_statuses=[ + client.V1ContainerStatus( + image="test-image", + image_id="test-image-id", + name="nemo-job-task", + ready=phase == "Running", + restart_count=0, + state=client.V1ContainerState( + running=client.V1ContainerStateRunning(started_at=datetime.datetime.now(datetime.timezone.utc)) + if phase == "Running" + else None, + terminated=client.V1ContainerStateTerminated( + exit_code=0, + reason="Completed", + ) + if phase == "Succeeded" + else None, + ), + ) + ], + ), + ) + + +@pytest.fixture +def workload_exchange_auth_config(): + return SimpleNamespace( + oidc=SimpleNamespace( + workload_token_exchange_enabled=True, + workload_audience="nemo-platform", + audience=None, + ) + ) + + @pytest.fixture def kubernetes_client_mock(): """Mock Kubernetes client for testing.""" @@ -909,14 +995,14 @@ def test_kubernetes_job_rejects_reserved_step_auth_env_vars(kubernetes_job, cpu_ def test_kubernetes_job_injects_projected_workload_identity_token_when_exchange_enabled( - kubernetes_job, cpu_execution_provider, test_step_pending + kubernetes_job, cpu_execution_provider, test_step_pending_with_auth_context ): kubernetes_job._execution_profile_config.workload_identity_token_expiration_seconds = 600 kubernetes_job._execution_profile_config.workload_identity_token_audience = "test-audience" auth_config = SimpleNamespace(oidc=SimpleNamespace(workload_token_exchange_enabled=True)) with patch("nmp.common.config.get_auth_config", return_value=auth_config): - kubernetes_job.schedule(cpu_execution_provider, test_step_pending) + kubernetes_job.schedule(cpu_execution_provider, test_step_pending_with_auth_context) call_args = kubernetes_job._batch_v1.create_namespaced_job.call_args job_body = call_args.kwargs["body"] @@ -925,6 +1011,7 @@ def test_kubernetes_job_injects_projected_workload_identity_token_when_exchange_ env_vars = {env.name: env.value for env in main_container.env} assert env_vars[WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR] == WORKLOAD_IDENTITY_TOKEN_FILE_PATH + assert NMP_PRINCIPAL_ENVVAR not in env_vars workload_identity_volume = next( volume for volume in pod_spec.volumes if volume.name == WORKLOAD_IDENTITY_VOLUME_NAME @@ -939,6 +1026,25 @@ def test_kubernetes_job_injects_projected_workload_identity_token_when_exchange_ assert mount.read_only is True +def test_kubernetes_job_does_not_mount_workload_identity_without_auth_context( + kubernetes_job, cpu_execution_provider, test_step_pending +): + auth_config = SimpleNamespace(oidc=SimpleNamespace(workload_token_exchange_enabled=True)) + + with patch("nmp.common.config.get_auth_config", return_value=auth_config): + kubernetes_job.schedule(cpu_execution_provider, test_step_pending) + + call_args = kubernetes_job._batch_v1.create_namespaced_job.call_args + job_body = call_args.kwargs["body"] + pod_spec = job_body.spec.template.spec + main_container = pod_spec.containers[0] + + env_vars = {env.name: env.value for env in main_container.env} + assert WORKLOAD_IDENTITY_TOKEN_FILE_ENVVAR not in env_vars + assert all(volume.name != WORKLOAD_IDENTITY_VOLUME_NAME for volume in pod_spec.volumes) + assert all(mount.name != WORKLOAD_IDENTITY_VOLUME_NAME for mount in main_container.volume_mounts) + + def test_schedule_job_with_args(kubernetes_job, cpu_execution_provider, test_step_pending): """Test job scheduling with custom args.""" @@ -2197,8 +2303,6 @@ def test_kubernetes_job_schedule_with_auth_context( def test_kubernetes_job_schedule_without_auth_context(kubernetes_job, cpu_execution_provider, test_step_pending): """Test that auth env vars are NOT set when auth_context is absent.""" - from nmp.common.auth.models import NMP_PRINCIPAL_ENVVAR - # Mock successful job creation kubernetes_job._batch_v1.create_namespaced_job.return_value = MagicMock() @@ -2224,6 +2328,242 @@ def test_kubernetes_job_schedule_without_auth_context(kubernetes_job, cpu_execut assert "NMP_JOB_LAUNCHER_OTLP_LOGS_PROTOCOL" not in env_var_names +def test_kubernetes_job_registers_pod_uid_workload_delegation( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.PENDING + ttl_seconds_active = 900 + kubernetes_job._execution_profile_config.ttl_seconds_active = ttl_seconds_active + pod = _kubernetes_pod("pod-uid-123") + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + before_sync = datetime.datetime.now(datetime.timezone.utc) + + expected_name = reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:test-namespace:default", + bound_reference_name=KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value="pod-uid-123", + ) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_register_workload_delegation") as register_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert update.status == PlatformJobStatus.ACTIVE + register_delegation.assert_called_once() + delegation = register_delegation.call_args.args[0] + assert delegation.name == expected_name + assert delegation.workload_subject == "system:serviceaccount:test-namespace:default" + assert delegation.workload_audience == "nemo-platform" + assert delegation.workload_workspace == test_step_pending_with_auth_context.workspace + assert delegation.job_id == test_step_pending_with_auth_context.job + assert delegation.attempt_id == test_step_pending_with_auth_context.attempt_id + assert delegation.step_id == test_step_pending_with_auth_context.id + assert delegation.auth_context.principal_id == "creator@example.com" + assert delegation.bound_reference_name == KUBERNETES_POD_UID_REFERENCE_NAME + assert delegation.bound_reference_value == "pod-uid-123" + assert delegation.opaque_subject_token_hash is None + expected_ttl_seconds = ttl_seconds_active + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS + assert ( + before_sync + datetime.timedelta(seconds=expected_ttl_seconds) + <= delegation.expires_at + <= datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(seconds=expected_ttl_seconds) + ) + + +def test_kubernetes_job_skips_workload_delegation_when_pod_uid_is_missing( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.PENDING + pod = _kubernetes_pod("pod-uid-123") + pod.metadata.uid = None + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_register_workload_delegation") as register_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert update.status == PlatformJobStatus.ACTIVE + register_delegation.assert_not_called() + + +def test_kubernetes_job_skips_workload_delegation_when_job_labels_are_missing( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.PENDING + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + del k8s_job.metadata.labels[JOB_STEP_ID_LABEL] + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_register_workload_delegation") as register_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert update.status == PlatformJobStatus.PENDING + register_delegation.assert_not_called() + + +def test_kubernetes_job_workload_delegation_conflict_is_not_recreated( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.PENDING + pod = _kubernetes_pod("pod-uid-123") + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object( + kubernetes_job, + "_register_workload_delegation", + side_effect=WorkloadDelegationConflictError("exists"), + ) as register_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + first_update = kubernetes_job.sync(test_step_pending_with_auth_context) + second_update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert first_update.status == PlatformJobStatus.ACTIVE + assert second_update.status == PlatformJobStatus.ACTIVE + register_delegation.assert_called_once() + + +def test_kubernetes_job_registration_failure_does_not_block_status_and_retries( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.PENDING + pod = _kubernetes_pod("pod-uid-123") + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object( + kubernetes_job, "_register_workload_delegation", side_effect=RuntimeError("store offline") + ) as register_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + first_update = kubernetes_job.sync(test_step_pending_with_auth_context) + second_update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert first_update.status == PlatformJobStatus.ACTIVE + assert second_update.status == PlatformJobStatus.ACTIVE + assert register_delegation.call_count == 2 + job_key = f"test-namespace/{name_for_step(test_step_pending_with_auth_context)}" + assert kubernetes_job._workload_delegations_by_job.get(job_key, set()) == set() + + +def test_kubernetes_job_revokes_pod_uid_workload_delegation_when_job_finishes( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.ACTIVE + pod = _kubernetes_pod("pod-uid-123", phase="Succeeded") + k8s_job = _kubernetes_job_for_step( + test_step_pending_with_auth_context, + status=client.V1JobStatus( + active=None, + completion_time=datetime.datetime.now(datetime.timezone.utc), + failed=None, + succeeded=1, + ), + ) + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + + expected_name = reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:test-namespace:default", + bound_reference_name=KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value="pod-uid-123", + ) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_revoke_workload_delegation") as revoke_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert update.status == PlatformJobStatus.COMPLETED + revoke_delegation.assert_called_once_with(expected_name) + + +def test_kubernetes_job_revokes_recorded_workload_delegation_when_pods_are_gone( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + test_step_pending_with_auth_context.status = PlatformJobStatus.ACTIVE + k8s_job = _kubernetes_job_for_step( + test_step_pending_with_auth_context, + status=client.V1JobStatus( + active=None, + completion_time=datetime.datetime.now(datetime.timezone.utc), + failed=None, + succeeded=1, + ), + ) + job_key = f"test-namespace/{name_for_step(test_step_pending_with_auth_context)}" + kubernetes_job._workload_delegations_by_job[job_key] = {"ref:recorded-delegation"} + kubernetes_job._batch_v1.read_namespaced_job.return_value = k8s_job + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[]) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_revoke_workload_delegation") as revoke_delegation, + patch("nmp.core.jobs.controllers.backends.kubernetes.kubernetes_job.update_all_tasks", return_value=False), + patch.object(kubernetes_job, "get_kube_job_events", return_value=[]), + ): + update = kubernetes_job.sync(test_step_pending_with_auth_context) + + assert update.status == PlatformJobStatus.COMPLETED + revoke_delegation.assert_called_once_with("ref:recorded-delegation") + assert job_key not in kubernetes_job._workload_delegations_by_job + + +def test_kubernetes_job_revokes_pod_uid_workload_delegation_when_job_deleted( + kubernetes_job, test_step_pending_with_auth_context, workload_exchange_auth_config +): + pod = _kubernetes_pod("pod-uid-123", phase="Succeeded") + k8s_job = _kubernetes_job_for_step(test_step_pending_with_auth_context) + kubernetes_job._core_v1.list_namespaced_pod.return_value = client.V1PodList(items=[pod]) + + expected_name = reference_delegation_name( + workload_audience="nemo-platform", + workload_subject="system:serviceaccount:test-namespace:default", + bound_reference_name=KUBERNETES_POD_UID_REFERENCE_NAME, + bound_reference_value="pod-uid-123", + ) + + with ( + patch("nmp.common.config.get_auth_config", return_value=workload_exchange_auth_config), + patch.object(kubernetes_job, "_revoke_workload_delegation") as revoke_delegation, + ): + kubernetes_job.terminate_job(k8s_job) + + revoke_delegation.assert_called_once_with(expected_name) + kubernetes_job._batch_v1.delete_namespaced_job.assert_called_once() + + def test_cleanup_steps_with_multi_step_job_only_first_step_complete(kubernetes_job): """Test cleanup_steps with a multi-step job where only the first step is complete. diff --git a/services/core/jobs/tests/controllers/test_workload_tokens.py b/services/core/jobs/tests/controllers/test_workload_tokens.py index f91d5f0b3b..55ff761954 100644 --- a/services/core/jobs/tests/controllers/test_workload_tokens.py +++ b/services/core/jobs/tests/controllers/test_workload_tokens.py @@ -1,17 +1,13 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import datetime import tarfile -import threading -import time -import httpx -import pytest from nmp.core.jobs.controllers.backends.workload_tokens import ( - OAuthPasswordGrantSubjectTokenIssuer, - SubjectToken, - SubjectTokenRefreshLoop, + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS, build_token_archive, + workload_delegation_expires_at, ) @@ -27,369 +23,20 @@ def test_build_token_archive_contains_read_only_token_file() -> None: assert extracted.read() == b"subject-token" -def test_oauth_password_grant_issuer_requests_subject_token(monkeypatch: pytest.MonkeyPatch) -> None: - captured: dict = {} +def test_workload_delegation_expires_at_adds_active_ttl_and_cleanup_buffer() -> None: + now = datetime.datetime(2026, 8, 10, 12, 0, tzinfo=datetime.timezone.utc) - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - captured["url"] = url - captured["data"] = data - captured["timeout"] = timeout - return httpx.Response(200, json={"access_token": "subject-token", "expires_in": 120}) + expires_at = workload_delegation_expires_at(ttl_seconds_active=900, now=now) - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) + assert expires_at == now + datetime.timedelta(seconds=900 + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS) - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - client_secret="secret", - username="svc-nemo", - password="app-password", - scope="openid email groups", - timeout=5.0, - ) - before = time.time() - - token = issuer.issue() - - assert token.value == "subject-token" - assert token.expires_at >= before + 120 - assert captured == { - "url": "https://idp.example.com/token", - "data": { - "grant_type": "password", - "client_id": "nemo-platform", - "client_secret": "secret", - "username": "svc-nemo", - "password": "app-password", - "scope": "openid email groups", - }, - "timeout": 5.0, - } - - -def test_oauth_password_grant_issuer_reports_idp_errors(monkeypatch: pytest.MonkeyPatch) -> None: - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - return httpx.Response( - 400, - json={"error": "invalid_grant", "error_description": "bad credentials"}, - headers={"content-type": "application/json"}, - ) - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - username="svc-nemo", - password="bad-password", - ) - - with pytest.raises(RuntimeError, match="invalid_grant - bad credentials"): - issuer.issue() - - -@pytest.mark.parametrize( - "response", - [ - httpx.Response(400, json=["invalid"], headers={"content-type": "application/json"}), - httpx.Response(400, json="invalid", headers={"content-type": "application/json"}), - httpx.Response(400, json=None, headers={"content-type": "application/json"}), - httpx.Response(400, content=b"not-json", headers={"content-type": "application/json"}), - ], -) -def test_oauth_password_grant_issuer_ignores_non_object_error_json( - monkeypatch: pytest.MonkeyPatch, response: httpx.Response -) -> None: - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - return response - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - username="svc-nemo", - password="bad-password", - ) - - with pytest.raises(RuntimeError, match="unknown_error - "): - issuer.issue() - - -def test_oauth_password_grant_issuer_rejects_non_loopback_http_before_sending_credentials( - monkeypatch: pytest.MonkeyPatch, -) -> None: - def fake_post(*args, **kwargs) -> httpx.Response: - raise AssertionError("token endpoint should not be called") - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="http://authentik-server:9000/application/o/token/", - client_id="nemo-platform", - username="svc-nemo", - password="app-password", - ) - - with pytest.raises(RuntimeError, match="token_endpoint must use https://"): - issuer.issue() - - -@pytest.mark.parametrize("token_endpoint", ["http://localhost:18080/token", "http://127.0.0.1:18080/token"]) -def test_oauth_password_grant_issuer_allows_loopback_http_for_local_development( - monkeypatch: pytest.MonkeyPatch, token_endpoint: str -) -> None: - captured: dict[str, object] = {} - - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - captured["url"] = url - return httpx.Response(200, json={"access_token": "subject-token", "expires_in": 120}) - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint=token_endpoint, - client_id="nemo-platform", - username="svc-nemo", - password="app-password", - ) - - assert issuer.issue().value == "subject-token" - assert captured["url"] == token_endpoint - - -def test_oauth_password_grant_issuer_uses_configured_default_expiry(monkeypatch: pytest.MonkeyPatch) -> None: - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - return httpx.Response(200, json={"access_token": "subject-token"}) - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - username="svc-nemo", - password="app-password", - default_expires_in_seconds=45, - ) - before = time.time() - - token = issuer.issue() - - assert token.value == "subject-token" - assert token.expires_at >= before + 45 - - -@pytest.mark.parametrize( - ("payload", "message"), - [ - ([], "Token endpoint response was not a JSON object"), - ({}, "Token endpoint response did not include a non-empty access_token"), - ({"access_token": ""}, "Token endpoint response did not include a non-empty access_token"), - ({"access_token": None}, "Token endpoint response did not include a non-empty access_token"), - ( - {"access_token": "subject-token", "expires_in": None}, - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - {"access_token": "subject-token", "expires_in": "120"}, - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - {"access_token": "subject-token", "expires_in": 0}, - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - {"access_token": "subject-token", "expires_in": -1}, - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - b'{"access_token": "subject-token", "expires_in": NaN}', - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - b'{"access_token": "subject-token", "expires_in": Infinity}', - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - b'{"access_token": "subject-token", "expires_in": -Infinity}', - "Token endpoint response did not include a positive numeric expires_in", - ), - ( - {"access_token": "subject-token", "expires_in": True}, - "Token endpoint response did not include a positive numeric expires_in", - ), - ], -) -def test_oauth_password_grant_issuer_rejects_invalid_success_response( - monkeypatch: pytest.MonkeyPatch, payload: object, message: str -) -> None: - def fake_post(url: str, *, data: dict, timeout: float) -> httpx.Response: - if isinstance(payload, bytes): - return httpx.Response(200, content=payload) - return httpx.Response(200, json=payload) - - monkeypatch.setattr("nmp.core.jobs.controllers.backends.workload_tokens.httpx.post", fake_post) - - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - username="svc-nemo", - password="app-password", - ) - with pytest.raises(RuntimeError, match=f"invalid_response - {message}"): - issuer.issue() +def test_workload_delegation_expires_at_normalizes_naive_now_to_utc() -> None: + now = datetime.datetime(2026, 8, 10, 12, 0) + expires_at = workload_delegation_expires_at(ttl_seconds_active=60, now=now) -def test_oauth_password_grant_issuer_repr_excludes_secrets() -> None: - issuer = OAuthPasswordGrantSubjectTokenIssuer( - token_endpoint="https://idp.example.com/token", - client_id="nemo-platform", - client_secret="super-sensitive-client-secret", - username="svc-nemo", - password="super-sensitive-password", + assert expires_at.tzinfo is datetime.timezone.utc + assert expires_at == now.replace(tzinfo=datetime.timezone.utc) + datetime.timedelta( + seconds=60 + WORKLOAD_DELEGATION_TTL_BUFFER_SECONDS ) - - issuer_repr = repr(issuer) - - assert "client_secret=" not in issuer_repr - assert "password=" not in issuer_repr - assert "super-sensitive-client-secret" not in issuer_repr - assert "super-sensitive-password" not in issuer_repr - - -def test_refresh_loop_refresh_once_writes_issued_token() -> None: - class FakeIssuer: - def issue(self) -> SubjectToken: - return SubjectToken(value="subject-token", expires_at=time.time() + 120) - - writes: list[str] = [] - refresher = SubjectTokenRefreshLoop(issuer=FakeIssuer(), write_token=writes.append) - - token = refresher.refresh_once() - - assert token.value == "subject-token" - assert writes == ["subject-token"] - - -def test_refresh_loop_can_restart_after_stop() -> None: - class FakeIssuer: - def __init__(self) -> None: - self.issued = 0 - - def issue(self) -> SubjectToken: - self.issued += 1 - return SubjectToken(value=f"subject-token-{self.issued}", expires_at=time.time() + 120) - - writes: list[str] = [] - wrote = threading.Event() - - def write_token(token: str) -> None: - writes.append(token) - wrote.set() - - refresher = SubjectTokenRefreshLoop( - issuer=FakeIssuer(), - write_token=write_token, - min_sleep_seconds=0.01, - ) - - refresher.start() - assert wrote.wait(timeout=1) - refresher.stop() - - wrote.clear() - refresher.start() - assert wrote.wait(timeout=1) - refresher.stop() - - assert writes[:2] == ["subject-token-1", "subject-token-2"] - - -def test_refresh_loop_stop_timeout_preserves_thread_and_prevents_late_write() -> None: - issue_started = threading.Event() - release_issue = threading.Event() - writes: list[str] = [] - - class BlockingIssuer: - def issue(self) -> SubjectToken: - issue_started.set() - if not release_issue.wait(timeout=1): - raise RuntimeError("issuer was not released") - return SubjectToken(value="late-subject-token", expires_at=time.time() + 120) - - refresher = SubjectTokenRefreshLoop( - issuer=BlockingIssuer(), - write_token=writes.append, - min_sleep_seconds=0.01, - ) - refresher._stop_timeout_seconds = 0.01 - - refresher.start() - assert issue_started.wait(timeout=1) - - with pytest.raises(RuntimeError, match="Timed out stopping workload identity subject token refresher"): - refresher.stop() - - thread = refresher._thread - assert thread is not None - assert thread.is_alive() - - release_issue.set() - thread.join(timeout=1) - assert not thread.is_alive() - - refresher.stop() - - assert refresher._thread is None - assert writes == [] - - -def test_refresh_loop_backs_off_failures_and_resets_after_success() -> None: - class FakeIssuer: - def __init__(self) -> None: - self.calls = 0 - - def issue(self) -> SubjectToken: - self.calls += 1 - if self.calls in {1, 2, 3, 5}: - raise RuntimeError("idp unavailable") - return SubjectToken(value="subject-token", expires_at=time.time()) - - class FakeStop(threading.Event): - def __init__(self) -> None: - super().__init__() - self.waits: list[float] = [] - self.stopped = False - - def is_set(self) -> bool: - return self.stopped - - def wait(self, timeout: float | None = None) -> bool: - assert timeout is not None - self.waits.append(timeout) - if len(self.waits) == 5: - self.stopped = True - return True - return False - - stop = FakeStop() - writes: list[str] = [] - refresher = SubjectTokenRefreshLoop( - issuer=FakeIssuer(), - write_token=writes.append, - min_sleep_seconds=1.0, - max_failure_backoff_seconds=4.0, - ) - refresher._stop = stop - - refresher._run() - - assert writes == ["subject-token"] - assert stop.waits == [1.0, 2.0, 4.0, 1.0, 1.0] - - -def test_subject_token_seconds_until_refresh_never_negative() -> None: - token = SubjectToken(value="expired", expires_at=time.time() - 10) - - assert token.seconds_until_refresh(margin_seconds=60) == 0.0 diff --git a/services/core/jobs/tests/test_config.py b/services/core/jobs/tests/test_config.py index d2c2e28dc3..fb243b19f3 100644 --- a/services/core/jobs/tests/test_config.py +++ b/services/core/jobs/tests/test_config.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch import pytest +import yaml from nemo_platform_plugin.capabilities import ProbeResult from nmp.common.config import Configuration, Runtime from nmp.core.jobs.app.providers import ( @@ -236,6 +237,22 @@ def test_jobs_config_merge_with_defaults_docker_additional_volumes(): assert additional[0].allow_create_volume is True +def test_authentik_compose_jobs_config_omits_docker_workload_identity_override(): + global_settings = yaml.safe_load( + pathlib.Path("contrib/auth/authentik/config/platform-compose-authentik.yaml").read_text(encoding="utf-8") + ) + + config = Configuration.global_settings_to_service_config(global_settings, JobsServiceConfig) + + workload_executor = next( + executor + for executor in config.executors + if executor.backend == "docker" and executor.provider == "cpu" and executor.profile == "workload" + ) + assert isinstance(workload_executor.config, DockerJobExecutionProfileConfig) + assert workload_executor.config.workload_identity.model_dump() == {} + + def test_docker_default_profiles(monkeypatch): # Clear caches to ensure fresh config read Configuration.clear_cache() diff --git a/tests/auth_idp/authentik_live.py b/tests/auth_idp/authentik_live.py index 1a6cda907d..2376918245 100644 --- a/tests/auth_idp/authentik_live.py +++ b/tests/auth_idp/authentik_live.py @@ -199,11 +199,6 @@ def prepare_authentik_compose_inputs(*, root: Path = AUTHENTIK_ROOT) -> None: } ] }, - "workload_identity": { - "token_endpoint": "https://nemo-gateway:8080/application/o/token/", - "username": "svc-nemo", - "password_env_var": "AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", - }, }, } ] diff --git a/tests/auth_idp/contracts/test_jobs.py b/tests/auth_idp/contracts/test_jobs.py index 169579f1bd..7dda2b2206 100644 --- a/tests/auth_idp/contracts/test_jobs.py +++ b/tests/auth_idp/contracts/test_jobs.py @@ -21,6 +21,7 @@ def test_provider_workload_job_runs_via_workload_profile( ): require_capability(auth_idp_case, "workspace_rbac") require_capability(auth_idp_case, "workload_job") + require_capability(auth_idp_case, "managed_workload_job_obo") e2e_setup_sdk = auth_idp_runtime.e2e_setup_sdk() for principal in auth_idp_runtime.workload_role_principals(): @@ -28,10 +29,11 @@ def test_provider_workload_job_runs_via_workload_profile( e2e_setup_sdk, workspace=auth_idp_workspace, principal=principal, - roles=["Viewer", "JobRunner"], + roles=["Viewer", "Editor", "JobRunner"], ) - job = e2e_setup_sdk.jobs.create( + job_submitter_sdk = auth_idp_runtime.workload_provider_sdk() + job = job_submitter_sdk.jobs.create( workspace=auth_idp_workspace, source=f"{auth_idp_case.id}-workload-job", spec={"test": "workload-job"}, @@ -44,11 +46,20 @@ def test_provider_workload_job_runs_via_workload_profile( "profile": "workload", "container": { "image": nmp_api_image(), - "entrypoint": ["nemo-platform"], + "entrypoint": ["sh", "-c"], "command": [ - "run", - "task", - "--task", + 'if [ -n "${NMP_PRINCIPAL:-}" ]; then ' + "echo 'Unexpected NMP_PRINCIPAL in managed workload'; exit 42; " + "fi; " + 'if [ -z "${NMP_WORKLOAD_IDENTITY_TOKEN_FILE:-}" ]; then ' + "echo 'Missing NMP_WORKLOAD_IDENTITY_TOKEN_FILE in managed workload'; exit 43; " + "fi; " + 'if [ ! -f "${NMP_WORKLOAD_IDENTITY_TOKEN_FILE}" ]; then ' + "echo 'Workload identity token file is missing'; exit 44; " + "fi; " + "echo 'Workload auth env: NMP_PRINCIPAL=absent " + "NMP_WORKLOAD_IDENTITY_TOKEN_FILE=present'; " + "exec nemo-platform run task --task " "nmp.hello_world.tasks.workload_workspace_get", ], }, @@ -70,4 +81,8 @@ def test_provider_workload_job_runs_via_workload_profile( assert all(log.job_step == "workload-workspace-get" for log in step_logs.data) assert all(log.job_task for log in step_logs.data) assert all(log.message.strip() for log in step_logs.data) + assert any( + "Workload auth env: NMP_PRINCIPAL=absent NMP_WORKLOAD_IDENTITY_TOKEN_FILE=present" in log.message + for log in step_logs.data + ) assert any(f"Successfully retrieved workspace: {auth_idp_workspace}" in log.message for log in step_logs.data) diff --git a/tests/auth_idp/static/test_authentik_kubernetes_demo.py b/tests/auth_idp/static/test_authentik_kubernetes_demo.py index 409f042c76..6f66539da1 100644 --- a/tests/auth_idp/static/test_authentik_kubernetes_demo.py +++ b/tests/auth_idp/static/test_authentik_kubernetes_demo.py @@ -28,12 +28,26 @@ ENVOY_CONTROLLER_ENV_URL = "https://nemo-platform-envoy.$(POD_NAMESPACE).svc.cluster.local:8080" AUTHENTIK_SERVICE_URL_TEMPLATE = '{{ include "nemo-platform-authentik.serviceUrl" (dict "root" . "serviceName" "authentik-server" "scheme" "http") }}' PUBLIC_GATEWAY_URL_TEMPLATE = '{{ include "nemo-platform-authentik.publicGatewayUrl" . }}' +AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS = { + "x-nmp-principal-id", + "x-nmp-principal-email", + "x-nmp-principal-groups", + "x-nmp-principal-on-behalf-of", + "x-nmp-principal-on-behalf-of-email", + "x-nmp-principal-on-behalf-of-groups", + "x-nmp-scopes", +} +SEALED_EXTERNAL_PRINCIPAL_HEADERS = AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS def _load_yaml(path: Path) -> dict: return yaml.safe_load(path.read_text(encoding="utf-8")) +def _exact_header_patterns(patterns: list[dict]) -> set[str]: + return {pattern["exact"] for pattern in patterns if "exact" in pattern} + + def _literal_run_commands(path: Path) -> set[tuple[str, ...]]: tree = ast.parse(path.read_text(encoding="utf-8")) commands: set[tuple[str, ...]] = set() @@ -797,6 +811,8 @@ def test_authentik_umbrella_values_configure_nemo_envoy_as_the_only_edge_proxy() ) in lua_code assert 'headers:remove("x-nmp-authorized")' in lua_code assert 'headers:remove("x-nmp-scopes")' in lua_code + for header in SEALED_EXTERNAL_PRINCIPAL_HEADERS: + assert f'headers:remove("{header}")' in lua_code assert "claim_to_headers" not in yaml.safe_dump(http_manager) assert all( @@ -818,10 +834,7 @@ def test_authentik_umbrella_values_configure_nemo_envoy_as_the_only_edge_proxy() "timeout": "5s", } allowed_headers = ext_authz["http_service"]["authorization_response"]["allowed_upstream_headers"]["patterns"] - assert {"exact": "x-nmp-principal-id"} in allowed_headers - assert {"exact": "x-nmp-principal-email"} in allowed_headers - assert {"exact": "x-nmp-principal-groups"} in allowed_headers - assert {"exact": "x-nmp-scopes"} in allowed_headers + assert AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS.issubset(_exact_header_patterns(allowed_headers)) protected_api_route = next(route for route in routes if route["match"] == {"prefix": "/apis/"}) assert "typed_per_filter_config" not in protected_api_route @@ -962,6 +975,22 @@ def test_authentik_umbrella_chart_applies_blueprint_with_waitable_helm_hook() -> assert "nmp.nvidia.com/blueprint-checksum" not in template +@pytest.mark.auth_idp_k8s +def test_authentik_tokenreview_rbac_does_not_grant_pod_or_job_reads() -> None: + template = (HELM_DIR / "templates" / "tokenreview-rbac.yaml").read_text(encoding="utf-8") + cluster_role = yaml.safe_load(template.split("---", maxsplit=1)[0]) + rules = cluster_role["rules"] + permissions = { + (tuple(rule.get("apiGroups", [])), tuple(rule.get("resources", [])), tuple(rule.get("verbs", []))) + for rule in rules + } + + assert (("authentication.k8s.io",), ("tokenreviews",), ("create",)) in permissions + granted_resources = {resource for rule in rules for resource in rule.get("resources", [])} + assert not {"pods", "pods/log", "jobs", "jobs/status"} & granted_resources + assert ".Values.integration.nemoPlatform.apiServiceAccountName" in template + + def test_authentik_kubernetes_runner_uses_helm_not_kustomize() -> None: run_sh = (AUTHENTIK_DIR / "run.sh").read_text(encoding="utf-8") live_test_path = Path("tests/auth_idp/k8s/test_authentik_kubernetes_live.py") @@ -1158,6 +1187,19 @@ def test_authentik_compose_runner_uses_nemo_scoped_ca_bundle() -> None: assert not any(line.startswith('REQUESTS_CA_BUNDLE="$(gateway_tls_cert_file)"') for line in stripped_lines) +def test_authentik_compose_e2e_config_uses_global_docker_workload_identity_switch_only() -> None: + from tests.auth_idp.authentik_live import AUTHENTIK_DOCKER_E2E_CONFIG + + config_overlay = AUTHENTIK_DOCKER_E2E_CONFIG.mark.args[2] + workload_executor = next( + executor + for executor in config_overlay["jobs"]["executors"] + if executor["backend"] == "docker" and executor["provider"] == "cpu" and executor["profile"] == "workload" + ) + + assert "workload_identity" not in workload_executor["config"] + + def test_authentik_umbrella_values_configure_workload_token_tls() -> None: values = _load_yaml(HELM_DIR / "values.yaml") tls_values = values["workloadTokenTls"] diff --git a/tests/auth_idp/static/test_provider_layout.py b/tests/auth_idp/static/test_provider_layout.py index 62931f898d..8db383c4c8 100644 --- a/tests/auth_idp/static/test_provider_layout.py +++ b/tests/auth_idp/static/test_provider_layout.py @@ -10,6 +10,21 @@ pytestmark = [pytest.mark.auth_idp] +AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS = { + "x-nmp-principal-id", + "x-nmp-principal-email", + "x-nmp-principal-groups", + "x-nmp-principal-on-behalf-of", + "x-nmp-principal-on-behalf-of-email", + "x-nmp-principal-on-behalf-of-groups", + "x-nmp-scopes", +} +SEALED_EXTERNAL_PRINCIPAL_HEADERS = AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS + + +def _exact_header_patterns(patterns: list[dict]) -> set[str]: + return {pattern["exact"] for pattern in patterns if "exact" in pattern} + def test_provider_name_discovery_does_not_require_grant_secret(monkeypatch): monkeypatch.delenv("AUTHENTIK_WORKLOAD_IDENTITY_PASSWORD", raising=False) @@ -88,6 +103,17 @@ def test_authentik_compose_defaults_support_direct_docker_compose_start(): assert "./.generated/blueprints" not in compose_text +def test_authentik_compose_platform_config_uses_global_docker_workload_identity_switch_only(): + config = yaml.safe_load(Path("contrib/auth/authentik/config/platform-compose-authentik.yaml").read_text()) + workload_executor = next( + executor + for executor in config["jobs"]["executors"] + if executor["backend"] == "docker" and executor["provider"] == "cpu" and executor["profile"] == "workload" + ) + + assert "workload_identity" not in workload_executor["config"] + + def test_authentik_compose_uses_liveness_for_container_health_and_routes_status_through_gateway(): compose = yaml.safe_load(Path("contrib/auth/authentik/compose/docker-compose.yml").read_text()) envoy = yaml.safe_load(Path("contrib/auth/authentik/gateway/envoy.yaml").read_text()) @@ -147,6 +173,8 @@ def test_authentik_compose_uses_liveness_for_container_health_and_routes_status_ ) in lua_code assert 'headers:remove("x-nmp-authorized")' in lua_code assert 'headers:remove("x-nmp-scopes")' in lua_code + for header in SEALED_EXTERNAL_PRINCIPAL_HEADERS: + assert f'headers:remove("{header}")' in lua_code assert "claim_to_headers" not in yaml.safe_dump(http_manager) assert all( @@ -169,10 +197,7 @@ def test_authentik_compose_uses_liveness_for_container_health_and_routes_status_ "timeout": "5s", } allowed_headers = ext_authz["http_service"]["authorization_response"]["allowed_upstream_headers"]["patterns"] - assert {"exact": "x-nmp-principal-id"} in allowed_headers - assert {"exact": "x-nmp-principal-email"} in allowed_headers - assert {"exact": "x-nmp-principal-groups"} in allowed_headers - assert {"exact": "x-nmp-scopes"} in allowed_headers + assert AUTH_CALLOUT_RESPONSE_PRINCIPAL_HEADERS.issubset(_exact_header_patterns(allowed_headers)) protected_api_route = next(route for route in routes if route["match"] == {"prefix": "/apis/"}) assert "typed_per_filter_config" not in protected_api_route @@ -228,9 +253,7 @@ def test_authentik_compose_uses_https_gateway_for_workloads(): for executor in config["jobs"]["executors"] if executor["provider"] == "cpu" and executor["profile"] == "workload" ) - assert workload_executor["config"]["workload_identity"]["token_endpoint"] == ( - "https://nemo-gateway:8080/application/o/token/" - ) + assert "workload_identity" not in workload_executor["config"] assert set(nemo["networks"]) == {"nemo-internal"} assert "nemo-direct" not in yaml.safe_dump(nemo) assert nemo["depends_on"]["gateway-tls-init"]["condition"] == "service_completed_successfully" diff --git a/tests/auth_idp/static/test_provider_manifest.py b/tests/auth_idp/static/test_provider_manifest.py index 059ee8d605..56480c67ac 100644 --- a/tests/auth_idp/static/test_provider_manifest.py +++ b/tests/auth_idp/static/test_provider_manifest.py @@ -112,7 +112,7 @@ def test_authentik_manifest_compose_and_kubernetes_runtime_capabilities_stay_in_ compose = runtimes["authentik-compose"] kubernetes = runtimes["authentik-kubernetes"] - assert compose - {"docker_subject_token_refresh"} == kubernetes - {"kubernetes_token_review"} + assert compose == kubernetes - {"kubernetes_token_review"} assert "interactive_user_token" not in compose assert "interactive_user_token" not in kubernetes assert "workload_provider_token" in compose