diff --git a/src/control-plane-services/event-ledger/cmd/api/service/v3.go b/src/control-plane-services/event-ledger/cmd/api/service/v3.go index 40c1d29748..a7475f960e 100644 --- a/src/control-plane-services/event-ledger/cmd/api/service/v3.go +++ b/src/control-plane-services/event-ledger/cmd/api/service/v3.go @@ -297,7 +297,7 @@ func (s *Server) processOTLPEvents(traceCtx context.Context, req *collectorlogsv for _, rl := range req.ResourceLogs { for _, sl := range rl.ScopeLogs { for _, lr := range sl.LogRecords { - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(traceCtx, lr) if err != nil { logger.WarnContext(traceCtx, "Skipping event", zap.Error(err)) result.FailureCount++ @@ -442,6 +442,25 @@ func eventContextToCanonical(eventContext ContextV3) (string, error) { return strings.Join(parts, ","), nil } +// bindNVCAClusterID makes an SIS-verified NVCA cluster identity authoritative +// over whatever cluster_id a request payload claims: a missing payload value +// is populated from it, and a mismatching one is rejected outright, so a PSAT +// valid for one cluster cannot write events attributed to another. Requests +// with no NVCA identity in context (SIS/Spot JWT callers) are unaffected. +func bindNVCAClusterID(ctx context.Context, payloadClusterID string) (string, error) { + identity, ok := middleware.NVCAIdentityFromContext(ctx) + if !ok { + return payloadClusterID, nil + } + if payloadClusterID == "" { + return identity.ClusterID, nil + } + if payloadClusterID != identity.ClusterID { + return "", fmt.Errorf("cluster_id %q does not match the authorized cluster", payloadClusterID) + } + return payloadClusterID, nil +} + // extractK8sEvent converts an OTLP log record to EventV3 // Expected OTLP attributes: // - event_name (string): Event type @@ -450,7 +469,7 @@ func eventContextToCanonical(eventContext ContextV3) (string, error) { // - Context fields (optional): instance_id, deployment_id, gpu_specification_id, cluster_id // - resource_id (optional): generic unique identifier for events that have no // other distinguishing context field (e.g. an ICMSRequest keyed by its request id). -func extractK8sEvent(lr *logsv1.LogRecord) (*EventV3, error) { +func extractK8sEvent(ctx context.Context, lr *logsv1.LogRecord) (*EventV3, error) { // Step 1: Convert OTLP protobuf attributes to map attrs := make(map[string]any) for _, attr := range lr.Attributes { @@ -464,11 +483,15 @@ func extractK8sEvent(lr *logsv1.LogRecord) (*EventV3, error) { } // Step 3: Convert wire format to internal context representation + clusterID, err := bindNVCAClusterID(ctx, wireFormat.ClusterID) + if err != nil { + return nil, err + } contextV3 := ContextV3{ InstanceID: wireFormat.InstanceID, DeploymentID: wireFormat.DeploymentID, GPUSpecificationID: wireFormat.GPUSpecificationID, - ClusterID: wireFormat.ClusterID, + ClusterID: clusterID, ResourceID: wireFormat.ResourceID, } @@ -522,7 +545,7 @@ func extractK8sEvent(lr *logsv1.LogRecord) (*EventV3, error) { // - namespace (required) // - Context fields (optional, camelCase): instanceId, deploymentId, gpuSpecificationId, clusterId // Note: CloudEvents spec forbids underscores in extension names, so we use camelCase -func extractCloudEvent(ce *cloudevents.Event) (*EventV3, error) { +func extractCloudEvent(ctx context.Context, ce *cloudevents.Event) (*EventV3, error) { // Validate required CloudEvents fields per spec (using CloudEvents field names in errors) if strings.TrimSpace(ce.ID()) == "" { return nil, errors.New("missing required field: id") @@ -541,11 +564,15 @@ func extractCloudEvent(ce *cloudevents.Event) (*EventV3, error) { } // Convert wire format to internal context representation + clusterID, err := bindNVCAClusterID(ctx, wireFormat.ClusterID) + if err != nil { + return nil, err + } contextV3 := ContextV3{ InstanceID: wireFormat.InstanceID, DeploymentID: wireFormat.DeploymentID, GPUSpecificationID: wireFormat.GPUSpecificationID, - ClusterID: wireFormat.ClusterID, + ClusterID: clusterID, ResourceID: wireFormat.ResourceID, } @@ -595,7 +622,7 @@ func (s *Server) processCloudEvents(traceCtx context.Context, cloudEvents []*clo continue } - event, err := extractCloudEvent(cloudEvent) + event, err := extractCloudEvent(traceCtx, cloudEvent) if err != nil { logger.WarnContext(traceCtx, "Skipping event", zap.Error(err)) result.FailureCount++ diff --git a/src/control-plane-services/event-ledger/cmd/api/service/v3_test.go b/src/control-plane-services/event-ledger/cmd/api/service/v3_test.go index 61f030e1a8..7d2ea48f70 100644 --- a/src/control-plane-services/event-ledger/cmd/api/service/v3_test.go +++ b/src/control-plane-services/event-ledger/cmd/api/service/v3_test.go @@ -40,6 +40,8 @@ import ( commonv1 "go.opentelemetry.io/proto/otlp/common/v1" logsv1 "go.opentelemetry.io/proto/otlp/logs/v1" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/middleware" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/logging" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/common/core/types" @@ -611,7 +613,7 @@ func TestExtractK8sEvent(t *testing.T) { "extra_field": "extra_value", }) - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(context.Background(), lr) require.NoError(t, err) // Check struct fields @@ -680,13 +682,57 @@ func TestExtractK8sEvent_ResourceID(t *testing.T) { "resource_id": "icms-abc", }) - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(context.Background(), lr) require.NoError(t, err) // resource_id participates in the context (sorted last), keeping the row unique. assert.Equal(t, "cluster_id=clus-1,resource_id=icms-abc", event.Context) } +// TestExtractK8sEvent_NVCAClusterBinding verifies that an SIS-verified NVCA +// cluster identity is authoritative over the payload: a matching cluster_id +// is accepted, a missing one is populated, and a mismatched one is rejected +// so a PSAT valid for one cluster cannot write events for another. +func TestExtractK8sEvent_NVCAClusterBinding(t *testing.T) { + nvcaCtx := middleware.WithNVCAIdentity(context.Background(), middleware.NVCAIdentity{ + Subject: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }) + + t.Run("matching payload cluster_id is accepted", func(t *testing.T) { + lr := createOTLPLogRecord("pod.ready", "tenant-123", "nvca", "pod-1", map[string]string{ + "cluster_id": "cluster-a", + }) + event, err := extractK8sEvent(nvcaCtx, lr) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) + + t.Run("missing payload cluster_id is populated from the verified identity", func(t *testing.T) { + lr := createOTLPLogRecord("pod.ready", "tenant-123", "nvca", "pod-1", nil) + event, err := extractK8sEvent(nvcaCtx, lr) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) + + t.Run("mismatched payload cluster_id is rejected", func(t *testing.T) { + lr := createOTLPLogRecord("pod.ready", "tenant-123", "nvca", "pod-1", map[string]string{ + "cluster_id": "cluster-b", + }) + _, err := extractK8sEvent(nvcaCtx, lr) + assert.Error(t, err) + }) + + t.Run("no NVCA identity leaves the payload cluster_id untouched", func(t *testing.T) { + lr := createOTLPLogRecord("pod.ready", "tenant-123", "sis", "pod-1", map[string]string{ + "cluster_id": "cluster-a", + }) + event, err := extractK8sEvent(context.Background(), lr) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) +} + // TestExtractK8sEvent_DistinctResourceIDsDoNotCollide verifies two resources with // the same non-resource context but different resource_id produce distinct contexts. func TestExtractK8sEvent_DistinctResourceIDsDoNotCollide(t *testing.T) { @@ -695,7 +741,7 @@ func TestExtractK8sEvent_DistinctResourceIDsDoNotCollide(t *testing.T) { "cluster_id": "clus-1", "resource_id": resourceID, }) - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(context.Background(), lr) require.NoError(t, err) return event.Context } @@ -712,7 +758,7 @@ func TestExtractK8sEvent_PodKeepsUnmappedAttrsInDetails(t *testing.T) { "icms_request_id": "icms-xyz", }) - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(context.Background(), lr) require.NoError(t, err) // Pod context stays the original shape and excludes the unmapped attribute. @@ -733,7 +779,7 @@ func TestExtractCloudEvent_SourceRequired(t *testing.T) { ce.SetSource("") // Empty source ce.SetExtension("namespace", "test-namespace") - _, err := extractCloudEvent(&ce) + _, err := extractCloudEvent(context.Background(), &ce) require.Error(t, err) assert.Contains(t, err.Error(), "missing required field: source") } @@ -746,7 +792,7 @@ func TestExtractCloudEvent_TypeRequired(t *testing.T) { ce.SetSource("/test") ce.SetExtension("namespace", "test-namespace") - _, err := extractCloudEvent(&ce) + _, err := extractCloudEvent(context.Background(), &ce) require.Error(t, err) assert.Contains(t, err.Error(), "missing required field: type") } @@ -759,7 +805,7 @@ func TestExtractCloudEvent_IdRequired(t *testing.T) { ce.SetSource("/test") ce.SetExtension("namespace", "test-namespace") - _, err := extractCloudEvent(&ce) + _, err := extractCloudEvent(context.Background(), &ce) require.Error(t, err) assert.Contains(t, err.Error(), "missing required field: id") } @@ -775,11 +821,61 @@ func TestExtractCloudEvent_ResourceID(t *testing.T) { ce.SetExtension("clusterId", "clus-1") ce.SetExtension("resourceId", "icms-1") - event, err := extractCloudEvent(&ce) + event, err := extractCloudEvent(context.Background(), &ce) require.NoError(t, err) assert.Equal(t, "cluster_id=clus-1,resource_id=icms-1", event.Context) } +// TestExtractCloudEvent_NVCAClusterBinding mirrors +// TestExtractK8sEvent_NVCAClusterBinding: bindNVCAClusterID is wired into +// both extractK8sEvent and extractCloudEvent, so both need the same +// match/populate/reject/no-identity coverage. +func TestExtractCloudEvent_NVCAClusterBinding(t *testing.T) { + nvcaCtx := middleware.WithNVCAIdentity(context.Background(), middleware.NVCAIdentity{ + Subject: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }) + + newEvent := func(clusterID string) cloudevents.Event { + ce := cloudevents.NewEvent() + ce.SetID("test-id") + ce.SetType("test.event") + ce.SetSource("/test") + ce.SetExtension("namespace", "tenant-123") + if clusterID != "" { + ce.SetExtension("clusterId", clusterID) + } + return ce + } + + t.Run("matching payload cluster_id is accepted", func(t *testing.T) { + ce := newEvent("cluster-a") + event, err := extractCloudEvent(nvcaCtx, &ce) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) + + t.Run("missing payload cluster_id is populated from the verified identity", func(t *testing.T) { + ce := newEvent("") + event, err := extractCloudEvent(nvcaCtx, &ce) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) + + t.Run("mismatched payload cluster_id is rejected", func(t *testing.T) { + ce := newEvent("cluster-b") + _, err := extractCloudEvent(nvcaCtx, &ce) + assert.Error(t, err) + }) + + t.Run("no NVCA identity leaves the payload cluster_id untouched", func(t *testing.T) { + ce := newEvent("cluster-a") + event, err := extractCloudEvent(context.Background(), &ce) + require.NoError(t, err) + assert.Contains(t, event.Context, "cluster_id=cluster-a") + }) +} + // ====================== // CloudEvents Endpoint Validation Tests // ====================== diff --git a/src/control-plane-services/event-ledger/cmd/api/startup/BUILD.bazel b/src/control-plane-services/event-ledger/cmd/api/startup/BUILD.bazel index 52ae2bf0a8..0d9cb403ed 100644 --- a/src/control-plane-services/event-ledger/cmd/api/startup/BUILD.bazel +++ b/src/control-plane-services/event-ledger/cmd/api/startup/BUILD.bazel @@ -16,6 +16,7 @@ go_library( "//src/control-plane-services/event-ledger/internal/data_access", "//src/control-plane-services/event-ledger/internal/interfaces", "//src/control-plane-services/event-ledger/internal/middleware", + "//src/control-plane-services/event-ledger/internal/nvca", "//src/control-plane-services/event-ledger/internal/observability/logging", "//src/control-plane-services/event-ledger/internal/observability/tracing", "//src/control-plane-services/event-ledger/internal/policy", diff --git a/src/control-plane-services/event-ledger/cmd/api/startup/run_service.go b/src/control-plane-services/event-ledger/cmd/api/startup/run_service.go index ac8b37bd74..f79d8ff487 100644 --- a/src/control-plane-services/event-ledger/cmd/api/startup/run_service.go +++ b/src/control-plane-services/event-ledger/cmd/api/startup/run_service.go @@ -46,6 +46,7 @@ import ( "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/data_access" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/interfaces" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/middleware" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/nvca" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/logging" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/tracing" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/policy" @@ -297,6 +298,23 @@ func runService(cfg config.Config) error { jwtOpts = &opts } + var introspector nvca.Introspector + if cfg.Auth.Introspection.Enabled { + introspectionCfg := cfg.Auth.Introspection.WithDefaults() + logger.Warn("nvca psat introspection enabled", zap.String("url", introspectionCfg.URL)) + + introspectionClient, err := nvca.NewClient( + introspectionCfg.URL, + time.Duration(introspectionCfg.TimeoutSeconds)*time.Second, + time.Duration(introspectionCfg.CacheTTLSeconds)*time.Second, + ) + if err != nil { + logger.Error("failed to create nvca introspection client", zap.Error(err)) + return fmt.Errorf("failed to create nvca introspection client: %w", err) + } + introspector = introspectionClient + } + requireLocalScopeCheck = cfg.SelfManaged authRouter.Use(middleware.NewAuthMiddleware( @@ -305,6 +323,7 @@ func runService(cfg config.Config) error { jwtOpts, jwkCache, cfg.SelfManaged, + introspector, logger, )) default: diff --git a/src/control-plane-services/event-ledger/internal/config/auth_config_test.go b/src/control-plane-services/event-ledger/internal/config/auth_config_test.go index 792772379b..0deb483f43 100644 --- a/src/control-plane-services/event-ledger/internal/config/auth_config_test.go +++ b/src/control-plane-services/event-ledger/internal/config/auth_config_test.go @@ -109,10 +109,10 @@ func TestValidateAuthConfig_JWTProvider(t *testing.T) { func TestValidateAuthConfig_PolicyProvider(t *testing.T) { tests := []struct { - name string - cfg AuthConfig + name string + cfg AuthConfig selfManaged bool - expectedErr error + expectedErr error }{ { name: "valid policy config", @@ -253,7 +253,7 @@ func TestValidateAuthConfig_PolicyProvider(t *testing.T) { }, { // In self-managed mode, OAuth2 fields are not required. - name: "valid config in self-managed mode - oauth2 fields not required", + name: "valid config in self-managed mode - oauth2 fields not required", selfManaged: true, cfg: AuthConfig{ Enabled: true, @@ -269,7 +269,7 @@ func TestValidateAuthConfig_PolicyProvider(t *testing.T) { }, { // Even in self-managed mode, always-required fields are still checked. - name: "self-managed mode does not bypass namespace check", + name: "self-managed mode does not bypass namespace check", selfManaged: true, cfg: AuthConfig{ Enabled: true, @@ -370,3 +370,48 @@ func TestValidateEndpointAuthConfig(t *testing.T) { }) } } + +func TestValidateAuthConfig_PolicyProviderIntrospection(t *testing.T) { + baseCfg := func() AuthConfig { + return AuthConfig{ + Enabled: true, + Provider: "policy", + JWKSetUrl: "https://example.com/.well-known/jwks.json", + Policy: PolicyConfig{ + PolicyEvaluatorAddr: "https://pdp.example.com", + Namespace: "test", + PolicyFQDN: "test.policy", + }, + } + } + + t.Run("introspection disabled requires no url", func(t *testing.T) { + cfg := baseCfg() + assert.NoError(t, ValidateAuthConfig(cfg, true)) + }) + + t.Run("introspection enabled without url fails startup", func(t *testing.T) { + cfg := baseCfg() + cfg.Introspection.Enabled = true + err := ValidateAuthConfig(cfg, true) + require.Error(t, err) + assert.Equal(t, ErrMissingIntrospectionURL, err) + }) + + t.Run("introspection enabled with url is valid", func(t *testing.T) { + cfg := baseCfg() + cfg.Introspection.Enabled = true + cfg.Introspection.URL = "https://sis.example.com/v1/nvca/tokens/introspect" + assert.NoError(t, ValidateAuthConfig(cfg, true)) + }) +} + +func TestIntrospectionConfigWithDefaults(t *testing.T) { + cfg := IntrospectionConfig{}.WithDefaults() + assert.Equal(t, 10, cfg.TimeoutSeconds) + assert.Equal(t, 300, cfg.CacheTTLSeconds) + + cfg = IntrospectionConfig{TimeoutSeconds: 5, CacheTTLSeconds: 60}.WithDefaults() + assert.Equal(t, 5, cfg.TimeoutSeconds) + assert.Equal(t, 60, cfg.CacheTTLSeconds) +} diff --git a/src/control-plane-services/event-ledger/internal/config/config.go b/src/control-plane-services/event-ledger/internal/config/config.go index 90a794009f..a95f0ecc0a 100644 --- a/src/control-plane-services/event-ledger/internal/config/config.go +++ b/src/control-plane-services/event-ledger/internal/config/config.go @@ -39,6 +39,7 @@ var ( ErrMissingPolicyNamespace = errors.New("policy: namespace is required") ErrMissingPolicyFQDN = errors.New("policy: policy-fqdn is required") ErrInvalidPolicyCredsRefreshInterval = errors.New("policy: creds-refresh-interval must be greater than 0") + ErrMissingIntrospectionURL = errors.New("auth: introspection.url is required when introspection is enabled") ) // Top-level config @@ -69,13 +70,38 @@ type PublisherConfig struct { type AuthConfig struct { Enabled bool - Provider string `mapstructure:"provider"` - JWKSetUrl string `mapstructure:"jwk-set-url"` - Issuer string `mapstructure:"issuer"` - Audience string `mapstructure:"audience"` - TenantClaim string `mapstructure:"tenant-claim"` - CacheRefreshInterval int `mapstructure:"cache-refresh-interval"` - Policy PolicyConfig `mapstructure:"policy"` + Provider string `mapstructure:"provider"` + JWKSetUrl string `mapstructure:"jwk-set-url"` + Issuer string `mapstructure:"issuer"` + Audience string `mapstructure:"audience"` + TenantClaim string `mapstructure:"tenant-claim"` + CacheRefreshInterval int `mapstructure:"cache-refresh-interval"` + Policy PolicyConfig `mapstructure:"policy"` + Introspection IntrospectionConfig `mapstructure:"introspection"` +} + +// IntrospectionConfig configures the SIS call used to verify NVCA's PSAT for +// callers that do not hold an OpenBao-issued JWT. It is deliberately separate +// from stack-level deployment gating (addons.eventLedger.enabled): a stack +// can enable the Event Ledger release without wiring introspection, and that +// must fail startup rather than silently accept unverified NVCA callers. +type IntrospectionConfig struct { + Enabled bool `mapstructure:"enabled"` + URL string `mapstructure:"url"` + TimeoutSeconds int `mapstructure:"timeout-seconds"` + CacheTTLSeconds int `mapstructure:"cache-ttl-seconds"` +} + +// WithDefaults fills in the timeout and cache TTL the design calls for: a +// 10-second SIS call timeout and a 5-minute introspection cache. +func (i IntrospectionConfig) WithDefaults() IntrospectionConfig { + if i.TimeoutSeconds <= 0 { + i.TimeoutSeconds = 10 + } + if i.CacheTTLSeconds <= 0 { + i.CacheTTLSeconds = 300 + } + return i } type PolicyConfig struct { @@ -146,6 +172,9 @@ func ValidateAuthConfig(cfg AuthConfig, selfManaged bool) error { return ErrInvalidPolicyCredsRefreshInterval } } + if cfg.Introspection.Enabled && cfg.Introspection.URL == "" { + return ErrMissingIntrospectionURL + } case "": return ErrMissingAuthProvider default: diff --git a/src/control-plane-services/event-ledger/internal/middleware/BUILD.bazel b/src/control-plane-services/event-ledger/internal/middleware/BUILD.bazel index 5e9fef8921..a83726b131 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/BUILD.bazel +++ b/src/control-plane-services/event-ledger/internal/middleware/BUILD.bazel @@ -9,6 +9,7 @@ go_library( "http_client.go", "jwt.go", "metrics.go", + "nvca_introspect.go", "policy.go", ], importpath = "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/middleware", @@ -16,6 +17,7 @@ go_library( deps = [ "//src/control-plane-services/event-ledger/cmd/api/error", "//src/control-plane-services/event-ledger/internal/config", + "//src/control-plane-services/event-ledger/internal/nvca", "//src/control-plane-services/event-ledger/internal/observability/logging", "//src/control-plane-services/event-ledger/internal/policy", "//src/control-plane-services/event-ledger/pkg/constants", @@ -54,6 +56,7 @@ go_test( embed = [":middleware"], deps = [ "//src/control-plane-services/event-ledger/internal/config", + "//src/control-plane-services/event-ledger/internal/nvca", "//src/control-plane-services/event-ledger/internal/observability/logging", "//src/control-plane-services/event-ledger/internal/policy", "//src/control-plane-services/event-ledger/pkg/testutils", diff --git a/src/control-plane-services/event-ledger/internal/middleware/jwt.go b/src/control-plane-services/event-ledger/internal/middleware/jwt.go index 30be7be9e7..1377a94861 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/jwt.go +++ b/src/control-plane-services/event-ledger/internal/middleware/jwt.go @@ -304,6 +304,13 @@ func requireScopes(requiredScopes Scopes, scopeRequirement ScopeRequirement) fun next.ServeHTTP(w, r) return } + // An SIS-introspected NVCA identity carries no scopes either. + // Only trust it on the write routes it was scoped for, never + // as a stand-in for an arbitrary required scope. + if _, ok := NVCAIdentityFromContext(parentCtx); ok && requiredScopes == WriteScopes { + next.ServeHTTP(w, r) + return + } logger.WarnContext(traceCtx, ErrMissingClaims) status := http.StatusUnauthorized // http.Error(w, ErrMissingClaims, status) @@ -450,7 +457,11 @@ func MaybeRequirePathTenant(enabled bool) mux.MiddlewareFunc { } } -func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.ResponseWriter, r *http.Request) (context.Context, error) { +// processJWTToken parses and validates a JWT from the request's Authorization +// header. When writeResponse is false, it returns the error without writing +// an HTTP response, so a caller can fall back to another verification path +// (e.g. SIS introspection) before deciding what to send the client. +func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.ResponseWriter, r *http.Request, writeResponse bool) (context.Context, error) { ctx := r.Context() // Safe guard against nil context if ctx == nil { @@ -467,12 +478,19 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response } errType := "Process JWT Token Error" + + respondUnauthorized := func(err error) { + if !writeResponse { + return + } + api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, http.StatusUnauthorized, err, w) + logging.LogHTTPResponse(traceCtx, ctxLogger, http.StatusUnauthorized, w.Header()) + } + if opts.JwksURL == "" { ctxLogger.WarnContext(traceCtx, ErrMissingJWKSURL) - status := http.StatusUnauthorized err := errors.New(ErrMissingJWKSURL) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } @@ -480,10 +498,8 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response authHeader := r.Header.Get("Authorization") if authHeader == "" { ctxLogger.WarnContext(traceCtx, ErrMissingAuthHeader) - status := http.StatusUnauthorized err := errors.New(ErrMissingAuthHeader) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } @@ -493,10 +509,8 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response tokenString := strings.TrimPrefix(authHeader, "Bearer ") if tokenString == authHeader { ctxLogger.WarnContext(traceCtx, ErrInvalidAuthFormat) - status := http.StatusUnauthorized err := errors.New(ErrInvalidAuthFormat) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } @@ -508,19 +522,15 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response token, err := parseJWTWithOptions(tokenString, claims, keyFunc, opts) if err != nil { ctxLogger.WarnContext(traceCtx, "invalid token", zap.Error(err)) - status := http.StatusUnauthorized err = fmt.Errorf("%s: %v", ErrInvalidToken, err) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } if !token.Valid { ctxLogger.WarnContext(traceCtx, ErrInvalidToken) - status := http.StatusUnauthorized err = errors.New(ErrInvalidToken) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } @@ -529,10 +539,8 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response authorizedTenants := tenantValuesFromClaim(claims[opts.TenantClaim]) if len(authorizedTenants) == 0 { ctxLogger.WarnContext(traceCtx, "missing or invalid tenant claim", zap.String("claim", opts.TenantClaim)) - status := http.StatusUnauthorized err = errors.New(ErrInvalidToken) - api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, status, err, w) - logging.LogHTTPResponse(traceCtx, ctxLogger, status, w.Header()) + respondUnauthorized(err) return nil, err } newCtx = context.WithValue(newCtx, tenantClaimsContextKey, authorizedTenants) @@ -548,7 +556,7 @@ func newParseJWTMiddleware(opts JWTParserOptions, jwkCache *jwk.Cache) mux.Middl return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - newContext, err := processJWTToken(opts, jwkCache, w, r) + newContext, err := processJWTToken(opts, jwkCache, w, r, true) if err != nil { return } diff --git a/src/control-plane-services/event-ledger/internal/middleware/nvca_introspect.go b/src/control-plane-services/event-ledger/internal/middleware/nvca_introspect.go new file mode 100644 index 0000000000..5e166e49e7 --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/middleware/nvca_introspect.go @@ -0,0 +1,128 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +SPDX-License-Identifier: Apache-2.0 + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package middleware + +import ( + "context" + "errors" + "net/http" + + "github.com/golang-jwt/jwt/v5" + "github.com/gorilla/mux" + "github.com/lestrrat-go/jwx/v2/jwk" + "go.uber.org/zap" + + api_error "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/cmd/api/error" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/nvca" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/logging" +) + +const nvcaIdentityContextKey contextKey = "nvca_identity" + +// NVCAIdentity is the caller identity established by introspecting NVCA's +// PSAT at SIS. ClusterID is the cluster SIS resolved for the token, and is +// authoritative over anything a request payload claims. +type NVCAIdentity struct { + Subject string + ClusterID string +} + +// WithNVCAIdentity attaches an NVCA identity to ctx. Production code reaches +// this only via a successful SIS introspection in +// newJWTWithPSATMiddleware; it is exported so other packages (and tests +// simulating an already-authenticated request) can do the same. +func WithNVCAIdentity(ctx context.Context, identity NVCAIdentity) context.Context { + return context.WithValue(ctx, nvcaIdentityContextKey, identity) +} + +// NVCAIdentityFromContext returns the NVCA identity established for this +// request by SIS introspection, if any. +func NVCAIdentityFromContext(ctx context.Context) (NVCAIdentity, bool) { + identity, ok := ctx.Value(nvcaIdentityContextKey).(NVCAIdentity) + return identity, ok +} + +// newJWTWithPSATMiddleware tries local OpenBao JWT verification first. When +// the bearer token is not an OpenBao token, it falls back to SIS +// introspection for NVCA's PSAT, following the same ordered chain ReVal +// uses. It never inspects the unverified aud claim to route between the two; +// each path performs full local or remote verification. +func newJWTWithPSATMiddleware(opts JWTParserOptions, jwkCache *jwk.Cache, introspector nvca.Introspector) mux.MiddlewareFunc { + if opts.Method == nil { + opts.Method = jwt.SigningMethodES256 + } + + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + traceCtx := r.Context() + logger := logging.GetLogger(traceCtx) + errType := "NVCA Introspection Error" + + // Capture the token before processJWTToken consumes and removes + // the Authorization header, so it is still available for the + // introspection fallback below. + token := bearerToken(r) + + newContext, err := processJWTToken(opts, jwkCache, w, r, false) + if err == nil { + next.ServeHTTP(w, r.WithContext(newContext)) + return + } + + if token == "" { + api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, http.StatusUnauthorized, errors.New(ErrMissingAuthHeader), w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusUnauthorized, w.Header()) + return + } + if len(token) > nvca.MaxTokenSize { + api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, http.StatusUnauthorized, nvca.ErrTokenTooLarge, w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusUnauthorized, w.Header()) + return + } + + result, ierr := introspector.Introspect(traceCtx, token) + if ierr != nil { + logger.ErrorContext(traceCtx, "nvca introspection call failed", zap.Error(ierr)) + api_error.GenerateErrorResponse(traceCtx, errType, "Service Unavailable", r.URL.Path, http.StatusServiceUnavailable, errors.New("introspection unavailable"), w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusServiceUnavailable, w.Header()) + return + } + if !result.Active { + logger.WarnContext(traceCtx, "nvca introspection returned inactive token") + api_error.GenerateErrorResponse(traceCtx, errType, "Unauthorized", r.URL.Path, http.StatusUnauthorized, errors.New(ErrInvalidToken), w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusUnauthorized, w.Header()) + return + } + if !nvca.IsValidNVCASubject(result.Sub) { + logger.WarnContext(traceCtx, "nvca introspection returned non-nvca subject") + api_error.GenerateErrorResponse(traceCtx, errType, "Forbidden", r.URL.Path, http.StatusForbidden, errors.New(ErrInsufficientPermissions), w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusForbidden, w.Header()) + return + } + if result.ClusterID == "" { + logger.WarnContext(traceCtx, "nvca introspection returned no cluster identity") + api_error.GenerateErrorResponse(traceCtx, errType, "Forbidden", r.URL.Path, http.StatusForbidden, errors.New(ErrInsufficientPermissions), w) + logging.LogHTTPResponse(traceCtx, logger, http.StatusForbidden, w.Header()) + return + } + + identity := NVCAIdentity{Subject: result.Sub, ClusterID: result.ClusterID} + next.ServeHTTP(w, r.WithContext(WithNVCAIdentity(traceCtx, identity))) + }) + } +} diff --git a/src/control-plane-services/event-ledger/internal/middleware/policy.go b/src/control-plane-services/event-ledger/internal/middleware/policy.go index 65d97a4f43..66c0ccec27 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/policy.go +++ b/src/control-plane-services/event-ledger/internal/middleware/policy.go @@ -24,6 +24,7 @@ import ( "strconv" "strings" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/nvca" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/logging" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/policy" "github.com/golang-jwt/jwt/v5" @@ -390,11 +391,20 @@ func chainMiddleware(first, second mux.MiddlewareFunc) mux.MiddlewareFunc { // Anything else is treated as an opaque API key and sent to policyClient // directly. policyClient's evaluation contract only accepts an API key, which // is why a JWT cannot be routed through it in self-managed deployments. -func NewAuthMiddleware(policyClient policy.Authorizer, serviceName string, jwtOpts *JWTParserOptions, jwkCache *jwk.Cache, selfManaged bool, logger *otelzap.Logger) mux.MiddlewareFunc { +// +// When introspector is non-nil, a JWT-shaped token that fails local OpenBao +// verification is retried against SIS's NVCA introspection endpoint before +// being rejected, following the same ordered chain ReVal uses for NVCA's PSAT. +func NewAuthMiddleware(policyClient policy.Authorizer, serviceName string, jwtOpts *JWTParserOptions, jwkCache *jwk.Cache, selfManaged bool, introspector nvca.Introspector, logger *otelzap.Logger) mux.MiddlewareFunc { apiKeyAuth := newPolicyMiddleware(policyClient, serviceName, logger) var jwtVerify mux.MiddlewareFunc - if jwtOpts != nil { + switch { + case jwtOpts == nil: + // no JWT verification configured + case introspector != nil: + jwtVerify = newJWTWithPSATMiddleware(*jwtOpts, jwkCache, introspector) + default: jwtVerify = NewParseJWTMiddleware(*jwtOpts, jwkCache) } if jwtVerify == nil { diff --git a/src/control-plane-services/event-ledger/internal/middleware/policy_test.go b/src/control-plane-services/event-ledger/internal/middleware/policy_test.go index 864af39c9e..004d23b6b3 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/policy_test.go +++ b/src/control-plane-services/event-ledger/internal/middleware/policy_test.go @@ -30,6 +30,7 @@ import ( "time" "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/config" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/nvca" policyclient "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/policy" pdpv1 "github.com/NVIDIA/nvcf/src/libraries/go/lib/pkg/nvkit/clients/pdp_types" "github.com/golang-jwt/jwt/v5" @@ -144,7 +145,7 @@ func signTokenWithClaims(t *testing.T, key *ecdsa.PrivateKey, claims jwt.MapClai func newAuthTestHandler(t *testing.T, jwtOpts *JWTParserOptions, jwkCache *jwk.Cache, client *stubPolicyClient, selfManaged bool, requiredScopes Scopes) http.Handler { t.Helper() logger := testLogger(t) - authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", jwtOpts, jwkCache, selfManaged, logger) + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", jwtOpts, jwkCache, selfManaged, nil, logger) scoped := MaybeRequireScopes(logger, true, requiredScopes, RequireAnyScopes) return authMiddleware(scoped(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) @@ -553,7 +554,7 @@ func TestNewAuthMiddlewareRejectsJWTShapedTokenWhenParsingFails(t *testing.T) { &config.HTTPClientConfig{}, ) - authMiddleware := NewAuthMiddleware(client, "test-service", &jwtOpts, jwk.NewCache(context.Background()), true, testLogger(t)) + authMiddleware := NewAuthMiddleware(client, "test-service", &jwtOpts, jwk.NewCache(context.Background()), true, nil, testLogger(t)) handlerCalled := false handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { handlerCalled = true @@ -571,7 +572,7 @@ func TestNewAuthMiddlewareRejectsJWTShapedTokenWhenParsingFails(t *testing.T) { } func TestNewAuthMiddlewareRejectsRequestsWithNilClientAndLogger(t *testing.T) { - authMiddleware := NewAuthMiddleware(nil, "test-service", nil, nil, true, nil) + authMiddleware := NewAuthMiddleware(nil, "test-service", nil, nil, true, nil, nil) handlerCalled := false handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -668,7 +669,7 @@ func TestSelfManagedJWTTenantClaimEnforced(t *testing.T) { t.Run(tc.name, func(t *testing.T) { client := &stubPolicyClient{result: allowResult(nil)} logger := testLogger(t) - authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, logger) + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, nil, logger) pathTenant := MaybeRequirePathTenant(true) scoped := MaybeRequireScopes(logger, true, ReadScopes, RequireAnyScopes) handler := authMiddleware(pathTenant(scoped(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -737,7 +738,7 @@ func TestManagedJWTStillDelegatesToPolicyDecisionPoint(t *testing.T) { jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) client := &stubPolicyClient{result: allowResult(nil)} - authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, false, testLogger(t)) + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, false, nil, testLogger(t)) var capturedCtx context.Context handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -756,3 +757,145 @@ func TestManagedJWTStillDelegatesToPolicyDecisionPoint(t *testing.T) { require.NotNil(t, capturedCtx) assert.Equal(t, "sis-api", GetClaims(capturedCtx)["sub"]) } + +type stubIntrospector struct { + result *nvca.IntrospectResult + err error + called bool +} + +func (s *stubIntrospector) Introspect(_ context.Context, _ string) (*nvca.IntrospectResult, error) { + s.called = true + return s.result, s.err +} + +// psatShapedToken is not a real JWT - it just has the three dot-separated, +// non-empty parts isJWTShapedToken looks for, so the dispatcher routes it to +// the JWT chain, where local OpenBao verification must fail before the SIS +// introspection fallback is tried. +const psatShapedToken = "psat.header.payload" + +func TestNVCAIntrospectionAuthorizesWriteRoute(t *testing.T) { + jwtOpts := NewJWTParserOptions("https://issuer.test/.well-known/jwks.json", nil, time.Minute, &config.HTTPClientConfig{}) + jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) + + introspector := &stubIntrospector{result: &nvca.IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }} + client := &stubPolicyClient{result: allowResult(nil)} + logger := testLogger(t) + + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, logger) + scoped := MaybeRequireScopes(logger, true, WriteScopes, RequireAnyScopes) + + var capturedCtx context.Context + handler := authMiddleware(scoped(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + capturedCtx = r.Context() + w.WriteHeader(http.StatusOK) + }))) + + req := httptest.NewRequest(http.MethodPost, "/v3/ledger/cloudevents", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + assert.True(t, introspector.called) + assert.False(t, client.called, "an NVCA PSAT must not be sent to the API-key evaluator") + + identity, ok := NVCAIdentityFromContext(capturedCtx) + require.True(t, ok) + assert.Equal(t, "cluster-a", identity.ClusterID) +} + +func TestNVCAIntrospectionDeniesReadRoute(t *testing.T) { + jwtOpts := NewJWTParserOptions("https://issuer.test/.well-known/jwks.json", nil, time.Minute, &config.HTTPClientConfig{}) + jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) + + introspector := &stubIntrospector{result: &nvca.IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }} + client := &stubPolicyClient{result: allowResult(nil)} + logger := testLogger(t) + + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, logger) + scoped := MaybeRequireScopes(logger, true, ReadScopes, RequireAnyScopes) + handler := authMiddleware(scoped(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }))) + + req := httptest.NewRequest(http.MethodGet, "/v3/ledger/namespace/nvcf/events", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusUnauthorized, recorder.Code, "an NVCA identity must not stand in for a read scope it was never issued") +} + +func TestNVCAIntrospectionRejectsInactiveToken(t *testing.T) { + jwtOpts := NewJWTParserOptions("https://issuer.test/.well-known/jwks.json", nil, time.Minute, &config.HTTPClientConfig{}) + jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) + + introspector := &stubIntrospector{result: &nvca.IntrospectResult{Active: false}} + client := &stubPolicyClient{result: allowResult(nil)} + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, testLogger(t)) + + handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodPost, "/v3/ledger/cloudevents", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusUnauthorized, recorder.Code) +} + +func TestNVCAIntrospectionRejectsNonNVCASubject(t *testing.T) { + jwtOpts := NewJWTParserOptions("https://issuer.test/.well-known/jwks.json", nil, time.Minute, &config.HTTPClientConfig{}) + jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) + + introspector := &stubIntrospector{result: &nvca.IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:some-other-workload", + ClusterID: "cluster-a", + }} + client := &stubPolicyClient{result: allowResult(nil)} + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, testLogger(t)) + + handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodPost, "/v3/ledger/cloudevents", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusForbidden, recorder.Code) +} + +func TestNVCAIntrospectionFailsClosedWhenSISUnavailable(t *testing.T) { + jwtOpts := NewJWTParserOptions("https://issuer.test/.well-known/jwks.json", nil, time.Minute, &config.HTTPClientConfig{}) + jwkCache := jwk.NewCache(context.Background(), jwk.WithRefreshWindow(time.Minute)) + + introspector := &stubIntrospector{err: errors.New("dial tcp: connection refused")} + client := &stubPolicyClient{result: allowResult(nil)} + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, testLogger(t)) + + handler := authMiddleware(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodPost, "/v3/ledger/cloudevents", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusServiceUnavailable, recorder.Code) +} diff --git a/src/control-plane-services/event-ledger/internal/nvca/BUILD.bazel b/src/control-plane-services/event-ledger/internal/nvca/BUILD.bazel new file mode 100644 index 0000000000..d89a94fa28 --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/nvca/BUILD.bazel @@ -0,0 +1,27 @@ +load("@rules_go//go:def.bzl", "go_library", "go_test") + +go_library( + name = "nvca", + srcs = ["introspect.go"], + importpath = "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/nvca", + visibility = ["//:__subpackages__"], + deps = [ + "@io_opentelemetry_go_contrib_instrumentation_net_http_otelhttp//:otelhttp", + ], +) + +alias( + name = "go_default_library", + actual = ":nvca", + visibility = ["//:__subpackages__"], +) + +go_test( + name = "nvca_test", + srcs = ["introspect_test.go"], + embed = [":nvca"], + deps = [ + "@com_github_stretchr_testify//assert", + "@com_github_stretchr_testify//require", + ], +) diff --git a/src/control-plane-services/event-ledger/internal/nvca/introspect.go b/src/control-plane-services/event-ledger/internal/nvca/introspect.go new file mode 100644 index 0000000000..f6dde2c55a --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/nvca/introspect.go @@ -0,0 +1,275 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +SPDX-License-Identifier: Apache-2.0 + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +// Package nvca introspects NVCA's Kubernetes projected service-account token +// (PSAT) at SIS, for callers that do not hold an OpenBao-issued JWT. It +// mirrors ReVal's ICMS introspection authorizer so Event Ledger and ReVal +// verify NVCA's identity the same way. +package nvca + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" +) + +// MaxTokenSize bounds how much bearer token material the introspection client +// will send. Self-managed cluster PSATs are well under 2 KiB; larger tokens +// are treated as abuse. +const MaxTokenSize = 2048 + +// ErrTokenTooLarge is returned when the bearer token exceeds MaxTokenSize. +var ErrTokenTooLarge = errors.New("bearer token exceeds maximum size of 2048 bytes") + +// maxCacheEntries bounds the introspection cache so high-cardinality token +// traffic can't grow it without limit. Realistic cardinality is one entry +// per distinct NVCA pod PSAT across all registered clusters, far under this. +const maxCacheEntries = 1024 + +const ( + // psatSubjectPrefix matches Kubernetes service-account token subjects. + psatSubjectPrefix = "system:serviceaccount:" + // expectedPSATServiceAccountName is the only ServiceAccount name accepted + // for PSAT subjects; the namespace is customer-configurable but the SA + // name is always `nvca`. + expectedPSATServiceAccountName = "nvca" + // spiffeSubjectPrefix matches SPIFFE SVID subjects. + spiffeSubjectPrefix = "spiffe://" + // spiffeNVCASegment must be the terminal path segment of an accepted + // SPIFFE SVID (matched with suffix so trailing-path attacks fail). + spiffeNVCASegment = "/nvca" +) + +// IsValidNVCASubject anchors identity to the NVCA workload rather than any +// service account that happens to run in the cluster: PSAT callers must be +// `system:serviceaccount::nvca`, SPIFFE callers must end with +// `/nvca`. Any other subject is rejected. +func IsValidNVCASubject(sub string) bool { + if strings.HasPrefix(sub, psatSubjectPrefix) { + parts := strings.SplitN(sub, ":", 4) + return len(parts) == 4 && parts[3] == expectedPSATServiceAccountName + } + if strings.HasPrefix(sub, spiffeSubjectPrefix) { + return strings.HasSuffix(sub, spiffeNVCASegment) + } + return false +} + +// IntrospectRequest is the body sent to SIS's introspection endpoint. +type IntrospectRequest struct { + Token string `json:"token"` +} + +// IntrospectResult is the response from SIS's NVCA token introspection +// endpoint (RFC 7662 shape, plus an NVCF-specific resolved cluster +// identifier that is not part of RFC 7662). +type IntrospectResult struct { + Active bool `json:"active"` + Sub string `json:"sub"` + ClusterID string `json:"cluster_id"` + Error string `json:"error,omitempty"` +} + +// Introspector verifies a bearer token by asking an external service whether +// it is currently valid. It is implemented by *Client and by test doubles. +type Introspector interface { + Introspect(ctx context.Context, token string) (*IntrospectResult, error) +} + +type cacheEntry struct { + result *IntrospectResult + expiresAt time.Time +} + +// Client calls SIS's POST /v1/nvca/tokens/introspect endpoint to verify +// NVCA's PSAT. Results are cached by a hash of the token (never the raw +// token) for cacheTTL, bounded by the token's own exp claim so a cached +// result never outlives the token it was computed for. +type Client struct { + introspectURL string + httpClient *http.Client + cacheTTL time.Duration + cacheMu sync.RWMutex + cache map[string]cacheEntry +} + +// NewClient builds an introspection client. introspectURL is required. A +// cacheTTL of 0 disables caching. +func NewClient(introspectURL string, timeout, cacheTTL time.Duration) (*Client, error) { + if strings.TrimSpace(introspectURL) == "" { + return nil, fmt.Errorf("nvca: introspect url is required") + } + if timeout <= 0 { + timeout = 10 * time.Second + } + return &Client{ + introspectURL: introspectURL, + httpClient: &http.Client{ + Timeout: timeout, + // Can't reuse middleware.GetSharedHTTPClient here: internal/middleware + // imports internal/nvca, so importing middleware back would cycle. + Transport: otelhttp.NewTransport(http.DefaultTransport, + otelhttp.WithSpanNameFormatter(func(_ string, _ *http.Request) string { + return "nvca.introspect" + }), + ), + }, + cacheTTL: cacheTTL, + cache: make(map[string]cacheEntry), + }, nil +} + +// Introspect implements Introspector. +func (c *Client) Introspect(ctx context.Context, token string) (*IntrospectResult, error) { + if len(token) > MaxTokenSize { + return nil, ErrTokenTooLarge + } + + key := cacheKey(token) + if cached, ok := c.cacheLookup(key); ok { + return cached, nil + } + + result, err := c.callIntrospect(ctx, token) + if err != nil { + return nil, err + } + + // Only a complete, subject-valid result is cached. A token that comes + // back inactive or with the wrong subject may pass moments later (clock + // skew, an nbf window), so it must be re-checked rather than pinned. A + // missing ClusterID must not be cached either: caching it would pin a + // transient, incomplete SIS response as a 403 for the full TTL even + // after SIS starts returning a complete one. + if result.Active && IsValidNVCASubject(result.Sub) && result.ClusterID != "" { + c.cacheStore(key, result, token) + } + + return result, nil +} + +func (c *Client) callIntrospect(ctx context.Context, token string) (*IntrospectResult, error) { + body, err := json.Marshal(IntrospectRequest{Token: token}) + if err != nil { + return nil, fmt.Errorf("marshal introspect request: %w", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.introspectURL, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("build introspect request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("call introspect endpoint: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read introspect response: %w", err) + } + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("introspect returned status %d", resp.StatusCode) + } + + var result IntrospectResult + if err := json.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("decode introspect response: %w", err) + } + return &result, nil +} + +// cacheKey returns a stable, non-reversible key for a token. Hashing keeps +// raw bearer material out of long-lived process memory. +func cacheKey(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + +// tokenExpiry decodes (without verifying) a JWT payload and returns the exp +// claim. Used only as an upper bound on cache TTL - the security boundary is +// the introspection result at SIS, not this local, unverified parse. +func tokenExpiry(token string) (time.Time, bool) { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return time.Time{}, false + } + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return time.Time{}, false + } + var claims struct { + Exp int64 `json:"exp"` + } + if err := json.Unmarshal(payload, &claims); err != nil || claims.Exp == 0 { + return time.Time{}, false + } + return time.Unix(claims.Exp, 0), true +} + +func (c *Client) cacheLookup(key string) (*IntrospectResult, bool) { + if c.cacheTTL <= 0 { + return nil, false + } + c.cacheMu.RLock() + entry, ok := c.cache[key] + c.cacheMu.RUnlock() + if !ok || !time.Now().Before(entry.expiresAt) { + return nil, false + } + return entry.result, true +} + +func (c *Client) cacheStore(key string, result *IntrospectResult, token string) { + if c.cacheTTL <= 0 { + return + } + ttl := c.cacheTTL + if exp, ok := tokenExpiry(token); ok { + if remaining := time.Until(exp); remaining < ttl { + ttl = remaining + } + } + if ttl <= 0 { + return + } + c.cacheMu.Lock() + if _, exists := c.cache[key]; !exists && len(c.cache) >= maxCacheEntries { + // At capacity: evict one entry rather than scanning the whole map. + // Go's range order is randomized, so this is an arbitrary eviction, + // not LRU - acceptable since entries are already TTL-bounded. + for k := range c.cache { + delete(c.cache, k) + break + } + } + c.cache[key] = cacheEntry{result: result, expiresAt: time.Now().Add(ttl)} + c.cacheMu.Unlock() +} diff --git a/src/control-plane-services/event-ledger/internal/nvca/introspect_test.go b/src/control-plane-services/event-ledger/internal/nvca/introspect_test.go new file mode 100644 index 0000000000..893f51e738 --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/nvca/introspect_test.go @@ -0,0 +1,181 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +SPDX-License-Identifier: Apache-2.0 + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package nvca + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsValidNVCASubject(t *testing.T) { + tests := []struct { + name string + sub string + want bool + }{ + {"psat with nvca service account", "system:serviceaccount:customer-ns:nvca", true}, + {"psat with other service account", "system:serviceaccount:customer-ns:default", false}, + {"psat missing service account segment", "system:serviceaccount:customer-ns", false}, + {"spiffe nvca svid", "spiffe://cluster.local/ns/customer-ns/sa/nvca", true}, + {"spiffe non-nvca svid", "spiffe://cluster.local/ns/customer-ns/sa/nvca-imposter", false}, + {"unrelated subject", "some-other-subject", false}, + {"empty subject", "", false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, IsValidNVCASubject(tc.sub)) + }) + } +} + +func signedTestToken(t *testing.T, exp time.Time) string { + t.Helper() + header := base64URLEncode(t, map[string]any{"alg": "none"}) + payload := base64URLEncode(t, map[string]any{"exp": exp.Unix()}) + return header + "." + payload + ".sig" +} + +func base64URLEncode(t *testing.T, v any) string { + t.Helper() + b, err := json.Marshal(v) + require.NoError(t, err) + return base64.RawURLEncoding.EncodeToString(b) +} + +func TestClientIntrospectRejectsOversizedToken(t *testing.T) { + client, err := NewClient("http://example.invalid/introspect", time.Second, 0) + require.NoError(t, err) + + oversized := strings.Repeat("a", MaxTokenSize+1) + _, err = client.Introspect(context.Background(), oversized) + assert.ErrorIs(t, err, ErrTokenTooLarge) +} + +func TestClientIntrospectCachesActiveValidSubjectOnly(t *testing.T) { + calls := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, time.Minute) + require.NoError(t, err) + + token := signedTestToken(t, time.Now().Add(time.Hour)) + + result, err := client.Introspect(context.Background(), token) + require.NoError(t, err) + assert.True(t, result.Active) + assert.Equal(t, "cluster-a", result.ClusterID) + assert.Equal(t, 1, calls) + + // Second call for the same token must hit the cache, not the server. + _, err = client.Introspect(context.Background(), token) + require.NoError(t, err) + assert.Equal(t, 1, calls, "expected the second introspection to be served from cache") +} + +func TestClientIntrospectDoesNotCacheInactiveResult(t *testing.T) { + calls := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(IntrospectResult{Active: false}) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, time.Minute) + require.NoError(t, err) + + token := signedTestToken(t, time.Now().Add(time.Hour)) + + _, err = client.Introspect(context.Background(), token) + require.NoError(t, err) + _, err = client.Introspect(context.Background(), token) + require.NoError(t, err) + assert.Equal(t, 2, calls, "an inactive result must never be cached, since clock skew can make the same token valid moments later") +} + +func TestClientIntrospectDoesNotCacheMissingClusterID(t *testing.T) { + calls := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:nvca", + // ClusterID intentionally omitted. + }) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, time.Minute) + require.NoError(t, err) + + token := signedTestToken(t, time.Now().Add(time.Hour)) + + _, err = client.Introspect(context.Background(), token) + require.NoError(t, err) + _, err = client.Introspect(context.Background(), token) + require.NoError(t, err) + assert.Equal(t, 2, calls, "a result missing ClusterID must not be cached, or a later complete SIS response stays hidden behind the stale cache entry until it expires") +} + +func TestClientIntrospectCacheEvictsAtCapacity(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req IntrospectRequest + _ = json.NewDecoder(r.Body).Decode(&req) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(IntrospectResult{ + Active: true, + Sub: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-" + req.Token, + }) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, time.Minute) + require.NoError(t, err) + + for i := 0; i <= maxCacheEntries; i++ { + token := signedTestToken(t, time.Now().Add(time.Hour)) + fmt.Sprintf(".%d", i) + _, err := client.Introspect(context.Background(), token) + require.NoError(t, err) + } + + client.cacheMu.RLock() + size := len(client.cache) + client.cacheMu.RUnlock() + assert.LessOrEqual(t, size, maxCacheEntries, "cache must never grow past maxCacheEntries") +}