diff --git a/internal/identity/biscuit_test.go b/internal/identity/biscuit_test.go index 3b5f6190..a0e4d5e4 100644 --- a/internal/identity/biscuit_test.go +++ b/internal/identity/biscuit_test.go @@ -31,6 +31,12 @@ import ( "github.com/libp2p/go-libp2p/core/peer" ) +// testTimeout bounds datalog evaluation in these tests. It is deliberately +// generous: WithMaxDuration is a wall-clock bound, and under -race on an +// oversubscribed CI runner a goroutine can be starved long past a sub-second +// budget purely by scheduling. No test here asserts anything about timing. +const testTimeout = time.Minute + func TestVerifyBiscuit_Expiration(t *testing.T) { pub, priv, err := ed25519.GenerateKey(rand.Reader) if err != nil { @@ -77,7 +83,7 @@ func TestVerifyBiscuit_Expiration(t *testing.T) { t.Fatalf("MintBiscuitToken failed: %v", err) } - _, err = VerifyBiscuit(biscuitData, dummyPeer, []ed25519.PublicKey{pub}, 500*time.Millisecond) + _, err = VerifyBiscuit(biscuitData, dummyPeer, []ed25519.PublicKey{pub}, testTimeout) if tt.expectError && err == nil { t.Errorf("Expected error due to expiration, got nil") } @@ -123,7 +129,7 @@ func TestMintBiscuitToken_ClaimsTranslation(t *testing.T) { t.Fatalf("Failed to unmarshal biscuit: %v", err) } - authorizer, err := b.Authorizer(pub, biscuit.WithWorldOptions(datalog.WithMaxDuration(500*time.Millisecond))) + authorizer, err := b.Authorizer(pub, biscuit.WithWorldOptions(datalog.WithMaxDuration(testTimeout))) if err != nil { t.Fatalf("Failed to get authorizer: %v", err) } @@ -216,7 +222,7 @@ func TestVerifyBiscuit_Concurrent(t *testing.T) { go func() { defer wg.Done() for j := 0; j < 100; j++ { - _, err := VerifyBiscuit(biscuitData, dummyPeer, []ed25519.PublicKey{pub}, 500*time.Millisecond) + _, err := VerifyBiscuit(biscuitData, dummyPeer, []ed25519.PublicKey{pub}, testTimeout) if err != nil { t.Errorf("Concurrent verification failed: %v", err) return @@ -260,7 +266,7 @@ func TestMintBiscuitToken(t *testing.T) { if err != nil { t.Fatal(err) } - authorizer, err := b.Authorizer(pub, biscuit.WithWorldOptions(datalog.WithMaxDuration(500*time.Millisecond))) + authorizer, err := b.Authorizer(pub, biscuit.WithWorldOptions(datalog.WithMaxDuration(testTimeout))) if err != nil { t.Fatal(err) } @@ -360,7 +366,7 @@ func TestMintBiscuitToken_VariousClaimsTypes(t *testing.T) { t.Fatalf("Unmarshal biscuit failed: %v", err) } - authorizer, err := b.Authorizer(pub, AuthorizerOptions(5*time.Second)...) + authorizer, err := b.Authorizer(pub, AuthorizerOptions(testTimeout)...) if err != nil { t.Fatalf("Authorizer failed: %v", err) } @@ -439,7 +445,7 @@ func TestMintBiscuitToken_WithPolicyRoles(t *testing.T) { t.Fatalf("Unmarshal biscuit failed: %v", err) } - authorizer, err := b.Authorizer(pub, AuthorizerOptions(5*time.Second)...) + authorizer, err := b.Authorizer(pub, AuthorizerOptions(testTimeout)...) if err != nil { t.Fatalf("Authorizer failed: %v", err) } @@ -468,7 +474,7 @@ func TestMintBiscuitToken_WithPolicyRoles(t *testing.T) { t.Fatalf("Unmarshal biscuit failed: %v", err) } - authorizer, err := b.Authorizer(pub, AuthorizerOptions(5*time.Second)...) + authorizer, err := b.Authorizer(pub, AuthorizerOptions(testTimeout)...) if err != nil { t.Fatalf("Authorizer failed: %v", err) } @@ -504,7 +510,7 @@ func TestMintBiscuitToken_LabelFacts(t *testing.T) { // Every declared label is present as its own fact. for k, v := range labels { - authorizer, err := b.Authorizer(pub, AuthorizerOptions(5*time.Second)...) + authorizer, err := b.Authorizer(pub, AuthorizerOptions(testTimeout)...) if err != nil { t.Fatalf("Authorizer failed: %v", err) } @@ -526,7 +532,7 @@ func TestMintBiscuitToken_LabelFacts(t *testing.T) { if err != nil { t.Fatalf("Unmarshal biscuit failed: %v", err) } - authorizer, err := b2.Authorizer(pub, AuthorizerOptions(5*time.Second)...) + authorizer, err := b2.Authorizer(pub, AuthorizerOptions(testTimeout)...) if err != nil { t.Fatalf("Authorizer failed: %v", err) } @@ -566,7 +572,7 @@ func TestVerifyAndExtractPeerID_MultipleTrustedKeys(t *testing.T) { // trustedPublicKeys has pub1 first, pub2 second (pub2 is the signer) trustedKeys := []ed25519.PublicKey{pub1, pub2} - extractedPeer, err := VerifyAndExtractPeerID(trustedKeys, biscuitData, 5*time.Second) + extractedPeer, err := VerifyAndExtractPeerID(trustedKeys, biscuitData, testTimeout) if err != nil { t.Fatalf("VerifyAndExtractPeerID failed with multiple trusted keys: %v", err) } @@ -603,15 +609,15 @@ func TestExtractPeerIDExpiry(t *testing.T) { t.Fatal(err) } - if _, err := VerifyAndExtractPeerID(trustedKeys, fresh, 5*time.Second); err != nil { + if _, err := VerifyAndExtractPeerID(trustedKeys, fresh, testTimeout); err != nil { t.Errorf("unexpired token rejected: %v", err) } - if _, err := VerifyAndExtractPeerID(trustedKeys, expired, 5*time.Second); err == nil { + if _, err := VerifyAndExtractPeerID(trustedKeys, expired, testTimeout); err == nil { t.Error("expired token accepted by the expiry-enforcing variant") } // The refresh flow depends on the exempt variant staying permissive. - if got, err := VerifyExpiredAndExtractPeerID(trustedKeys, expired, 5*time.Second); err != nil { + if got, err := VerifyExpiredAndExtractPeerID(trustedKeys, expired, testTimeout); err != nil { t.Errorf("expired token rejected by the refresh variant: %v", err) } else if got != dummyPeer { t.Errorf("got peer %s, want %s", got, dummyPeer) diff --git a/internal/node/a2a_service_test.go b/internal/node/a2a_service_test.go index 5c234b1d..d13e3db9 100644 --- a/internal/node/a2a_service_test.go +++ b/internal/node/a2a_service_test.go @@ -79,15 +79,18 @@ func TestA2AEgressHookNonA2APassthrough(t *testing.T) { } func TestA2AEgressHookMalformedLabels(t *testing.T) { - rec := httptest.NewRecorder() - req := httptest.NewRequest("POST", "/sam/12D3KooWpeer/a2a/agent/", nil) - req.Header.Set(api.HeaderSamRequiredLabels, "not-a-label") - _, ok := applyEgressMiddleware(nil, rec, req) - if ok { - t.Fatal("malformed labels must be refused") - } - if rec.Code != http.StatusBadRequest { - t.Fatalf("status = %d, want 400", rec.Code) + // ",," is the fail-open shape: it must not parse to "no requirement". + for _, header := range []string{"not-a-label", ",,"} { + rec := httptest.NewRecorder() + req := httptest.NewRequest("POST", "/sam/12D3KooWpeer/a2a/agent/", nil) + req.Header.Set(api.HeaderSamRequiredLabels, header) + _, ok := applyEgressMiddleware(nil, rec, req) + if ok { + t.Fatalf("labels header %q must be refused", header) + } + if rec.Code != http.StatusBadRequest { + t.Fatalf("labels header %q: status = %d, want 400", header, rec.Code) + } } } diff --git a/internal/node/labels_gate.go b/internal/node/labels_gate.go index 6adcbe92..3c19efd0 100644 --- a/internal/node/labels_gate.go +++ b/internal/node/labels_gate.go @@ -32,8 +32,16 @@ import ( // parseRequiredLabels splits the X-Sam-Required-Labels header value // (comma-separated "key=value" pairs) into a label map; any malformed entry // rejects the whole request. +// +// Only a blank specification means "no requirement". A specification that +// carries content but names no label — ",," or a lone separator — is rejected +// instead of being read as unconstrained: an empty requirement set switches +// the label gate off entirely (see VerifyPeerLabels and rankProviders), so +// silently deriving one from a caller's non-blank input would turn a +// fail-closed control into a fail-open one. Empty entries *alongside* real +// ones stay tolerated, so a trailing comma is still harmless. func parseRequiredLabels(h string) (map[string]string, error) { - if h == "" { + if strings.TrimSpace(h) == "" { return nil, nil } var out map[string]string @@ -61,6 +69,9 @@ func parseRequiredLabels(h string) (map[string]string, error) { } out[k] = v } + if len(out) == 0 { + return nil, fmt.Errorf("invalid required labels %q: expected at least one key=value pair", h) + } return out, nil } diff --git a/internal/node/openai_facade_test.go b/internal/node/openai_facade_test.go index 4fc1b3b6..748b2047 100644 --- a/internal/node/openai_facade_test.go +++ b/internal/node/openai_facade_test.go @@ -478,6 +478,92 @@ func TestParseRequiredLabels(t *testing.T) { } } +// An empty requirement set turns the label gate off (VerifyPeerLabels returns +// immediately, and rankProviders skips the label filter), so the only input +// allowed to produce one is a genuinely blank specification. +func TestParseRequiredLabels_BlankSpecMeansNoRequirement(t *testing.T) { + for _, in := range []string{"", " ", " ", " "} { + got, err := parseRequiredLabels(in) + if err != nil || got != nil { + t.Errorf("parseRequiredLabels(%q): got %v, %v; want nil, nil", in, got, err) + } + } +} + +// Regression: a specification that carries content but names no label used to +// parse to zero requirements with no error, silently disabling a fail-closed +// control for a caller who asked to be constrained. It must be an error. +func TestParseRequiredLabels_ContentfulSpecNamingNoLabelFailsClosed(t *testing.T) { + for _, in := range []string{",", ",,", ",,,", " , ", " , ", " ,, "} { + got, err := parseRequiredLabels(in) + if err == nil { + t.Errorf("parseRequiredLabels(%q): got %v, nil; want an error, because zero requirements means an unconstrained request", in, got) + } + if got != nil { + t.Errorf("parseRequiredLabels(%q): labels must be nil on error, got %v", in, got) + } + } +} + +// The stricter rule must not reach empty entries that merely sit next to real +// ones: a trailing comma is a normal way to write a list and stays harmless. +func TestParseRequiredLabels_EmptyEntriesBesideRealOnesStayTolerated(t *testing.T) { + tests := []struct { + in string + want map[string]string + }{ + {"region=eu,", map[string]string{"region": "eu"}}, + {",region=eu", map[string]string{"region": "eu"}}, + {" region=eu , , team=platform ,,", map[string]string{"region": "eu", "team": "platform"}}, + } + for _, tt := range tests { + got, err := parseRequiredLabels(tt.in) + if err != nil { + t.Errorf("parseRequiredLabels(%q): unexpected error %v", tt.in, err) + continue + } + if len(got) != len(tt.want) { + t.Errorf("parseRequiredLabels(%q): got %v, want %v", tt.in, got, tt.want) + continue + } + for k, v := range tt.want { + if got[k] != v { + t.Errorf("parseRequiredLabels(%q): got[%q] = %q, want %q", tt.in, k, got[k], v) + } + } + } +} + +// The boundary behaviour the parser fix exists for: a caller who asked to be +// constrained is refused, not quietly served by an arbitrary provider. +func TestFacade_Completions_ContentfulRequiredLabelsNamingNoLabelIsRejected(t *testing.T) { + f := newTestFacade() + f.discover = func(_ context.Context) ([]*api.DiscoveredProvider, error) { + return []*api.DiscoveredProvider{{PeerId: "peerUS", SrvName: "srvUS"}}, nil + } + f.remoteModels = func(_ context.Context, _, _ string) ([]string, error) { + return []string{"m1"}, nil + } + forwarded := false + f.forward = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + forwarded = true + w.WriteHeader(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", + strings.NewReader(`{"model":"m1","messages":[{"role":"user","content":"hi"}]}`)) + req.Header.Set(api.HeaderSamRequiredLabels, ",,") + rec := httptest.NewRecorder() + f.handleCompletions(rec, req) + + if rec.Code != http.StatusBadRequest { + t.Fatalf("status: got %d, want %d (body %s)", rec.Code, http.StatusBadRequest, rec.Body.String()) + } + if forwarded { + t.Error("request must not reach a provider: the caller asked for a label constraint that named no label") + } +} + func TestRankProviders_LabelMatchIsExactAndCaseSensitive(t *testing.T) { f := newTestFacade() provider := modelProvider{peerID: "peerEU", service: "srv", labels: map[string]string{"region": "eu-de"}}