Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 19 additions & 13 deletions internal/identity/biscuit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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")
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand All @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
Expand Down
21 changes: 12 additions & 9 deletions internal/node/a2a_service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
}

Expand Down
13 changes: 12 additions & 1 deletion internal/node/labels_gate.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
}

Expand Down
86 changes: 86 additions & 0 deletions internal/node/openai_facade_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"}}
Expand Down
Loading