diff --git a/src/libraries/go/lib/pkg/auth/nvcaintrospect/BUILD.bazel b/src/libraries/go/lib/pkg/auth/nvcaintrospect/BUILD.bazel new file mode 100644 index 0000000000..b0fcc308fe --- /dev/null +++ b/src/libraries/go/lib/pkg/auth/nvcaintrospect/BUILD.bazel @@ -0,0 +1,42 @@ +# SPDX-FileCopyrightText: Copyright (c) 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. + +load("@rules_go//go:def.bzl", "go_library", "go_test") + +go_library( + name = "nvcaintrospect", + srcs = ["introspect.go"], + importpath = "github.com/NVIDIA/nvcf/src/libraries/go/lib/pkg/auth/nvcaintrospect", + visibility = ["//visibility:public"], + deps = [ + "@io_opentelemetry_go_contrib_instrumentation_net_http_otelhttp//:otelhttp", + "@io_opentelemetry_go_otel//:otel", + "@io_opentelemetry_go_otel//codes", + ], +) + +go_test( + name = "nvcaintrospect_test", + srcs = ["introspect_test.go"], + embed = [":nvcaintrospect"], + deps = [ + "@com_github_stretchr_testify//assert", + "@com_github_stretchr_testify//require", + "@io_opentelemetry_go_otel//:otel", + "@io_opentelemetry_go_otel//codes", + "@io_opentelemetry_go_otel_sdk//trace", + "@io_opentelemetry_go_otel_sdk//trace/tracetest", + ], +) diff --git a/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect.go b/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect.go new file mode 100644 index 0000000000..2c17abef05 --- /dev/null +++ b/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect.go @@ -0,0 +1,344 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 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 nvcaintrospect verifies NVCA's Kubernetes projected service-account +// token (PSAT) or SPIFFE SVID against a remote RFC 7662-shaped introspection +// endpoint, for callers that hold neither an OpenBao-issued JWT nor a +// Starfleet SSA token. It is the shared primitive behind ReVal's and Event +// Ledger's own introspection clients; service-specific authorization +// wiring (ReVal's Authorizer adapter, Event Ledger's cluster binding) stays +// in each service. +package nvcaintrospect + +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" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/codes" +) + +const instrumentationName = "nvcaintrospect" + +// MaxTokenSize bounds how much bearer token material the 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") + +// maxResponseSize bounds how much of the introspection response body the +// client will read. The response is a handful of short string fields; a +// faulty or compromised endpoint could otherwise stream an unbounded body. +const maxResponseSize = 64 * 1024 + +// errResponseTooLarge is returned when the introspection response exceeds +// maxResponseSize. +var errResponseTooLarge = errors.New("introspect response exceeds maximum size") + +// maxCacheEntries bounds the introspection cache so high-cardinality token +// traffic can't grow it without limit. +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 the introspection endpoint. +type IntrospectRequest struct { + Token string `json:"token"` +} + +// IntrospectResult is the response from a 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"` + Aud string `json:"aud,omitempty"` + Iss string `json:"iss,omitempty"` + ClusterID string `json:"cluster_id"` + TokenType string `json:"token_type,omitempty"` + Error string `json:"error,omitempty"` +} + +// cacheEntry stores an introspection result (allow or deny) with its +// expiration wall-clock. +type cacheEntry struct { + result *IntrospectResult + expiresAt time.Time +} + +// Client calls a remote NVCA token introspection endpoint. 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. +// +// Both an active token with a valid NVCA subject and an active token with an +// invalid subject are cached: a token's subject is immutable once issued, so +// either verdict for a specific token is safe to reuse. An inactive result is +// never cached, since clock skew or an nbf window can make the same token +// valid moments later. +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("nvcaintrospect: introspect url is required") + } + if timeout <= 0 { + timeout = 10 * time.Second + } + // Clone DefaultTransport to include the HTTP proxy env vars. Fall back to + // DefaultTransport as-is if it's been replaced with a non-*http.Transport + // RoundTripper (e.g. by an httpmock-style test helper), since the single- + // value assertion would otherwise panic. + var transport http.RoundTripper = http.DefaultTransport + if t, ok := http.DefaultTransport.(*http.Transport); ok { + transport = t.Clone() + } + return &Client{ + introspectURL: introspectURL, + httpClient: &http.Client{ + Timeout: timeout, + Transport: otelhttp.NewTransport(transport, + otelhttp.WithSpanNameFormatter(func(_ string, _ *http.Request) string { + return "nvcaintrospect.introspect" + }), + ), + // The introspection endpoint should never redirect. Following one + // would resend the bearer token (via the transport's Authorization + // header, if the caller sets one) to whatever host the redirect + // names, so refuse rather than follow. + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + cacheTTL: cacheTTL, + cache: make(map[string]cacheEntry), + }, nil +} + +// Introspect verifies token against the configured endpoint, using the cache +// when enabled. +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 + } + + if shouldCache(result) { + c.cacheStore(key, result, token) + } + + return result, nil +} + +// shouldCache reports whether an introspection result is safe to reuse for a +// token's remaining lifetime. +// +// An inactive result is never cached: clock skew or an nbf window can make +// the same token valid moments later. A result with an empty Sub is never +// cached either: that's an incomplete response from the introspection +// endpoint, not evidence of an invalid identity. A non-empty, invalid-subject +// result is cached regardless of ClusterID: a token's subject can't change +// once issued, so that denial is permanent. A valid-subject result with an +// empty ClusterID is not cached: that's also an incomplete response (a +// caller requiring ClusterID would reject it), and caching it would pin that +// rejection for the full TTL even after the endpoint starts returning a +// complete response. +func shouldCache(result *IntrospectResult) bool { + if !result.Active { + return false + } + if result.Sub == "" { + return false + } + if !IsValidNVCASubject(result.Sub) { + return true + } + return result.ClusterID != "" +} + +func (c *Client) callIntrospect(ctx context.Context, token string) (result *IntrospectResult, err error) { + ctx, span := otel.Tracer(instrumentationName).Start(ctx, "nvcaintrospect.call") + defer func() { + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() + }() + + 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(io.LimitReader(resp.Body, maxResponseSize+1)) + if err != nil { + return nil, fmt.Errorf("read introspect response: %w", err) + } + if len(respBody) > maxResponseSize { + return nil, errResponseTooLarge + } + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("introspect returned status %d", resp.StatusCode) + } + + var parsed IntrospectResult + if err = json.Unmarshal(respBody, &parsed); err != nil { + return nil, fmt.Errorf("decode introspect response: %w", err) + } + return &parsed, 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 the remote endpoint, 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 + } + cp := *entry.result + return &cp, 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 + } + stored := *result + 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: &stored, expiresAt: time.Now().Add(ttl)} + c.cacheMu.Unlock() +} diff --git a/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect_test.go b/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect_test.go new file mode 100644 index 0000000000..31348c10b1 --- /dev/null +++ b/src/libraries/go/lib/pkg/auth/nvcaintrospect/introspect_test.go @@ -0,0 +1,339 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 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 nvcaintrospect + +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" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/codes" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" +) + +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 TestClientIntrospectRejectsOversizedResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(strings.Repeat("a", maxResponseSize+1))) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, 0) + require.NoError(t, err) + + _, err = client.Introspect(context.Background(), signedTestToken(t, time.Now().Add(time.Hour))) + assert.ErrorIs(t, err, errResponseTooLarge) +} + +func TestClientIntrospectDoesNotFollowRedirects(t *testing.T) { + other := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(IntrospectResult{Active: true, Sub: "system:serviceaccount:customer-ns:nvca", ClusterID: "cluster-a"}) + })) + defer other.Close() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, other.URL, http.StatusFound) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, 0) + require.NoError(t, err) + + _, err = client.Introspect(context.Background(), signedTestToken(t, time.Now().Add(time.Hour))) + require.Error(t, err, "a redirect response must not be silently followed and treated as success") +} + +// stubRoundTripper is a http.RoundTripper that is not a *http.Transport, to +// simulate a host application replacing http.DefaultTransport (as +// github.com/jarcoal/httpmock does when activated). +type stubRoundTripper struct{} + +func (stubRoundTripper) RoundTrip(*http.Request) (*http.Response, error) { + return nil, fmt.Errorf("stubRoundTripper: not implemented") +} + +func TestNewClientDoesNotPanicWhenDefaultTransportIsReplaced(t *testing.T) { + prevTransport := http.DefaultTransport + http.DefaultTransport = stubRoundTripper{} + defer func() { http.DefaultTransport = prevTransport }() + + require.NotPanics(t, func() { + _, err := NewClient("http://example.invalid", time.Second, 0) + require.NoError(t, err) + }) +} + +func TestClientIntrospectRecordsErrorOnSpan(t *testing.T) { + recorder := tracetest.NewSpanRecorder() + prevTP := otel.GetTracerProvider() + otel.SetTracerProvider(sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(recorder))) + defer otel.SetTracerProvider(prevTP) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + client, err := NewClient(server.URL, time.Second, 0) + require.NoError(t, err) + + _, err = client.Introspect(context.Background(), signedTestToken(t, time.Now().Add(time.Hour))) + require.Error(t, err) + + var found bool + for _, span := range recorder.Ended() { + if span.Name() != "nvcaintrospect.call" { + continue + } + found = true + assert.Equal(t, codes.Error, span.Status().Code, "a callIntrospect failure must be recorded as an error on its span") + } + require.True(t, found, "expected an ended nvcaintrospect.call span") +} + +func TestClientIntrospectCachesActiveValidSubject(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) + + _, 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 TestClientIntrospectCachesActiveInvalidSubject(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:some-other-workload", + }) + })) + 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, 1, calls, "an active token's subject is immutable, so an invalid-subject result must be cached") +} + +func TestClientIntrospectDoesNotCacheValidSubjectMissingClusterID(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 valid-subject result missing ClusterID is incomplete and must not be cached, or a later complete response stays hidden behind the stale entry") +} + +func TestClientIntrospectDoesNotCacheMissingSub(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, + ClusterID: "cluster-a", + // Sub 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, "an active result with an empty Sub is incomplete, not an invalid-subject denial, and must not be cached") +} + +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 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") +} + +func TestClientIntrospectCachedResultIsIsolatedFromCallerMutation(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + 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) + result.ClusterID = "tampered" + + again, err := client.Introspect(context.Background(), token) + require.NoError(t, err) + assert.Equal(t, "cluster-a", again.ClusterID, "mutating a returned result must not corrupt the cached value") +}