diff --git a/management/internals/modules/agentnetwork/catalog/catalog.go b/management/internals/modules/agentnetwork/catalog/catalog.go index 3c7b995e5..b58743798 100644 --- a/management/internals/modules/agentnetwork/catalog/catalog.go +++ b/management/internals/modules/agentnetwork/catalog/catalog.go @@ -81,6 +81,10 @@ type Provider struct { // surface — the proxy middleware then falls back to URL sniffing // or skips request-side enrichment. ParserID string + // RouterVendors declares every parser surface a gateway route can serve. + // Leave empty for single-surface providers, where ParserID remains the + // router discriminator for backward compatibility. + RouterVendors []string // PricingSurfaces names the cost-meter pricing surfaces this // provider's Models are priced under ("openai", "anthropic", // "bedrock" — the llm.Parser surface the request parser stamps as @@ -116,8 +120,7 @@ type Provider struct { // Discovery, when non-nil, describes how to ask this vendor which // models the operator's own credential can actually reach, so the // provider form can offer a live list instead of only the hand-curated - // Models above. Nil for entries with no listing endpoint (gateways - // vary too much) — those keep free-text entry. + // Models above. Nil entries keep free-text entry. Discovery *Discovery } @@ -154,10 +157,13 @@ const ( // one from the caller is also what keeps this from being an open proxy: the // only hosts management will dial are the ones written here. type Discovery struct { - Host string - Path string - Query string - Shape ListingShape + Host string + Path string + Query string + Shape ListingShape + // ExactModelsOnly omits wildcard patterns from listings when NetBird's + // provider model rows cannot represent the vendor's matching semantics. + ExactModelsOnly bool // Headers are static headers the vendor requires beyond the credential // (Anthropic versions its API through one and rejects a request without // it). The auth header itself comes from AuthHeaderName/Template. @@ -635,6 +641,34 @@ var providers = []Provider{ }, Models: []Model{}, }, + { + ID: "agentgateway", + Kind: KindGateway, + Name: "agentgateway", + Description: "Bring your own agentgateway with trusted NetBird identity stamped on every request", + DefaultHost: "", + AuthHeaderName: "Authorization", + AuthHeaderTemplate: "Bearer ${API_KEY}", + DefaultContentType: "application/json", + BrandColor: "#8023C3", + // Agentgateway accepts both OpenAI and Anthropic request shapes. + // Leave ParserID empty so the proxy detects the shape from the URL. + ParserID: "", + RouterVendors: []string{"openai", "anthropic"}, + PricingSurfaces: []string{"openai", "anthropic"}, + Discovery: &Discovery{ + Path: "/v1/models", + Shape: ShapeOpenAIData, + ExactModelsOnly: true, + }, + IdentityInjection: &IdentityInjection{ + HeaderPair: &HeaderPairInjection{ + EndUserIDHeader: "x-netbird-user-id", + TagsHeader: "x-netbird-groups", + }, + }, + Models: []Model{}, + }, { ID: "portkey", Kind: KindGateway, diff --git a/management/internals/modules/agentnetwork/catalog/catalog_test.go b/management/internals/modules/agentnetwork/catalog/catalog_test.go index e4e887e6f..8abd3a312 100644 --- a/management/internals/modules/agentnetwork/catalog/catalog_test.go +++ b/management/internals/modules/agentnetwork/catalog/catalog_test.go @@ -5,6 +5,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/http/api" ) // TestClaudeLineupSelectable pins the models Claude Code resolves to by @@ -34,3 +36,51 @@ func TestClaudeLineupSelectable(t *testing.T) { } } } + +func TestAgentgatewayCatalogEntry(t *testing.T) { + entry, ok := Lookup("agentgateway") + require.True(t, ok, "agentgateway must be available in the provider catalog") + + assert.Equal(t, KindGateway, entry.Kind, "agentgateway must be grouped with AI gateways") + assert.Empty(t, entry.DefaultHost, "operators must provide their agentgateway proxy URL") + assert.Equal(t, "Authorization", entry.AuthHeaderName) + assert.Equal(t, "Bearer ${API_KEY}", entry.AuthHeaderTemplate) + assert.Equal(t, "application/json", entry.DefaultContentType) + assert.Empty(t, entry.ParserID, "URL detection must select the OpenAI or Anthropic parser") + assert.Equal(t, []string{"openai", "anthropic"}, entry.RouterVendors, + "agentgateway must accept both parser surfaces") + assert.Equal(t, []string{"openai", "anthropic"}, entry.PricingSurfaces, + "agentgateway models can use either pricing surface") + assert.Empty(t, entry.Models, "an empty model list makes agentgateway a catch-all route") + require.NotNil(t, entry.Discovery) + assert.Empty(t, entry.Discovery.Host, "discovery must use the configured proxy URL") + assert.Equal(t, "/v1/models", entry.Discovery.Path) + assert.Equal(t, ShapeOpenAIData, entry.Discovery.Shape) + assert.True(t, entry.Discovery.ExactModelsOnly, + "wildcard model semantics are not supported by NetBird") + + require.NotNil(t, entry.IdentityInjection) + require.NotNil(t, entry.IdentityInjection.HeaderPair) + assert.Nil(t, entry.IdentityInjection.JSONMetadata) + assert.False(t, entry.IdentityInjection.HeaderPair.Customizable, + "NetBird identity header names are part of the integration contract") + assert.Equal(t, "x-netbird-user-id", entry.IdentityInjection.HeaderPair.EndUserIDHeader) + assert.Equal(t, "x-netbird-groups", entry.IdentityInjection.HeaderPair.TagsHeader) + assert.False(t, entry.IdentityInjection.HeaderPair.EndUserIDInBody) + assert.False(t, entry.IdentityInjection.HeaderPair.TagsInBody) +} + +func TestAgentgatewayCatalogAPIResponse(t *testing.T) { + entry, ok := Lookup("agentgateway") + require.True(t, ok) + + resp := entry.ToAPIResponse() + assert.Equal(t, "agentgateway", resp.Id) + assert.Equal(t, api.AgentNetworkCatalogProviderKindGateway, resp.Kind) + assert.Empty(t, resp.Models) + require.NotNil(t, resp.IdentityInjection) + require.NotNil(t, resp.IdentityInjection.HeaderPair) + assert.False(t, resp.IdentityInjection.HeaderPair.Customizable) + assert.Equal(t, "x-netbird-user-id", resp.IdentityInjection.HeaderPair.EndUserIdHeader) + assert.Equal(t, "x-netbird-groups", resp.IdentityInjection.HeaderPair.TagsHeader) +} diff --git a/management/internals/modules/agentnetwork/modeldiscovery/discovery.go b/management/internals/modules/agentnetwork/modeldiscovery/discovery.go index 253cc63b3..c9f2b09df 100644 --- a/management/internals/modules/agentnetwork/modeldiscovery/discovery.go +++ b/management/internals/modules/agentnetwork/modeldiscovery/discovery.go @@ -52,9 +52,8 @@ const ( ) // ErrNoDiscovery is returned for a catalog entry that declares no listing -// endpoint. Gateways vary too much to have one, and the caller should fall -// back to the catalog list plus free-text entry rather than treating this as -// a failure. +// endpoint. The caller should fall back to the catalog list plus free-text +// entry rather than treating this as a failure. var ErrNoDiscovery = errors.New("provider has no model-discovery endpoint") // ErrInvalidRequest marks a discovery failure caused by the caller's own input @@ -356,6 +355,9 @@ func decorate(entry catalog.Provider, ids []listedModel) []Model { if listed.id == "" { continue } + if entry.Discovery.ExactModelsOnly && strings.Contains(listed.id, "*") { + continue + } if _, dup := seen[listed.id]; dup { continue } diff --git a/management/internals/modules/agentnetwork/modeldiscovery/discovery_test.go b/management/internals/modules/agentnetwork/modeldiscovery/discovery_test.go index 133bd5148..59b21a2fe 100644 --- a/management/internals/modules/agentnetwork/modeldiscovery/discovery_test.go +++ b/management/internals/modules/agentnetwork/modeldiscovery/discovery_test.go @@ -59,6 +59,13 @@ const openAIListing = `{"object":"list","data":[ {"id":"gpt-4o","object":"model","created":1715367049,"owned_by":"system"} ]}` +const agentgatewayListing = `{"object":"list","data":[ + {"id":"gpt-4o-mini","object":"model","created":1785166485,"owned_by":"openai"}, + {"id":"claude-haiku-4-5","object":"model","created":1785166485,"owned_by":"anthropic"}, + {"id":"openai/*","object":"model","created":1785166485,"owned_by":"openai"}, + {"id":"*-latest","object":"model","created":1785166485,"owned_by":"openai"} +]}` + const anthropicListing = `{"data":[ {"type":"model","id":"claude-haiku-4-5-20251001","display_name":"Claude Haiku 4.5"}, {"type":"model","id":"claude-sonnet-4-6","display_name":"Claude Sonnet 4.6"} @@ -97,6 +104,26 @@ func TestFetchOpenAIListing(t *testing.T) { } } +func TestFetchAgentgatewayListing(t *testing.T) { + cl, tr := newStubClient(http.StatusOK, agentgatewayListing) + + models, err := cl.Fetch(context.Background(), Request{ + CatalogID: "agentgateway", + UpstreamURL: "https://gateway.example.com", + APIKey: "virtual-key", + }) + require.NoError(t, err) + + assert.Equal(t, "https://gateway.example.com/v1/models", tr.got.URL.String()) + assert.Equal(t, "Bearer virtual-key", tr.got.Header.Get("Authorization"), + "agentgateway model discovery must use the configured virtual key") + assert.Equal(t, []string{"gpt-4o-mini", "claude-haiku-4-5"}, ids(models), + "model patterns must not be offered as exact NetBird authorization rows") + for _, m := range models { + assert.True(t, m.PricingKnown, "known upstream model must use NetBird catalog pricing: %s", m.ID) + } +} + func TestFetchAnthropicSendsTheVersionHeader(t *testing.T) { cl, tr := newStubClient(http.StatusOK, anthropicListing) diff --git a/management/internals/modules/agentnetwork/synthesizer.go b/management/internals/modules/agentnetwork/synthesizer.go index 66a19acd9..b838ac547 100644 --- a/management/internals/modules/agentnetwork/synthesizer.go +++ b/management/internals/modules/agentnetwork/synthesizer.go @@ -352,6 +352,7 @@ type routerConfig struct { type routerProviderRoute struct { ID string `json:"id"` Vendor string `json:"vendor,omitempty"` + Vendors []string `json:"vendors,omitempty"` Models []string `json:"models"` UpstreamScheme string `json:"upstream_scheme"` UpstreamHost string `json:"upstream_host"` @@ -461,6 +462,7 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][] cfg.Providers = append(cfg.Providers, routerProviderRoute{ ID: p.ID, Vendor: providerVendor(p), + Vendors: providerVendors(p), Models: providerModelIDs(p), UpstreamScheme: scheme, UpstreamHost: host, @@ -525,6 +527,17 @@ func providerVendor(p *types.Provider) string { return entry.ParserID } +// providerVendors returns the parser surfaces a multi-surface gateway route +// accepts. Single-surface providers keep using the singular vendor field so +// existing proxy versions and configurations retain their wire shape. +func providerVendors(p *types.Provider) []string { + entry, ok := catalog.Lookup(p.ProviderID) + if !ok || len(entry.RouterVendors) == 0 { + return nil + } + return append([]string(nil), entry.RouterVendors...) +} + // providerModelIDs returns the model identifiers exposed by the // provider, deduplicated and in the operator's declared order. Empty // slice when no models are configured — the router treats that as diff --git a/management/internals/modules/agentnetwork/synthesizer_test.go b/management/internals/modules/agentnetwork/synthesizer_test.go index 352d36646..6aeafadbd 100644 --- a/management/internals/modules/agentnetwork/synthesizer_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_test.go @@ -6,9 +6,9 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" @@ -497,6 +497,55 @@ func TestSynthesizeServices_IdentityInject_LiteLLM(t *testing.T) { assert.Equal(t, "x-litellm-tags", entry.HeaderPair.TagsHeader) } +func TestBuildIdentityInjectConfigJSON_Agentgateway(t *testing.T) { + provider := &types.Provider{ + ID: "prov-agentgateway", + ProviderID: "agentgateway", + } + + raw, err := buildIdentityInjectConfigJSON( + []*types.Provider{provider}, + map[string][]string{provider.ID: []string{"grp-eng"}}, + ) + require.NoError(t, err) + + var cfg identityInjectConfig + require.NoError(t, json.Unmarshal(raw, &cfg)) + require.Len(t, cfg.Providers, 1) + + rule := cfg.Providers[0] + assert.Equal(t, provider.ID, rule.ProviderID) + require.NotNil(t, rule.HeaderPair) + assert.Nil(t, rule.JSONMetadata) + assert.Equal(t, "x-netbird-user-id", rule.HeaderPair.EndUserIDHeader) + assert.Equal(t, "x-netbird-groups", rule.HeaderPair.TagsHeader) + assert.False(t, rule.HeaderPair.EndUserIDInBody) + assert.False(t, rule.HeaderPair.TagsInBody) +} + +func TestBuildRouterConfigJSON_AgentgatewayVendors(t *testing.T) { + provider := &types.Provider{ + ID: "prov-agentgateway", + ProviderID: "agentgateway", + UpstreamURL: "https://gateway.example.com", + APIKey: "virtual-key", + } + + raw, err := buildRouterConfigJSON( + []*types.Provider{provider}, + map[string][]string{provider.ID: {"grp-eng"}}, + nil, + ) + require.NoError(t, err) + + var cfg routerConfig + require.NoError(t, json.Unmarshal(raw, &cfg)) + require.Len(t, cfg.Providers, 1) + assert.Empty(t, cfg.Providers[0].Vendor, + "the singular vendor remains empty for a multi-surface gateway") + assert.Equal(t, []string{"openai", "anthropic"}, cfg.Providers[0].Vendors) +} + // TestSynthesizeServices_IdentityInject_Bifrost_OperatorOverrides // covers the customizable HeaderPair contract. The Bifrost catalog // entry sets HeaderPair.Customizable=true with x-bf-dim-* defaults diff --git a/proxy/internal/middleware/builtin/llm_router/factory.go b/proxy/internal/middleware/builtin/llm_router/factory.go index 81b8727f1..70a2179b5 100644 --- a/proxy/internal/middleware/builtin/llm_router/factory.go +++ b/proxy/internal/middleware/builtin/llm_router/factory.go @@ -36,7 +36,10 @@ type ProviderRoute struct { // request on a same-vendor route so catch-all gateways of a different // vendor can't swallow it. Empty disables vendor filtering for this // route. - Vendor string `json:"vendor,omitempty"` + Vendor string `json:"vendor,omitempty"` + // Vendors lists every parser surface a multi-surface gateway accepts. + // Vendor remains supported for existing single-surface configurations. + Vendors []string `json:"vendors,omitempty"` Models []string `json:"models"` UpstreamScheme string `json:"upstream_scheme"` UpstreamHost string `json:"upstream_host"` diff --git a/proxy/internal/middleware/builtin/llm_router/middleware.go b/proxy/internal/middleware/builtin/llm_router/middleware.go index b8d4b001b..6381f01c7 100644 --- a/proxy/internal/middleware/builtin/llm_router/middleware.go +++ b/proxy/internal/middleware/builtin/llm_router/middleware.go @@ -409,7 +409,7 @@ func stripBedrockNamespace(out *middleware.Output) { // peer, return matchOutcomeUnauthorised so the caller can emit // the dedicated no_authorised_provider deny code. // 3. Vendor precedence: when the request carries a detected vendor -// (llm.provider) and at least one candidate is the same vendor, +// (llm.provider) and at least one candidate declares that vendor, // drop the rest — a vendor-tagged request must never cross to // another vendor's route (e.g. an Anthropic call landing on an // OpenAI-compatible gateway that also claims the model). @@ -432,9 +432,9 @@ func (m *Middleware) matchRoute(model, vendor, reqPath string, userGroups []stri // Vendor pinning runs BEFORE the group filter so a request the parser // tagged with a vendor can never cross to another vendor's route — not - // even an authorised one. Narrow to same-vendor routes when any - // model-matched route declares that vendor; setups with no vendor tag on - // any route fall through unchanged. After narrowing, if no same-vendor + // even an authorised one. Narrow to supporting routes when any + // model-matched route declares that vendor; setups with no matching vendor + // declaration fall through unchanged. After narrowing, if no supporting // route authorises the caller, that's matchOutcomeUnauthorised (no // cross-vendor fallback). if vendor != "" { @@ -805,21 +805,31 @@ func authorisingGroupsCSV(routeGroups, userGroups []string) string { return strings.Join(out, ",") } -// matchingVendor returns the subset of routes whose Vendor equals the -// request's detected vendor. Routes with an empty Vendor never match — an -// untagged route can't be asserted to speak the request's surface, so it -// stays out of the vendor-filtered set (but remains eligible via the -// fall-through when no route matches the vendor at all). +// matchingVendor returns the routes that declare the request's detected +// vendor through either the legacy singular field or the multi-vendor field. +// Untagged routes remain eligible only when no route declares the vendor. func matchingVendor(routes []ProviderRoute, vendor string) []ProviderRoute { var out []ProviderRoute for _, r := range routes { - if r.Vendor == vendor { + if routeSupportsVendor(r, vendor) { out = append(out, r) } } return out } +func routeSupportsVendor(route ProviderRoute, vendor string) bool { + if route.Vendor == vendor { + return true + } + for _, candidate := range route.Vendors { + if candidate == vendor { + return true + } + } + return false +} + // explicitlyClaiming returns the subset of routes whose Models list // names the model exactly. Catch-all routes (empty Models) are excluded, // so callers can prefer a provider that genuinely declares the model over diff --git a/proxy/internal/middleware/builtin/llm_router/middleware_test.go b/proxy/internal/middleware/builtin/llm_router/middleware_test.go index 5a1d32480..8612d8f18 100644 --- a/proxy/internal/middleware/builtin/llm_router/middleware_test.go +++ b/proxy/internal/middleware/builtin/llm_router/middleware_test.go @@ -412,6 +412,50 @@ func TestRouter_VendorKeepsOpenAIOffAnthropic(t *testing.T) { assert.Equal(t, "api.openai.com", out.Mutations.RewriteUpstream.Host, "openai vendor must pin to the openai route despite anthropic being declared first") } +func TestRouter_MultiVendorGatewayAcceptsBothSurfaces(t *testing.T) { + gateway := ProviderRoute{ + ID: "agentgateway", + Vendors: []string{"openai", "anthropic"}, + Models: nil, + AllowedGroupIDs: []string{defaultTestGroup}, + UpstreamScheme: "https", + UpstreamHost: "gateway.example.com", + } + other := ProviderRoute{ + ID: "other-vendor", + Vendor: "mistral", + Models: nil, + AllowedGroupIDs: []string{defaultTestGroup}, + UpstreamScheme: "https", + UpstreamHost: "mistral.example.com", + } + mw := New(Config{Providers: []ProviderRoute{other, gateway}}) + + for _, tc := range []struct { + name string + vendor string + model string + path string + }{ + {name: "OpenAI", vendor: "openai", model: "gpt-4o-mini", path: "/v1/chat/completions"}, + {name: "Anthropic", vendor: "anthropic", model: "claude-sonnet-4-5", path: "/v1/messages"}, + } { + t.Run(tc.name, func(t *testing.T) { + out, err := mw.Invoke(context.Background(), newInputVendorModelURL(tc.vendor, tc.model, tc.path)) + require.NoError(t, err) + require.NotNil(t, out) + assert.Equal(t, middleware.DecisionAllow, out.Decision, + "supported vendor must route through the multi-surface gateway") + require.NotNil(t, out.Mutations) + require.NotNil(t, out.Mutations.RewriteUpstream) + assert.Equal(t, "gateway.example.com", out.Mutations.RewriteUpstream.Host) + + provider, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID) + assert.Equal(t, "agentgateway", provider) + }) + } +} + // TestRouter_VendorAbsentFallsBackToModelPath confirms vendor filtering is // inert when the request carries no detected vendor: routing then relies on // model/path as before. @@ -692,6 +736,23 @@ func TestRouter_FactoryRejectsBadJSON(t *testing.T) { require.Error(t, err, "malformed JSON config must be rejected at chain build time") } +func TestRouter_FactoryDecodesLegacyAndMultiVendorFields(t *testing.T) { + raw := []byte(`{"providers":[` + + `{"id":"legacy","vendor":"openai","models":[],"upstream_scheme":"https","upstream_host":"openai.example.com","auth_header_name":"Authorization","auth_header_value":"Bearer legacy","allowed_group_ids":["group"]},` + + `{"id":"multi","vendors":["openai","anthropic"],"models":[],"upstream_scheme":"https","upstream_host":"gateway.example.com","auth_header_name":"Authorization","auth_header_value":"Bearer multi","allowed_group_ids":["group"]}` + + `]}`) + + resolved, err := Factory{}.New(raw) + require.NoError(t, err) + router, ok := resolved.(*Middleware) + require.True(t, ok, "factory must return the concrete router middleware") + require.Len(t, router.cfg.Providers, 2) + assert.Equal(t, "openai", router.cfg.Providers[0].Vendor, + "the legacy singular field must keep decoding") + assert.Equal(t, []string{"openai", "anthropic"}, router.cfg.Providers[1].Vendors, + "the multi-vendor field must decode both supported surfaces") +} + func TestRouter_FactoryAcceptsEmptyShapes(t *testing.T) { cases := [][]byte{nil, []byte(""), []byte(" "), []byte("null"), []byte("{}"), []byte("[]")} for _, raw := range cases { diff --git a/proxy/internal/middleware/chain.go b/proxy/internal/middleware/chain.go index 45d32cdb0..9eed93678 100644 --- a/proxy/internal/middleware/chain.go +++ b/proxy/internal/middleware/chain.go @@ -264,7 +264,7 @@ func applyMutations(ctx context.Context, d *Dispatcher, spec Spec, r *http.Reque if m == nil { return } - add, remove, blocked := FilterHeaderMutations(m) + add, remove, blocked := filterHeaderMutations(m, spec.ID) for _, h := range blocked { d.metrics.IncHeaderMutationBlocked(ctx, spec.ID, h) } diff --git a/proxy/internal/middleware/chain_test.go b/proxy/internal/middleware/chain_test.go index 929ccee08..ffb23271e 100644 --- a/proxy/internal/middleware/chain_test.go +++ b/proxy/internal/middleware/chain_test.go @@ -2,6 +2,7 @@ package middleware import ( "context" + "net/http" "strconv" "testing" @@ -278,6 +279,64 @@ func TestChain_ApplyMutations_RewriteGatedOnCanMutate(t *testing.T) { assert.Nil(t, rewrite, "rewrite must be filtered when CanMutate=false") } +func TestChain_IdentityInjectReplacesReservedNetBirdHeaders(t *testing.T) { + mw := &fakeMiddleware{ + id: "llm_identity_inject", + slot: SlotOnRequest, + mutationsSupported: true, + canMutate: true, + mutations: &Mutations{ + HeadersRemove: []string{"x-netbird-user-id", "x-netbird-groups"}, + HeadersAdd: []KV{ + {Key: "x-netbird-user-id", Value: "trusted-user"}, + {Key: "x-netbird-groups", Value: "trusted-group"}, + }, + }, + } + c := chainFor(t, mw) + req, err := http.NewRequest(http.MethodGet, "https://gateway.example.com/v1/models", nil) + require.NoError(t, err) + req.Header.Set("x-netbird-user-id", "spoofed-user") + req.Header.Set("x-netbird-groups", "spoofed-group") + + denied, _, _, err := c.RunRequest(context.Background(), req, &Input{}, NewAccumulator(0)) + require.NoError(t, err) + assert.Nil(t, denied, "identity injection must not deny the request") + assert.Equal(t, "trusted-user", req.Header.Get("x-netbird-user-id"), + "the built-in identity middleware must replace a spoofed user header") + assert.Equal(t, "trusted-group", req.Header.Get("x-netbird-groups"), + "the built-in identity middleware must replace spoofed groups") +} + +func TestChain_OtherMiddlewareCannotReplaceReservedNetBirdHeaders(t *testing.T) { + mw := &fakeMiddleware{ + id: "untrusted-middleware", + slot: SlotOnRequest, + mutationsSupported: true, + canMutate: true, + mutations: &Mutations{ + HeadersRemove: []string{"x-netbird-user-id", "x-netbird-groups"}, + HeadersAdd: []KV{ + {Key: "x-netbird-user-id", Value: "replacement-user"}, + {Key: "x-netbird-groups", Value: "replacement-group"}, + }, + }, + } + c := chainFor(t, mw) + req, err := http.NewRequest(http.MethodGet, "https://gateway.example.com/v1/models", nil) + require.NoError(t, err) + req.Header.Set("x-netbird-user-id", "original-user") + req.Header.Set("x-netbird-groups", "original-group") + + denied, _, _, err := c.RunRequest(context.Background(), req, &Input{}, NewAccumulator(0)) + require.NoError(t, err) + assert.Nil(t, denied, "blocked mutations must not deny the request") + assert.Equal(t, "original-user", req.Header.Get("x-netbird-user-id"), + "other middleware must remain unable to mutate reserved identity headers") + assert.Equal(t, "original-group", req.Header.Get("x-netbird-groups"), + "other middleware must remain unable to mutate reserved identity headers") +} + // TestChain_RunRequest_PropagatesUserGroups asserts the chain forwards // Input.UserGroups verbatim through cloneInputFor so policy-aware // middlewares (e.g. llm_policy_check) can authorise without an extra diff --git a/proxy/internal/middleware/headerpolicy.go b/proxy/internal/middleware/headerpolicy.go index d041ad1e1..b1fa564c2 100644 --- a/proxy/internal/middleware/headerpolicy.go +++ b/proxy/internal/middleware/headerpolicy.go @@ -2,6 +2,8 @@ package middleware import "strings" +const trustedIdentityMiddlewareID = "llm_identity_inject" + var denyHeaders = []string{ "Authorization", "Connection", @@ -78,18 +80,22 @@ func isHeaderFieldName(name string) bool { // header names so the dispatcher can increment the blocked-header // metric. func FilterHeaderMutations(m *Mutations) (filteredAdd []KV, filteredRemove []string, blocked []string) { + return filterHeaderMutations(m, "") +} + +func filterHeaderMutations(m *Mutations, middlewareID string) (filteredAdd []KV, filteredRemove []string, blocked []string) { if m == nil { return nil, nil, nil } for _, kv := range m.HeadersAdd { - if IsHeaderMutable(kv.Key) { + if IsHeaderMutable(kv.Key) || isTrustedIdentityHeader(middlewareID, kv.Key) { filteredAdd = append(filteredAdd, kv) continue } blocked = append(blocked, kv.Key) } for _, name := range m.HeadersRemove { - if IsHeaderMutable(name) { + if IsHeaderMutable(name) || isTrustedIdentityHeader(middlewareID, name) { filteredRemove = append(filteredRemove, name) continue } @@ -97,3 +103,11 @@ func FilterHeaderMutations(m *Mutations) (filteredAdd []KV, filteredRemove []str } return filteredAdd, filteredRemove, blocked } + +func isTrustedIdentityHeader(middlewareID, name string) bool { + if middlewareID != trustedIdentityMiddlewareID { + return false + } + return strings.EqualFold(name, "x-netbird-user-id") || + strings.EqualFold(name, "x-netbird-groups") +} diff --git a/proxy/internal/middleware/headerpolicy_test.go b/proxy/internal/middleware/headerpolicy_test.go new file mode 100644 index 000000000..7daa93eec --- /dev/null +++ b/proxy/internal/middleware/headerpolicy_test.go @@ -0,0 +1,26 @@ +package middleware + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestFilterHeaderMutationsDoesNotTrustReservedHeaders(t *testing.T) { + mutations := &Mutations{ + HeadersAdd: []KV{ + {Key: "x-request-label", Value: "allowed"}, + {Key: "x-netbird-user-id", Value: "spoofed-user"}, + }, + HeadersRemove: []string{"x-request-label", "x-netbird-groups"}, + } + + filteredAdd, filteredRemove, blocked := FilterHeaderMutations(mutations) + + assert.Equal(t, []KV{{Key: "x-request-label", Value: "allowed"}}, filteredAdd, + "the public filter should retain mutable additions") + assert.Equal(t, []string{"x-request-label"}, filteredRemove, + "the public filter should retain mutable removals") + assert.ElementsMatch(t, []string{"x-netbird-user-id", "x-netbird-groups"}, blocked, + "the public filter must not grant the identity middleware exception") +}