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..305886f7cc 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 @@ -282,10 +282,11 @@ func (s *Server) PostK8sEventV3(w http.ResponseWriter, r *http.Request) { // EventProcessingResult holds the results of processing a batch of events type EventProcessingResult struct { - SuccessCount int - FailureCount int - ProcessedEvents []ProcessedEventSummary - LastError error + SuccessCount int + FailureCount int + ProcessedEvents []ProcessedEventSummary + LastError error + AuthorizationErr error } // processOTLPEvents extracts and stores K8s events from OTLP log records @@ -297,8 +298,13 @@ 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 { + if errors.Is(err, errClusterAuthorization) { + logger.WarnContext(traceCtx, "Aborting batch", zap.Error(err)) + result.AuthorizationErr = err + return result + } logger.WarnContext(traceCtx, "Skipping event", zap.Error(err)) result.FailureCount++ result.LastError = err @@ -442,6 +448,30 @@ func eventContextToCanonical(eventContext ContextV3) (string, error) { return strings.Join(parts, ","), nil } +// errClusterAuthorization is checked with errors.Is by the batch processors +// to abort on a bindNVCAClusterID rejection, instead of skipping it like an +// ordinary per-event validation error. +var errClusterAuthorization = errors.New("cluster_id does not match the authorized cluster") + +// 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("%w: %q", errClusterAuthorization, payloadClusterID) + } + return payloadClusterID, nil +} + // extractK8sEvent converts an OTLP log record to EventV3 // Expected OTLP attributes: // - event_name (string): Event type @@ -450,7 +480,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 +494,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 +556,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 +575,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,8 +633,13 @@ func (s *Server) processCloudEvents(traceCtx context.Context, cloudEvents []*clo continue } - event, err := extractCloudEvent(cloudEvent) + event, err := extractCloudEvent(traceCtx, cloudEvent) if err != nil { + if errors.Is(err, errClusterAuthorization) { + logger.WarnContext(traceCtx, "Aborting batch", zap.Error(err)) + result.AuthorizationErr = err + return result + } logger.WarnContext(traceCtx, "Skipping event", zap.Error(err)) result.FailureCount++ result.LastError = err @@ -722,6 +765,12 @@ func (s *Server) storeK8sEvent(traceCtx context.Context, event *EventV3) error { func (s *Server) sendEventResponse(w http.ResponseWriter, traceCtx context.Context, result EventProcessingResult) { logger := logging.GetLogger(traceCtx) + if result.AuthorizationErr != nil { + logger.WarnContext(traceCtx, "Rejecting batch", zap.Error(result.AuthorizationErr)) + sendProblemDetail(w, http.StatusForbidden, "Forbidden", result.AuthorizationErr.Error()) + return + } + // Prepare response based on results if result.SuccessCount == 0 && result.FailureCount == 0 { logger.ErrorContext(traceCtx, "No events found in OTLP request") 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..a1f51a4da8 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" @@ -605,13 +607,41 @@ func TestPostK8sEventV3_StatsFailureAllFail(t *testing.T) { assert.Contains(t, w.Body.String(), "stats down") } +func TestPostK8sEventV3_ClusterAuthorizationMismatchAbortsBatch(t *testing.T) { + mockDB := &mockDBHandlerV3{} + server := newServerWithMock(t, mockDB) + + req := newOTLPRequest( + createOTLPLogRecord("pod.ready", "ns", "src", "pod-1", nil), + createOTLPLogRecord("pod.ready", "ns", "src", "pod-2", map[string]string{"cluster_id": "cluster-b"}), + ) + body, err := proto.Marshal(req) + require.NoError(t, err) + + logger := testutils.InitTestLogger(t) + httpReq := httptest.NewRequest("POST", "/v3/ledger/k8s-events", bytes.NewReader(body)) + httpReq.Header.Set("Content-Type", "application/x-protobuf") + ctx := context.WithValue(httpReq.Context(), logging.LoggerKey, logging.NewTraceLogger(httpReq.Context(), logger)) + ctx = middleware.WithNVCAIdentity(ctx, middleware.NVCAIdentity{ + Subject: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }) + httpReq = httpReq.WithContext(ctx) + + w := httptest.NewRecorder() + server.PostK8sEventV3(w, httpReq) + + assert.Equal(t, http.StatusForbidden, w.Code) + assert.Empty(t, mockDB.storedEvents, "a cluster-authorization failure must abort the batch before any DB write") +} + // Test extractK8sEvent func TestExtractK8sEvent(t *testing.T) { lr := createOTLPLogRecord("pod.ready", "tenant-123", "kubernetes", "pod-456", map[string]string{ "extra_field": "extra_value", }) - event, err := extractK8sEvent(lr) + event, err := extractK8sEvent(context.Background(), lr) require.NoError(t, err) // Check struct fields @@ -680,13 +710,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 +769,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 +786,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 +807,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 +820,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 +833,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 +849,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 // ====================== @@ -843,6 +967,53 @@ func TestPostCloudEventV3_BatchMissingSpecversion(t *testing.T) { assert.Contains(t, w.Body.String(), "specversion") } +func TestPostCloudEventV3_ClusterAuthorizationMismatchAbortsBatch(t *testing.T) { + body := []byte(`[ + { + "specversion": "1.0", + "type": "pod.ready", + "source": "/test", + "id": "event-1", + "namespace": "ns" + }, + { + "specversion": "1.0", + "type": "pod.ready", + "source": "/test", + "id": "event-2", + "namespace": "ns", + "clusterid": "cluster-b" + } + ]`) + + mockDB := &mockDBHandlerV3{} + logger := testutils.InitTestLogger(t) + server := NewServer( + Connections{DbHandlerV2: mockDB}, + logger, + nil, + "test", + &config.HTTPClientConfig{}, + config.PaginationConfig{}, + config.StatsConfig{}, + ) + + req := httptest.NewRequest("POST", "/v3/ledger/cloudevents", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/cloudevents-batch+json") + ctx := context.WithValue(req.Context(), logging.LoggerKey, logging.NewTraceLogger(req.Context(), logger)) + ctx = middleware.WithNVCAIdentity(ctx, middleware.NVCAIdentity{ + Subject: "system:serviceaccount:customer-ns:nvca", + ClusterID: "cluster-a", + }) + req = req.WithContext(ctx) + + w := httptest.NewRecorder() + server.PostCloudEventV3(w, req) + + assert.Equal(t, http.StatusForbidden, w.Code) + assert.Empty(t, mockDB.storedEvents, "a cluster-authorization failure must abort the batch before any DB write") +} + func TestPostCloudEventV3_BatchRejectsNullEvent(t *testing.T) { w, _ := executeCloudEventsRequest(t, []byte(`[null]`), "application/cloudevents-batch+json") 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..cf46a5ccd3 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: @@ -447,12 +466,12 @@ func runService(cfg config.Config) error { } authRouter.Handle("/v3/ledger/k8s-events", - middleware.MaybeRequireScopes(logger, requireLocalScopeCheck, middleware.WriteScopes, middleware.RequireAnyScopes)(wrapper(http.HandlerFunc(server.PostK8sEventV3))), + middleware.MaybeRequireScopesAllowNVCA(logger, requireLocalScopeCheck, middleware.WriteScopes, middleware.RequireAnyScopes)(wrapper(http.HandlerFunc(server.PostK8sEventV3))), ).Methods("POST", "OPTIONS") // CloudEvents receiver endpoint authRouter.Handle("/v3/ledger/cloudevents", - middleware.MaybeRequireScopes(logger, requireLocalScopeCheck, middleware.WriteScopes, middleware.RequireAnyScopes)(http.HandlerFunc(server.PostCloudEventV3)), + middleware.MaybeRequireScopesAllowNVCA(logger, requireLocalScopeCheck, middleware.WriteScopes, middleware.RequireAnyScopes)(http.HandlerFunc(server.PostCloudEventV3)), ).Methods("POST", "OPTIONS") // V3 Stats endpoint - retrieve aggregated stats for a namespace diff --git a/src/control-plane-services/event-ledger/go.mod b/src/control-plane-services/event-ledger/go.mod index 4c441b732f..5d9f65960a 100644 --- a/src/control-plane-services/event-ledger/go.mod +++ b/src/control-plane-services/event-ledger/go.mod @@ -54,7 +54,7 @@ require ( ) require ( - github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260728185909-afca4ec2fb26 + github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260923212141-ea12b8777d46 github.com/beorn7/perks v1.0.1 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect diff --git a/src/control-plane-services/event-ledger/go.sum b/src/control-plane-services/event-ledger/go.sum index 821a73626c..f4032b0169 100644 --- a/src/control-plane-services/event-ledger/go.sum +++ b/src/control-plane-services/event-ledger/go.sum @@ -2,6 +2,10 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMT github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/Masterminds/squirrel v1.5.4 h1:uUcX/aBc8O7Fg9kaISIUsHXdKuqehiXAMQTYX8afzqM= github.com/Masterminds/squirrel v1.5.4/go.mod h1:NNaOrjSoIDfDA40n7sr2tPNZRfjzjA400rg+riTZj10= +github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260728185909-afca4ec2fb26 h1:P9OSmvy6MQPqfzRgaLl0XpvK5UYHAFtNnQQb+yk+6b0= +github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260728185909-afca4ec2fb26/go.mod h1:nj3yBW2weO0qzi1ML45tBqviLa2Y/KG9qCalxQuJVB8= +github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260923212141-ea12b8777d46 h1:Ln94dXu7bqznN+fLv0wuRZiSlX8FQEGUoPyFFrKP9DY= +github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260923212141-ea12b8777d46/go.mod h1:WUFjMVtVWK1PAw7WBXM0nWhYbtlE/4zyfjkj+lHUEtY= github.com/benbjohnson/clock v1.1.0/go.mod h1:J11/hYXuz8f4ySSvYwY0FKfm+ezbsZBKZxNJlLklBHA= github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= @@ -367,5 +371,3 @@ gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= -github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260728185909-afca4ec2fb26 h1:P9OSmvy6MQPqfzRgaLl0XpvK5UYHAFtNnQQb+yk+6b0= -github.com/NVIDIA/nvcf/src/libraries/go/lib v0.0.0-20260728185909-afca4ec2fb26/go.mod h1:nj3yBW2weO0qzi1ML45tBqviLa2Y/KG9qCalxQuJVB8= 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/cliargs.go b/src/control-plane-services/event-ledger/internal/config/cliargs.go index 0b15259b08..40bc294ff3 100644 --- a/src/control-plane-services/event-ledger/internal/config/cliargs.go +++ b/src/control-plane-services/event-ledger/internal/config/cliargs.go @@ -80,6 +80,10 @@ func (c *CliArgs) SetupAuth() { c.int64("auth.policy.creds-refresh-interval", 300, "Interval in seconds to periodically refresh credentials from file (recommended: 300)", true) c.string("auth.policy.subject-field", "subject", "Policy input field name for JWT subject", true) c.string("auth.policy.api-key-field", "apiKey", "Policy input field name for API key tokens", true) + c.bool("auth.introspection.enabled", false, "Enable SIS introspection of NVCA's PSAT for callers without an OpenBao JWT", true) + c.string("auth.introspection.url", "", "SIS token introspection endpoint URL", true) + c.int("auth.introspection.timeout-seconds", 10, "SIS introspection call timeout in seconds", true) + c.int("auth.introspection.cache-ttl-seconds", 300, "SIS introspection result cache TTL in seconds", true) } // SetupDatabase defines database providers arguments and configuration settings 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/config/config_test.go b/src/control-plane-services/event-ledger/internal/config/config_test.go index 8d050989ac..4972cc917c 100644 --- a/src/control-plane-services/event-ledger/internal/config/config_test.go +++ b/src/control-plane-services/event-ledger/internal/config/config_test.go @@ -236,6 +236,34 @@ func TestCloudEventsEnabledConfiguration_EventLedgerEnvPrefix(t *testing.T) { assert.False(t, cfg.Publisher.Cloudevents.Enabled) } +func TestAuthIntrospectionConfiguration_EventLedgerEnvPrefix(t *testing.T) { + t.Setenv("EVENT_LEDGER_AUTH_INTROSPECTION_ENABLED", "true") + t.Setenv("EVENT_LEDGER_AUTH_INTROSPECTION_URL", "http://api.sis.svc.cluster.local:8080/v1/nvca/tokens/introspect") + + v := viper.New() + v.AutomaticEnv() + v.SetEnvPrefix("EVENT_LEDGER") + v.SetEnvKeyReplacer(strings.NewReplacer("-", "_", ".", "_")) + + rootCmd := &cobra.Command{ + Use: "test", + RunE: func(cmd *cobra.Command, args []string) error { + return nil + }, + } + + logger := &otelzap.SugaredLogger{} + cliArgs := NewCliArgs(rootCmd, v, logger) + cliArgs.SetupAuth() + + require.NoError(t, rootCmd.Execute()) + + var cfg Config + require.NoError(t, v.Unmarshal(&cfg)) + assert.True(t, cfg.Auth.Introspection.Enabled) + assert.Equal(t, "http://api.sis.svc.cluster.local:8080/v1/nvca/tokens/introspect", cfg.Auth.Introspection.URL) +} + func TestAuthPolicyCredentialsRefreshIntervalDefault(t *testing.T) { v := viper.New() 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..c549a5f17b 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/jwt.go +++ b/src/control-plane-services/event-ledger/internal/middleware/jwt.go @@ -243,7 +243,20 @@ func MaybeRequireScopes(logger *otelzap.Logger, authEnabled bool, requiredScopes if !authEnabled { return next } - return requireScopes(requiredScopes, scopeRequirement)(next) + return requireScopes(requiredScopes, scopeRequirement, false)(next) + } +} + +// MaybeRequireScopesAllowNVCA is MaybeRequireScopes, but also trusts an +// SIS-introspected NVCA identity for write routes. Use only where the +// handler also binds the identity to a specific cluster (see +// bindNVCAClusterID in cmd/api/service/v3.go). +func MaybeRequireScopesAllowNVCA(logger *otelzap.Logger, authEnabled bool, requiredScopes Scopes, scopeRequirement ScopeRequirement) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + if !authEnabled { + return next + } + return requireScopes(requiredScopes, scopeRequirement, true)(next) } } @@ -287,7 +300,7 @@ func getScopesFromClaims(claims jwt.MapClaims) ([]string, bool) { return result, len(result) > 0 } -func requireScopes(requiredScopes Scopes, scopeRequirement ScopeRequirement) func(http.Handler) http.Handler { +func requireScopes(requiredScopes Scopes, scopeRequirement ScopeRequirement, allowNVCAIdentity bool) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { parentCtx := r.Context() @@ -304,6 +317,20 @@ func requireScopes(requiredScopes Scopes, scopeRequirement ScopeRequirement) fun next.ServeHTTP(w, r) return } + if _, nvcaOK := NVCAIdentityFromContext(parentCtx); nvcaOK { + // NVCA's PSAT carries no scopes; only trust it on routes + // that opt in and bind it to a specific cluster (see + // bindNVCAClusterID in cmd/api/service/v3.go). + if allowNVCAIdentity && requiredScopes == WriteScopes { + next.ServeHTTP(w, r) + return + } + logger.WarnContext(traceCtx, ErrInsufficientPermissions) + status := http.StatusForbidden + api_error.GenerateErrorResponse(traceCtx, errType, "Forbidden", r.URL.Path, status, errors.New(ErrInsufficientPermissions), w) + logging.LogHTTPResponse(traceCtx, logger, status, w.Header()) + return + } logger.WarnContext(traceCtx, ErrMissingClaims) status := http.StatusUnauthorized // http.Error(w, ErrMissingClaims, status) @@ -450,7 +477,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,23 +498,36 @@ 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()) + } + + // writeResponse=false means a caller (the PSAT fallback path) has another + // verification method to try before treating this as a real failure, so + // log quietly here and let that caller log once both paths are done. + logFailure := ctxLogger.WarnContext + if !writeResponse { + logFailure = ctxLogger.DebugContext + } + if opts.JwksURL == "" { - ctxLogger.WarnContext(traceCtx, ErrMissingJWKSURL) - status := http.StatusUnauthorized + logFailure(traceCtx, ErrMissingJWKSURL) 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 } // Get the token from the Authorization header authHeader := r.Header.Get("Authorization") if authHeader == "" { - ctxLogger.WarnContext(traceCtx, ErrMissingAuthHeader) - status := http.StatusUnauthorized + logFailure(traceCtx, ErrMissingAuthHeader) 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 } @@ -492,35 +536,29 @@ 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 + logFailure(traceCtx, ErrInvalidAuthFormat) 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 } // Get the key function for token verification - keyFunc := newJWKKeyFunc(traceCtx, opts, jwkCache, ctxLogger) + keyFunc := newJWKKeyFunc(traceCtx, opts, jwkCache, ctxLogger, !writeResponse) // Parse and validate the token claims := jwt.MapClaims{} token, err := parseJWTWithOptions(tokenString, claims, keyFunc, opts) if err != nil { - ctxLogger.WarnContext(traceCtx, "invalid token", zap.Error(err)) - status := http.StatusUnauthorized + logFailure(traceCtx, "invalid token", zap.Error(err)) 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 + logFailure(traceCtx, ErrInvalidToken) 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 } @@ -528,11 +566,9 @@ func processJWTToken(opts JWTParserOptions, jwkCache *jwk.Cache, w http.Response if opts.TenantClaim != "" { 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 + logFailure(traceCtx, "missing or invalid tenant claim", zap.String("claim", opts.TenantClaim)) 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 +584,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 } @@ -558,12 +594,17 @@ func newParseJWTMiddleware(opts JWTParserOptions, jwkCache *jwk.Cache) mux.Middl } } -func newJWKKeyFunc(ctx context.Context, opts JWTParserOptions, jwkCache *jwk.Cache, ctxLogger *logging.TraceLogger) jwt.Keyfunc { +func newJWKKeyFunc(ctx context.Context, opts JWTParserOptions, jwkCache *jwk.Cache, ctxLogger *logging.TraceLogger, quiet bool) jwt.Keyfunc { // Create a safe context if nil if ctx == nil { ctx = context.Background() } + logFailure := ctxLogger.ErrorContext + if quiet { + logFailure = ctxLogger.DebugContext + } + return func(token *jwt.Token) (interface{}, error) { if token == nil { return nil, errors.New("nil token provided to key function") @@ -571,26 +612,26 @@ func newJWKKeyFunc(ctx context.Context, opts JWTParserOptions, jwkCache *jwk.Cac keySet, err := fetchJwk(ctx, opts, jwkCache, ctxLogger) if err != nil { - ctxLogger.ErrorContext(ctx, "failed to fetch jwk", zap.Error(err)) + logFailure(ctx, "failed to fetch jwk", zap.Error(err)) return nil, errors.New(ErrFetchingJwk) } kid, err := extractKidFromTokenHeaders(token.Header) if err != nil { - ctxLogger.ErrorContext(ctx, "failed to extract kid from token headers", zap.Error(err)) + logFailure(ctx, "failed to extract kid from token headers", zap.Error(err)) // this is already assured to be of type nverror return nil, err } key, ok := keySet.LookupKeyID(kid) if !ok { - ctxLogger.ErrorContext(ctx, "jwk not found for kid", zap.String("kid", kid)) + logFailure(ctx, "jwk not found for kid", zap.String("kid", kid)) return nil, errors.New(ErrMissingJwk) } var cryptoKey interface{} if err := key.Raw(&cryptoKey); err != nil { - ctxLogger.ErrorContext(ctx, "failed to get raw crypto key", zap.Error(err)) + logFailure(ctx, "failed to get raw crypto key", zap.Error(err)) return nil, fmt.Errorf("failed to get raw crypto key: %w", err) } diff --git a/src/control-plane-services/event-ledger/internal/middleware/jwt_test.go b/src/control-plane-services/event-ledger/internal/middleware/jwt_test.go index 1c45d76810..b8f44865ba 100644 --- a/src/control-plane-services/event-ledger/internal/middleware/jwt_test.go +++ b/src/control-plane-services/event-ledger/internal/middleware/jwt_test.go @@ -209,11 +209,12 @@ func TestRequireScopes(t *testing.T) { logger := otelzap.New(zaptest.NewLogger(t)) tests := []struct { - name string - claims jwt.MapClaims - requiredScopes Scopes - scopeRequirement ScopeRequirement - expectedStatus int + name string + claims jwt.MapClaims + requiredScopes Scopes + scopeRequirement ScopeRequirement + allowNVCAIdentity bool + expectedStatus int }{ { name: "Has all required scopes", @@ -277,7 +278,7 @@ func TestRequireScopes(t *testing.T) { }) // Create middleware - middleware := requireScopes(tt.requiredScopes, tt.scopeRequirement) + middleware := requireScopes(tt.requiredScopes, tt.scopeRequirement, tt.allowNVCAIdentity) wrappedHandler := middleware(testHandler) // Create a test request @@ -305,6 +306,58 @@ func TestRequireScopes(t *testing.T) { } } +func TestRequireScopesNVCAIdentity(t *testing.T) { + logger := otelzap.New(zaptest.NewLogger(t)) + + tests := []struct { + name string + requiredScopes Scopes + allowNVCAIdentity bool + expectedStatus int + }{ + { + name: "NVCA identity on an allowed write route passes", + requiredScopes: WriteScopes, + allowNVCAIdentity: true, + expectedStatus: http.StatusOK, + }, + { + name: "NVCA identity on a write route that didn't opt in is forbidden, not unauthorized", + requiredScopes: WriteScopes, + allowNVCAIdentity: false, + expectedStatus: http.StatusForbidden, + }, + { + name: "NVCA identity on a non-write route is forbidden, not unauthorized", + requiredScopes: ReadScopes, + allowNVCAIdentity: true, + expectedStatus: http.StatusForbidden, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + testHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + middleware := requireScopes(tt.requiredScopes, RequireAnyScopes, tt.allowNVCAIdentity) + wrappedHandler := middleware(testHandler) + + req := httptest.NewRequest("POST", "/test", nil) + traceLogger := logging.NewTraceLogger(req.Context(), logger) + ctx := context.WithValue(req.Context(), logging.LoggerKey, traceLogger) + ctx = WithNVCAIdentity(ctx, NVCAIdentity{Subject: "system:serviceaccount:customer-ns:nvca", ClusterID: "cluster-a"}) + req = req.WithContext(ctx) + + rec := httptest.NewRecorder() + wrappedHandler.ServeHTTP(rec, req) + + assert.Equal(t, tt.expectedStatus, rec.Code) + }) + } +} + func TestMaybeRequireScopes(t *testing.T) { logger := otelzap.New(zaptest.NewLogger(t)) 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..070523f144 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,8 @@ 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" + "github.com/NVIDIA/nvcf/src/control-plane-services/event-ledger/internal/observability/logging" 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" @@ -38,7 +40,10 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/uptrace/opentelemetry-go-extra/otelzap" + "go.uber.org/zap" + "go.uber.org/zap/zapcore" "go.uber.org/zap/zaptest" + "go.uber.org/zap/zaptest/observer" "google.golang.org/protobuf/types/known/structpb" ) @@ -144,7 +149,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 +558,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 +576,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 +673,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 +742,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 +761,179 @@ 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 := MaybeRequireScopesAllowNVCA(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 TestNVCAIntrospectionValidPSATDoesNotLogAtWarnOrAbove(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)} + + observedCore, observedLogs := observer.New(zapcore.DebugLevel) + logger := otelzap.New(zap.New(observedCore)) + + authMiddleware := NewAuthMiddleware(client, "nv-cloud-functions", &jwtOpts, jwkCache, true, introspector, logger) + scoped := MaybeRequireScopesAllowNVCA(logger, true, WriteScopes, RequireAnyScopes) + handler := authMiddleware(scoped(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }))) + + req := httptest.NewRequest(http.MethodPost, "/v3/ledger/cloudevents", nil) + req.Header.Set("Authorization", "Bearer "+psatShapedToken) + ctx := context.WithValue(req.Context(), logging.LoggerKey, logging.NewTraceLogger(req.Context(), logger)) + req = req.WithContext(ctx) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + + for _, entry := range observedLogs.All() { + assert.Lessf(t, entry.Level, zapcore.WarnLevel, "expected local-verification-failure path for a valid PSAT must not log at warn or above, got %q at %s", entry.Message, entry.Level) + } +} + +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.StatusForbidden, recorder.Code, "an NVCA identity must not stand in for a read scope it was never issued, but it is authenticated so this is 403, not 401") +} + +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..100ceda563 --- /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 = [ + "//src/libraries/go/lib/pkg/auth/nvcaintrospect", + ], +) + +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..f5e83938ab --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/nvca/introspect.go @@ -0,0 +1,70 @@ +/* +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 wraps +// the shared nvcaintrospect.Client so Event Ledger and ReVal verify NVCA's +// identity the same way. +package nvca + +import ( + "context" + "time" + + "github.com/NVIDIA/nvcf/src/libraries/go/lib/pkg/auth/nvcaintrospect" +) + +// IntrospectRequest and IntrospectResult are the shared nvcaintrospect types, +// aliased here so existing callers and tests in this package are unaffected. +type IntrospectRequest = nvcaintrospect.IntrospectRequest +type IntrospectResult = nvcaintrospect.IntrospectResult + +// MaxTokenSize and ErrTokenTooLarge alias the shared package's token-size limit. +const MaxTokenSize = nvcaintrospect.MaxTokenSize + +var ErrTokenTooLarge = nvcaintrospect.ErrTokenTooLarge + +// IsValidNVCASubject aliases the shared package's NVCA subject validator. +var IsValidNVCASubject = nvcaintrospect.IsValidNVCASubject + +// 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) +} + +// Client calls SIS's POST /v1/nvca/tokens/introspect endpoint to verify +// NVCA's PSAT, using the shared nvcaintrospect.Client for the actual call, +// caching, and subject validation. +type Client struct { + client *nvcaintrospect.Client +} + +// NewClient builds an introspection client. introspectURL is required. A +// cacheTTL of 0 disables caching. +func NewClient(introspectURL string, timeout, cacheTTL time.Duration) (*Client, error) { + client, err := nvcaintrospect.NewClient(introspectURL, timeout, cacheTTL) + if err != nil { + return nil, err + } + return &Client{client: client}, nil +} + +// Introspect implements Introspector. +func (c *Client) Introspect(ctx context.Context, token string) (*IntrospectResult, error) { + return c.client.Introspect(ctx, token) +} 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..c73cd1ec49 --- /dev/null +++ b/src/control-plane-services/event-ledger/internal/nvca/introspect_test.go @@ -0,0 +1,153 @@ +/* +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" + "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") +} +