diff --git a/docs/agent-networks/01-end-to-end-flows.md b/docs/agent-networks/01-end-to-end-flows.md index 7264f3768..f498f8736 100644 --- a/docs/agent-networks/01-end-to-end-flows.md +++ b/docs/agent-networks/01-end-to-end-flows.md @@ -109,7 +109,7 @@ sequenceDiagram Chk->>Inj: continue Inj->>Inj: inject NetBird identity headers per provider config Inj->>Grd: continue - Grd->>Grd: enforce model allowlist + Grd->>Grd: enforce per-provider allowlist (fail-closed backstop) Grd->>Up: forward (over WireGuard) Up-->>Resp: response (JSON or SSE stream) Resp->>Resp: parse usage tokens, completion @@ -135,6 +135,15 @@ sequenceDiagram (`redact_pii = settings.RedactPii`). Phones, emails, credit cards, PII names — see `redact.go` for the full set. See [`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md). +- The model allowlist is enforced in TWO places. `CheckLLMPolicyLimits` + is authoritative: it resolves the policy that governs this + (provider, caller-groups) and denies (`deny_code = llm_policy.model_blocked`) + when no applicable policy permits the model — so an allowlist scoped to + one group/provider never leaks to another, and an un-guardrailed policy + is genuinely unrestricted. `llm_guardrail` is a per-provider fail-closed + backstop: it only carries an allowlist for a provider every authorising + policy restricts, and blocks unknown/undetermined models even when + management is unreachable. - SSE streaming requires special handling on the response side; the parser must handle partial chunks without buffering the whole stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md). diff --git a/docs/agent-networks/modules/21-management-agentnetwork.md b/docs/agent-networks/modules/21-management-agentnetwork.md index b64c1ba20..cc74206e9 100644 --- a/docs/agent-networks/modules/21-management-agentnetwork.md +++ b/docs/agent-networks/modules/21-management-agentnetwork.md @@ -122,7 +122,7 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest | on_request | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** | | on_request | 2 | `llm_limit_check` | `{}` | – | | on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** | - | on_request | 4 | `llm_guardrail` | `{"model_allowlist"?, "prompt_capture":{enabled,redact_pii}}` | – | + | on_request | 4 | `llm_guardrail` | `{"provider_allowlists"?: {providerID: []model}, "prompt_capture":{enabled,redact_pii}}` | – | | on_response | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | – | | on_response | 6 | `cost_meter` | `{}` | – | | on_response | 7 | `llm_response_parser` | `{"capture_completion": , "redact_pii"?: true}` | – | diff --git a/docs/agent-networks/modules/31-proxy-middleware-builtin.md b/docs/agent-networks/modules/31-proxy-middleware-builtin.md index 904de6424..efe1bc4ce 100644 --- a/docs/agent-networks/modules/31-proxy-middleware-builtin.md +++ b/docs/agent-networks/modules/31-proxy-middleware-builtin.md @@ -244,7 +244,7 @@ no mocks. Tests: `TestChain_AllowPath_StampsAttributionAndRecordsCounter` | `llm_router` | `{providers: [{id, models, upstream_scheme, upstream_host, upstream_path?, auth_header_name, auth_header_value, allowed_group_ids}]}` | | `llm_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` | | `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` | -| `llm_guardrail` | `{model_allowlist: []string, prompt_capture: {enabled, redact_pii}}` | +| `llm_guardrail` | `{provider_allowlists: {providerID: []string}, prompt_capture: {enabled, redact_pii}}` — allowlist keyed by resolved provider id; a provider absent from the map is unrestricted (fail-closed backstop; authoritative per-policy/group check is management's `CheckLLMPolicyLimits`) | | `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` | | `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) | | `llm_limit_record` | `{}` — same pattern as `llm_limit_check` | diff --git a/e2e/agentnetwork/guardrail_multipolicy_test.go b/e2e/agentnetwork/guardrail_multipolicy_test.go new file mode 100644 index 000000000..f78c93d28 --- /dev/null +++ b/e2e/agentnetwork/guardrail_multipolicy_test.go @@ -0,0 +1,232 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// TestGuardrailMultiPolicyModelAllowlist is the end-to-end regression guard for +// the multi-policy interactions of the model-allowlist guardrail: what happens +// when several enabled policies govern an account and some carry a guardrail +// while others don't. The earlier implementation merged every policy's allowlist +// into ONE account-wide union enforced flat on every request, which produced two +// defects this test pins: +// +// - false-ALLOW (cross-group leak): a model allowlisted only for another +// group's policy became usable by any caller, because the union ignored +// which policy/group actually authorised the request; and +// - false-DENY: a policy with NO guardrail (intended unrestricted) still had +// its traffic blocked by some other policy's allowlist. +// +// The fix scopes enforcement to the matched policy/group in management +// (SelectPolicyForRequest) with a per-provider fail-closed backstop at the +// proxy. The account here has one client in grpMain and three policies over two +// catch-all upstreams (the mock vLLM answers any model 200, so a 403 can only be +// a policy decision, never a routing miss): +// +// - polMain : grpMain -> pRestricted, guardrail allowlisting modelSelected +// - polOther: grpOther -> pRestricted, guardrail allowlisting modelOther +// - polOpen : grpMain -> pOpen, NO guardrail (unrestricted) +// +// Providers declare their models so routing is deterministic. Over the tunnel, +// as the grpMain client: +// +// - modelSelected on pRestricted is served (200) — allowed by grpMain's policy; +// - modelOther on pRestricted is denied 403 (llm_policy.model_blocked) — it is +// allowlisted only for grpOther and must NOT leak to grpMain; and +// - openModel on pOpen is served (200) — the un-guardrailed policy leaves that +// provider unrestricted and must NOT be blocked by another policy's list. +// +// Under the old account-wide union the middle case returned 200 (the leak) and +// the last case returned 403 (the false-deny); both are inverted here. +func TestGuardrailMultiPolicyModelAllowlist(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + modelSelected = "e2e-selected" + modelOther = "e2e-other" + openModel = "e2e-open" + ) + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-main"}) + require.NoError(t, err, "create main group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) }) + + grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-other"}) + require.NoError(t, err, "create other group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-guardrail-mp-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grpMain.Id}, // client joins grpMain only + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + staticKey := "static-e2e-token" + models := func(ids ...string) *[]api.AgentNetworkProviderModel { + out := make([]api.AgentNetworkProviderModel, 0, len(ids)) + for _, id := range ids { + out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001}) + } + return &out + } + + // pRestricted declares the two guardrailed models so routing is deterministic + // (model -> provider). Created first, so it carries the bootstrap cluster. + pRestricted, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "restricted", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: models(modelSelected, modelOther), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + }) + require.NoError(t, err, "create restricted provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pRestricted.Id) }) + + pOpen, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "open", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: models(openModel), + }) + require.NoError(t, err, "create open provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pOpen.Id) }) + + mkGuardrail := func(name, model string) api.AgentNetworkGuardrail { + var gr api.AgentNetworkGuardrailRequest + gr.Name = name + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{model} + g, gerr := srv.CreateGuardrail(ctx, gr) + require.NoError(t, gerr, "create guardrail %s", name) + t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) }) + return g + } + gMain := mkGuardrail("e2e-guardrail-mp-main", modelSelected) + gOther := mkGuardrail("e2e-guardrail-mp-other", modelOther) + + enabled := true + // polMain: grpMain restricted to modelSelected on pRestricted. + polMain, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-main", + Enabled: &enabled, + SourceGroups: []string{grpMain.Id}, + DestinationProviderIds: []string{pRestricted.Id}, + GuardrailIds: &[]string{gMain.Id}, + }) + require.NoError(t, err, "create main policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) }) + + // polOther: grpOther restricted to modelOther on the SAME provider. The + // client is not in grpOther, so modelOther must never be usable by it. + polOther, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-other", + Enabled: &enabled, + SourceGroups: []string{grpOther.Id}, + DestinationProviderIds: []string{pRestricted.Id}, + GuardrailIds: &[]string{gOther.Id}, + }) + require.NoError(t, err, "create other policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) }) + + // polOpen: grpMain on pOpen with NO guardrail — unrestricted. + polOpen, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-open", + Enabled: &enabled, + SourceGroups: []string{grpMain.Id}, + DestinationProviderIds: []string{pOpen.Id}, + }) + require.NoError(t, err, "create open policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOpen.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-mp-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(model string) (int, string) { + code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "") + require.NoError(t, cerr, "request must reach the proxy") + return code, body + } + // sendUntil200 absorbs first-call tunnel/DNS jitter on the freshly warmed tunnel. + sendUntil200 := func(model string) (int, string) { + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + code, body = send(model) + if code == 200 { + break + } + time.Sleep(5 * time.Second) + } + return code, body + } + + t.Run("selected model allowed for its group", func(t *testing.T) { + code, body := sendUntil200(modelSelected) + assert.Equal(t, 200, code, + "grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + }) + + t.Run("other group's model does not leak", func(t *testing.T) { + // modelOther is allowlisted only for grpOther. The grpMain client must be + // denied by management's per-policy/group check — not waved through by an + // account-wide union. This is the security-critical wrong-ALLOW guard. + code, body := send(modelOther) + assert.Equal(t, 403, code, + "another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + assert.Contains(t, body, "llm_policy.model_blocked", + "denial must be a model-allowlist decision; body: %s", body) + }) + + t.Run("unguarded policy leaves its provider unrestricted", func(t *testing.T) { + // polOpen carries no guardrail, so pOpen is unrestricted for grpMain. The + // old account-wide union would have blocked openModel (it is on no + // allowlist); it must now be served — the false-DENY guard. + code, body := sendUntil200(openModel) + assert.Equal(t, 200, code, + "an un-guardrailed policy's provider must not be blocked by another policy's allowlist; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + }) +} diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index d88e0d77c..9c1f979d5 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -86,12 +86,20 @@ type Manager interface { // PolicySelectionInput is the per-request selection envelope. The // proxy populates it from CapturedData (account, user, groups) plus -// the provider llm_router resolved. +// the provider llm_router resolved and the model it extracted. type PolicySelectionInput struct { AccountID string UserID string GroupIDs []string ProviderID string + // Model is the (already-normalised) upstream model identifier the proxy + // extracted from the request. The proxy's request parser strips + // path-routed provider decoration (Bedrock region/version, Vertex + // @version) before emitting it, so a plain case-insensitive compare + // against a guardrail allowlist is sufficient here. Empty when the model + // could not be determined — treated as "not permitted" by any policy that + // restricts models (fail closed), mirroring the proxy guardrail. + Model string } // PolicySelectionResult names the policy that "pays" for this request diff --git a/management/internals/modules/agentnetwork/policyselect.go b/management/internals/modules/agentnetwork/policyselect.go index 9203a1910..f2495d6c5 100644 --- a/management/internals/modules/agentnetwork/policyselect.go +++ b/management/internals/modules/agentnetwork/policyselect.go @@ -5,6 +5,7 @@ import ( "fmt" "math" "sort" + "strings" "time" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" @@ -35,6 +36,11 @@ const ( denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded" //nolint:gosec // account deny code label, not a credential denyCodeAccountBudgetCapExceeded = "llm_account.budget_cap_exceeded" + // denyCodeModelBlocked is returned when policies govern the request's + // (provider, caller-groups) but none of them permits the requested model + // under its guardrail allowlist. Matches the proxy guardrail's code so the + // two enforcement layers surface the same label. + denyCodeModelBlocked = "llm_policy.model_blocked" ) // consumptionCache holds the consumption counters prefetched for one @@ -159,6 +165,34 @@ func (m *managerImpl) SelectPolicyForRequest(ctx context.Context, in PolicySelec } candidates := filterApplicablePolicies(policies, in) + // Model-allowlist gate (per-policy, per-group). Among the policies that + // authorise this (provider, caller-groups), keep only those whose guardrails + // permit the requested model; a policy with no allowlist-enabled guardrail is + // unrestricted and always permits. When policies govern this request but none + // permits the model, deny. This is the authoritative allowlist decision: + // because it is scoped to the matched policies, a model allowlisted for one + // group/provider never leaks to another (no account-wide union), and a policy + // that carries no guardrail is genuinely unrestricted rather than being caught + // by some other policy's allowlist. + // Only consult guardrails when at least one candidate policy references one; + // a policy with no guardrail is unrestricted, so when none of them carries a + // guardrail the gate is a no-op and we skip the store read entirely. + if len(candidates) > 0 && anyPolicyHasGuardrails(candidates) { + guardrailsByID, gErr := m.loadGuardrailsByID(ctx, in.AccountID) + if gErr != nil { + return nil, gErr + } + permitted := filterModelPermittedPolicies(candidates, guardrailsByID, in.Model) + if len(permitted) == 0 { + return &PolicySelectionResult{ + Allow: false, + DenyCode: denyCodeModelBlocked, + DenyReason: modelBlockedReason(in.Model), + }, nil + } + candidates = permitted + } + // Prefetch every consumption counter the ceiling + candidate policies will // read, in a single store round-trip, then score against the cache. cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now) @@ -250,6 +284,93 @@ func filterApplicablePolicies(policies []*types.Policy, in PolicySelectionInput) return out } +// anyPolicyHasGuardrails reports whether any policy references at least one +// guardrail, so the selector can skip loading guardrails when none do. +func anyPolicyHasGuardrails(policies []*types.Policy) bool { + for _, p := range policies { + if p != nil && len(p.GuardrailIDs) > 0 { + return true + } + } + return false +} + +// loadGuardrailsByID loads the account's guardrails indexed by ID. Used by the +// model-allowlist gate to resolve each candidate policy's attached guardrails. +func (m *managerImpl) loadGuardrailsByID(ctx context.Context, accountID string) (map[string]*types.Guardrail, error) { + guardrails, err := m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, fmt.Errorf("list account guardrails: %w", err) + } + byID := make(map[string]*types.Guardrail, len(guardrails)) + for _, g := range guardrails { + if g != nil { + byID[g.ID] = g + } + } + return byID, nil +} + +// filterModelPermittedPolicies returns the subset of policies whose guardrails +// permit the model. Order is preserved so downstream scoring is unaffected. +func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, model string) []*types.Policy { + out := make([]*types.Policy, 0, len(policies)) + for _, p := range policies { + if policyPermitsModel(p, byID, model) { + out = append(out, p) + } + } + return out +} + +// policyPermitsModel reports whether a policy permits the requested model under +// its attached guardrails. A policy with no allowlist-enabled guardrail is +// unrestricted (permits any model, including an empty/undetermined one). A +// policy with one or more allowlist-enabled guardrails permits the model only +// when it appears in the union of those allowlists; an empty/undetermined model +// never matches a non-empty allowlist, so such a policy fails closed — the same +// contract the proxy guardrail enforces for path-routed providers. +func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model string) bool { + if p == nil { + return false + } + wanted := normaliseModelID(model) + restricted := false + for _, gID := range p.GuardrailIDs { + g, ok := byID[gID] + if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled { + continue + } + restricted = true + if wanted == "" { + continue + } + for _, allowed := range g.Checks.ModelAllowlist.Models { + if normaliseModelID(allowed) == wanted { + return true + } + } + } + return !restricted +} + +// normaliseModelID lowercases and trims a model identifier so the allowlist +// compare is case-insensitive and trim-tolerant. Mirrors the proxy guardrail's +// normaliseModel so both layers agree on what "same model" means. +func normaliseModelID(model string) string { + return strings.ToLower(strings.TrimSpace(model)) +} + +// modelBlockedReason builds the human-readable deny reason for a model-allowlist +// rejection. The model is quoted when known; an undetermined model is reported +// as such so the access log distinguishes "wrong model" from "no model". +func modelBlockedReason(model string) string { + if normaliseModelID(model) == "" { + return "request model could not be determined for the policy allowlist" + } + return fmt.Sprintf("model %q is not permitted by any applicable policy allowlist", model) +} + // candidate is the per-policy intermediate the selector ranks. A // policy that's been exhausted on any enabled cap never makes it // into this slice; the selector's deny envelope carries the latest diff --git a/management/internals/modules/agentnetwork/policyselect_model_test.go b/management/internals/modules/agentnetwork/policyselect_model_test.go new file mode 100644 index 000000000..9097ac150 --- /dev/null +++ b/management/internals/modules/agentnetwork/policyselect_model_test.go @@ -0,0 +1,226 @@ +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/store" +) + +// guardedPolicy builds an enabled, uncapped policy that authorises sourceGroups +// to reach providerID under the given guardrails. Uncapped keeps the selector's +// headroom scoring trivial so these tests isolate the model-allowlist gate. +func guardedPolicy(id, account string, sourceGroups []string, providerID string, guardrailIDs ...string) *types.Policy { + return &types.Policy{ + ID: id, + AccountID: account, + Enabled: true, + SourceGroups: sourceGroups, + DestinationProviderIDs: []string{providerID}, + GuardrailIDs: guardrailIDs, + CreatedAt: time.Now().UTC(), + } +} + +// allowlistGuardrail builds a guardrail whose model allowlist is enabled and +// carries the given models. +func allowlistGuardrail(id, account string, models ...string) *types.Guardrail { + return &types.Guardrail{ + ID: id, + AccountID: account, + Checks: types.GuardrailChecks{ + ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true, Models: models}, + }, + } +} + +func expectPolicies(mockStore *store.MockStore, account string, policies ...*types.Policy) { + mockStore.EXPECT(). + GetAccountAgentNetworkPolicies(gomock.Any(), gomock.Any(), account). + Return(policies, nil) +} + +func expectGuardrails(mockStore *store.MockStore, account string, guardrails ...*types.Guardrail) { + mockStore.EXPECT(). + GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), account). + Return(guardrails, nil) +} + +// TestSelectPolicy_ModelBlockedByAllowlist proves the authoritative allowlist +// decision: a policy authorises the (provider, group) but restricts the model, +// and the requested model isn't on the list, so the request is denied. +func TestSelectPolicy_ModelBlockedByAllowlist(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + UserID: "user-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.False(t, res.Allow, "a model outside the only applicable policy's allowlist must be denied") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode, "deny code must be model_blocked") + assert.NotEmpty(t, res.DenyReason, "deny reason must be populated") +} + +// TestSelectPolicy_ModelAllowedByAllowlist is the allow counterpart: the model +// is on the applicable policy's allowlist, so selection proceeds normally. +func TestSelectPolicy_ModelAllowedByAllowlist(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o", "claude-opus-4")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + UserID: "user-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "a model on the applicable policy's allowlist must be allowed") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} + +// TestSelectPolicy_CaseInsensitiveModelMatch proves the compare tolerates case +// and surrounding whitespace, matching the proxy guardrail's normalisation. +func TestSelectPolicy_CaseInsensitiveModelMatch(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", " GPT-4o ")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "gpt-4o", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "case/whitespace variants must match the allowlist entry") +} + +// TestSelectPolicy_UnguardedPolicyIsUnrestricted is the false-deny fix: when +// two policies authorise the same (provider, group) and one carries no +// guardrail, that un-guardrailed policy makes the request unrestricted — it is +// NOT caught by the other policy's allowlist. +func TestSelectPolicy_UnguardedPolicyIsUnrestricted(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + restricted := guardedPolicy("pol-restricted", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + open := guardedPolicy("pol-open", "acc-1", []string{"grp-eng"}, "prov-1") // no guardrail + expectPolicies(mockStore, "acc-1", restricted, open) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "an un-guardrailed policy for the same (provider, group) must leave the request unrestricted") + assert.Equal(t, "pol-open", res.SelectedPolicyID, "the unrestricted policy must be the one that pays") +} + +// TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups is the false-allow fix: a +// model allowlisted only for grp-b must not be usable by a caller in grp-a, +// even though both groups' policies target the same provider. The selector only +// considers policies applicable to the caller's groups, so grp-b's allowlist +// never enters grp-a's decision. +func TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + polA := guardedPolicy("pol-a", "acc-1", []string{"grp-a"}, "prov-1", "g-a") + polB := guardedPolicy("pol-b", "acc-1", []string{"grp-b"}, "prov-1", "g-b") + expectPolicies(mockStore, "acc-1", polA, polB) + expectGuardrails(mockStore, "acc-1", + allowlistGuardrail("g-a", "acc-1", "gpt-4o"), + allowlistGuardrail("g-b", "acc-1", "claude-opus-4"), + ) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-a"}, + ProviderID: "prov-1", + Model: "claude-opus-4", // only allowed for grp-b + }) + require.NoError(t, err) + assert.False(t, res.Allow, "grp-b's allowlisted model must not leak to a grp-a caller") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode) +} + +// TestSelectPolicy_UndeterminedModelFailsClosed proves the fail-closed contract +// mirrors the proxy: with a restricted applicable policy and an empty model +// (e.g. a path-routed shape the parser couldn't map), the request is denied. +func TestSelectPolicy_UndeterminedModelFailsClosed(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "", // undetermined + }) + require.NoError(t, err) + assert.False(t, res.Allow, "an undetermined model must fail closed against a restricted policy") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode) +} + +// TestSelectPolicy_DisabledAllowlistDoesNotRestrict proves a guardrail whose +// model allowlist is disabled imposes no model restriction, even though the +// policy references it. +func TestSelectPolicy_DisabledAllowlistDoesNotRestrict(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + disabled := &types.Guardrail{ + ID: "g-1", + AccountID: "acc-1", + Checks: types.GuardrailChecks{ + ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}, + }, + } + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", disabled) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "anything-goes", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "a disabled allowlist must not restrict the model") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} diff --git a/management/internals/modules/agentnetwork/synthesizer.go b/management/internals/modules/agentnetwork/synthesizer.go index 95fe91773..1dbb2ca65 100644 --- a/management/internals/modules/agentnetwork/synthesizer.go +++ b/management/internals/modules/agentnetwork/synthesizer.go @@ -235,7 +235,14 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID) applyAccountCollectionControls(&mergedGuardrails, settings) - guardrailJSON, err := marshalGuardrailConfig(mergedGuardrails) + // The proxy guardrail is the fail-closed backstop, scoped per provider so it + // never applies one policy's allowlist to another provider's traffic. The + // authoritative per-policy/group model decision lives in management's + // SelectPolicyForRequest; a provider only lands in this map when EVERY policy + // authorising it restricts models, so an un-guardrailed policy leaves its + // providers unrestricted here. + providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID) + guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture) if err != nil { return nil, err } @@ -780,10 +787,12 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt // guardrailConfig is the JSON shape the proxy-side llm_guardrail // middleware expects. Mirrors the proxy registration documented in -// the management→proxy contract. +// the management→proxy contract. provider_allowlists is keyed by the +// resolved provider id llm_router stamps; a provider absent from the map is +// unrestricted at the proxy layer. type guardrailConfig struct { - ModelAllowlist []string `json:"model_allowlist,omitempty"` - PromptCapture guardrailPromptCapture `json:"prompt_capture"` + ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"` + PromptCapture guardrailPromptCapture `json:"prompt_capture"` } type guardrailPromptCapture struct { @@ -828,12 +837,12 @@ func applyAccountCollectionControls(merged *MergedGuardrails, settings *types.Se merged.PromptCapture.RedactPii = settings.RedactPii || merged.PromptCapture.RedactPii } -func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) { +func marshalGuardrailConfig(providerAllowlists map[string][]string, capture MergedPromptCapture) ([]byte, error) { cfg := guardrailConfig{ - ModelAllowlist: merged.ModelAllowlist, + ProviderAllowlists: providerAllowlists, PromptCapture: guardrailPromptCapture{ - Enabled: merged.PromptCapture.Enabled, - RedactPii: merged.PromptCapture.RedactPii, + Enabled: capture.Enabled, + RedactPii: capture.RedactPii, }, } out, err := json.Marshal(cfg) @@ -843,6 +852,81 @@ func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) { return out, nil } +// buildProviderAllowlists computes the per-provider model allowlist the proxy +// guardrail enforces as a fail-closed backstop. A provider is included ONLY when +// every enabled policy that authorises it restricts models (has at least one +// allowlist-enabled guardrail); its list is then the union of those policies' +// allowed models. If any authorising policy leaves models unrestricted, the +// provider is omitted entirely — the proxy treats an absent provider as +// unrestricted and defers to management's per-policy/group decision, so a +// mixed set of policies can neither leak a model across providers nor wrongly +// block an un-guardrailed policy's traffic. +func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]string { + type providerAcc struct { + models map[string]struct{} + anyUnrestricted bool + } + accs := make(map[string]*providerAcc) + for _, p := range policies { + if p == nil { + continue + } + restricted, models := policyModelAllowlist(p, byID) + for _, providerID := range p.DestinationProviderIDs { + if providerID == "" { + continue + } + acc, ok := accs[providerID] + if !ok { + acc = &providerAcc{models: make(map[string]struct{})} + accs[providerID] = acc + } + if !restricted { + acc.anyUnrestricted = true + continue + } + for _, m := range models { + acc.models[m] = struct{}{} + } + } + } + out := make(map[string][]string, len(accs)) + for providerID, acc := range accs { + if acc.anyUnrestricted { + continue + } + models := make([]string, 0, len(acc.models)) + for m := range acc.models { + models = append(models, m) + } + sort.Strings(models) + out[providerID] = models + } + return out +} + +// policyModelAllowlist reports whether a policy restricts models (has at least +// one allowlist-enabled guardrail) and the union of allowed models across those +// guardrails. Models are returned verbatim; the proxy factory lowercases and +// trims them at decode time, matching the runtime compare. +func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bool, []string) { + restricted := false + var models []string + for _, gID := range p.GuardrailIDs { + g, ok := byID[gID] + if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled { + continue + } + restricted = true + for _, m := range g.Checks.ModelAllowlist.Models { + if m != "" { + models = append(models, m) + } + } + } + return restricted, models +} + // buildAccountService composes the per-account gateway Service. The // target carries the noop placeholder URL — the router middleware // rewrites every request to the matched provider's upstream before the @@ -991,11 +1075,10 @@ func unionSourceGroups(policies []*types.Policy) []string { // expectations and is intentionally distinct from // types.GuardrailChecks so we can evolve either side independently. type MergedGuardrails struct { - ModelAllowlist []string `json:"model_allowlist,omitempty"` - TokenLimits MergedTokenLimits `json:"token_limits"` - Budget MergedBudget `json:"budget"` - PromptCapture MergedPromptCapture `json:"prompt_capture"` - Retention MergedRetention `json:"retention"` + TokenLimits MergedTokenLimits `json:"token_limits"` + Budget MergedBudget `json:"budget"` + PromptCapture MergedPromptCapture `json:"prompt_capture"` + Retention MergedRetention `json:"retention"` } type MergedTokenLimits struct { @@ -1030,59 +1113,38 @@ type MergedRetention struct { Days int `json:"days"` } -// mergeGuardrails computes the effective guardrail spec applied at the -// proxy, given the referencing policies and the account's guardrail -// catalogue. Policy enabled-ness is the caller's responsibility — only -// enabled policies should be passed in. +// mergeGuardrails computes the prompt-capture portion of the effective +// guardrail spec, given the referencing policies and the account's guardrail +// catalogue. Policy enabled-ness is the caller's responsibility — only enabled +// policies should be passed in. // -// Merge rules: -// - Model allowlist: union of allowlists across policies that enable it. -// - Token / Budget: most-restrictive (min of non-zero caps) per window. -// - Prompt capture: enabled if any policy enables it; redact_pii sticks -// if any enabling policy turns it on. -// - Retention: enabled if any enables it; smallest non-zero days wins. +// The model allowlist is NOT merged here: it is enforced per-policy/group in +// management (SelectPolicyForRequest) and shipped to the proxy per-provider via +// buildProviderAllowlists, so a single account-wide union would be both +// incorrect (leaks a model across groups/providers) and redundant. Token, +// budget, and retention likewise live off guardrails now (Policy.Limits and +// account Settings), so prompt capture is all that remains to merge. +// +// Merge rule — prompt capture: enabled if any policy enables it; redact_pii +// sticks if any enabling policy turns it on. func mergeGuardrails(policies []*types.Policy, byID map[string]*types.Guardrail) MergedGuardrails { merged := MergedGuardrails{} - allowlist := make(map[string]struct{}) - allowlistEnabled := false - for _, policy := range policies { for _, gID := range policy.GuardrailIDs { g, ok := byID[gID] if !ok || g == nil { continue } - mergeGuardrail(g, &merged, allowlist, &allowlistEnabled) + mergeGuardrail(g, &merged) } } - - if allowlistEnabled { - merged.ModelAllowlist = make([]string, 0, len(allowlist)) - for m := range allowlist { - merged.ModelAllowlist = append(merged.ModelAllowlist, m) - } - sort.Strings(merged.ModelAllowlist) - } return merged } -// mergeGuardrail folds a single guardrail's enabled checks into the -// running merge: model-allowlist models join the shared set (and flip -// allowlistEnabled), and prompt-capture / redact-pii stick once any -// enabling guardrail turns them on. -// -// TokenLimits, Budget, and Retention have moved off guardrails — token -// and budget caps now live on the Policy itself (Policy.Limits) and -// retention moves to account-level Settings — so they are not merged here. -func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails, allowlist map[string]struct{}, allowlistEnabled *bool) { - if g.Checks.ModelAllowlist.Enabled { - *allowlistEnabled = true - for _, m := range g.Checks.ModelAllowlist.Models { - if m != "" { - allowlist[m] = struct{}{} - } - } - } +// mergeGuardrail folds a single guardrail's prompt-capture settings into the +// running merge: enabled / redact-pii stick once any enabling guardrail turns +// them on. +func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) { if g.Checks.PromptCapture.Enabled { merged.PromptCapture.Enabled = true if g.Checks.PromptCapture.RedactPii { diff --git a/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go b/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go index 8ed4910da..8b13f78ac 100644 --- a/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go @@ -81,8 +81,8 @@ func TestSynthesizeServices_RealStore_PromptCaptureAccountIsSoleControl(t *testi require.Len(t, services, 1) cfg := decodeServiceGuardrailConfig(t, services[0]) - assert.Equal(t, []string{"gpt-5.4"}, cfg.ModelAllowlist, - "model allowlist is a pure policy guardrail and must always reach the config") + assert.Equal(t, map[string][]string{"prov-1": {"gpt-5.4"}}, cfg.ProviderAllowlists, + "model allowlist is a pure policy guardrail and must reach the per-provider config") assert.False(t, cfg.PromptCapture.Enabled, "prompt capture must be off when the account toggle is off, even with a capture-enabled guardrail") } @@ -172,7 +172,7 @@ func TestSynthesizeServices_RealStore_NoGuardrail_CaptureOff(t *testing.T) { require.Len(t, services, 1, "exactly one synth service expected") cfg := decodeServiceGuardrailConfig(t, services[0]) - assert.Empty(t, cfg.ModelAllowlist, "no guardrail → no allowlist") + assert.Empty(t, cfg.ProviderAllowlists, "no guardrail → provider unrestricted (absent from map)") assert.False(t, cfg.PromptCapture.Enabled, "no guardrail → prompt capture off by default") assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default") } diff --git a/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go b/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go new file mode 100644 index 000000000..f4c1ff361 --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go @@ -0,0 +1,77 @@ +package agentnetwork + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" +) + +// policyForProviders builds an enabled policy authorising the given providers +// under the given guardrails (both optional). Groups are irrelevant to +// buildProviderAllowlists, which keys purely on destination provider. +func policyForProviders(id string, guardrailIDs []string, providerIDs ...string) *types.Policy { + return &types.Policy{ + ID: id, + Enabled: true, + DestinationProviderIDs: providerIDs, + GuardrailIDs: guardrailIDs, + } +} + +func TestBuildProviderAllowlists(t *testing.T) { + byID := map[string]*types.Guardrail{ + "g-4o": allowlistGuardrail("g-4o", "acc-1", "gpt-4o"), + "g-opus": allowlistGuardrail("g-opus", "acc-1", "claude-opus-4"), + "g-disabled": {ID: "g-disabled", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}}}, + } + + t.Run("all authorising policies restrict yields per-provider union", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", []string{"g-opus"}, "prov-x"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, map[string][]string{"prov-x": {"claude-opus-4", "gpt-4o"}}, got, + "a provider every policy restricts carries the sorted union of their models") + }) + + t.Run("any un-guardrailed policy leaves the provider unrestricted (omitted)", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", nil, "prov-x"), // no guardrail + } + got := buildProviderAllowlists(policies, byID) + assert.NotContains(t, got, "prov-x", + "a provider reachable by an un-guardrailed policy must be omitted so the proxy treats it as unrestricted") + }) + + t.Run("a disabled allowlist counts as unrestricted", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-disabled"}, "prov-x"), + } + got := buildProviderAllowlists(policies, byID) + assert.NotContains(t, got, "prov-x", + "a policy whose only guardrail has a disabled allowlist is unrestricted") + }) + + t.Run("providers are isolated from one another", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", []string{"g-opus"}, "prov-y"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, []string{"gpt-4o"}, got["prov-x"], "prov-x keeps only its own model") + assert.Equal(t, []string{"claude-opus-4"}, got["prov-y"], "prov-y keeps only its own model") + }) + + t.Run("one policy authorising two providers restricts both", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x", "prov-y"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, []string{"gpt-4o"}, got["prov-x"]) + assert.Equal(t, []string{"gpt-4o"}, got["prov-y"]) + }) +} diff --git a/management/internals/modules/agentnetwork/synthesizer_test.go b/management/internals/modules/agentnetwork/synthesizer_test.go index 206d1d12a..04948fdbe 100644 --- a/management/internals/modules/agentnetwork/synthesizer_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_test.go @@ -1031,8 +1031,14 @@ func TestSynthesizeServices_GuardrailMerge_AllowlistUnion_LimitsRestrictive(t *t var cfg guardrailConfig require.NoError(t, json.Unmarshal(guardrailJSON, &cfg), "guardrail config must unmarshal cleanly") - assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ModelAllowlist, - "model allowlist union must keep both models") + // Both policies restrict the same provider, so the proxy-side per-provider + // backstop carries the union of their models. This coarse union is + // deliberate: it can still let grp-a reach grp-b's model at the proxy layer, + // which management's per-policy/group check (SelectPolicyForRequest) is the + // one that rejects. The backstop only guarantees nothing outside the + // provider's union slips through when management is unreachable. + assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ProviderAllowlists["prov-1"], + "per-provider allowlist union must keep both models") } func TestSynthesizeServices_BackfillsMissingSessionKeys(t *testing.T) { diff --git a/management/internals/shared/grpc/proxy.go b/management/internals/shared/grpc/proxy.go index b289f8c71..8f24de116 100644 --- a/management/internals/shared/grpc/proxy.go +++ b/management/internals/shared/grpc/proxy.go @@ -285,6 +285,7 @@ func (s *ProxyServiceServer) CheckLLMPolicyLimits(ctx context.Context, req *prot UserID: req.GetUserId(), GroupIDs: req.GetGroupIds(), ProviderID: req.GetProviderId(), + Model: req.GetModel(), }) if err != nil { log.WithContext(ctx).Errorf("select policy for request: %v", err) diff --git a/proxy/internal/middleware/builtin/llm_guardrail/factory.go b/proxy/internal/middleware/builtin/llm_guardrail/factory.go index 6dd2a8e8d..3409ce0c0 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/factory.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/factory.go @@ -10,11 +10,22 @@ import ( ) // Config is the JSON-decoded shape accepted by the factory. The -// runtime path consumes the normalised allowlist; raw config is not +// runtime path consumes the normalised allowlists; raw config is not // retained beyond construction. type Config struct { - ModelAllowlist []string `json:"model_allowlist"` - PromptCapture PromptCapture `json:"prompt_capture"` + // ProviderAllowlists maps a resolved provider id (the value llm_router + // stamps as KeyLLMResolvedProviderID) to that provider's model allowlist. A + // provider present here is restricted to the listed models; a provider + // absent is unrestricted. The synthesiser only lists a provider when EVERY + // policy authorising it enables an allowlist, so a provider reachable by any + // un-guardrailed policy is intentionally absent (unrestricted) here and the + // precise per-policy/group decision is left to management. Keeping the gate + // per-provider — rather than one account-wide union — is what stops a model + // allowlisted for one provider from leaking onto another and stops an + // un-guardrailed policy's traffic from being blocked by an unrelated + // policy's allowlist. + ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"` + PromptCapture PromptCapture `json:"prompt_capture"` } // PromptCapture toggles the optional prompt capture + redaction step @@ -55,20 +66,28 @@ func isEmptyJSON(raw []byte) bool { } // normaliseConfig lowercases and trims allowlist entries so the runtime -// match is case-insensitive. Empty entries are dropped. +// match is case-insensitive. Empty entries are dropped. A provider whose +// entries all drop out keeps an empty (non-nil) list — an allowlist that +// permits nothing — which is the intended "deny every model" for that +// provider, distinct from the provider being absent (unrestricted). func normaliseConfig(cfg Config) Config { - if len(cfg.ModelAllowlist) == 0 { + if len(cfg.ProviderAllowlists) == 0 { + cfg.ProviderAllowlists = nil return cfg } - cleaned := make([]string, 0, len(cfg.ModelAllowlist)) - for _, entry := range cfg.ModelAllowlist { - n := normaliseModel(entry) - if n == "" { - continue + cleaned := make(map[string][]string, len(cfg.ProviderAllowlists)) + for provider, models := range cfg.ProviderAllowlists { + list := make([]string, 0, len(models)) + for _, entry := range models { + n := normaliseModel(entry) + if n == "" { + continue + } + list = append(list, n) } - cleaned = append(cleaned, n) + cleaned[provider] = list } - cfg.ModelAllowlist = cleaned + cfg.ProviderAllowlists = cleaned return cfg } diff --git a/proxy/internal/middleware/builtin/llm_guardrail/middleware.go b/proxy/internal/middleware/builtin/llm_guardrail/middleware.go index eded877ac..671906763 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/middleware.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/middleware.go @@ -83,8 +83,9 @@ func (m *Middleware) MutationsSupported() bool { return false } // prompt capture only affects the metadata emitted alongside an allow. func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) { model, modelPresent := lookupMetadata(in.Metadata, middleware.KeyLLMModel) + providerID, _ := lookupMetadata(in.Metadata, middleware.KeyLLMResolvedProviderID) - if denial := m.evaluateAllowlist(model, modelPresent); denial != nil { + if denial := m.evaluateAllowlist(providerID, model, modelPresent); denial != nil { return denial, nil } @@ -110,20 +111,38 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar // is a no-op. func (m *Middleware) Close() error { return nil } -// evaluateAllowlist returns a deny Output when the configured allowlist -// rejects the model. A nil return means the request should proceed. -func (m *Middleware) evaluateAllowlist(model string, modelPresent bool) *middleware.Output { - if len(m.cfg.ModelAllowlist) == 0 { +// evaluateAllowlist returns a deny Output when the allowlist for the request's +// resolved provider rejects the model. A nil return means the request should +// proceed. The allowlist is scoped to the provider llm_router resolved: a +// provider absent from the config is unrestricted, so an un-guardrailed policy's +// traffic is never caught by an unrelated provider's allowlist. +func (m *Middleware) evaluateAllowlist(providerID, model string, modelPresent bool) *middleware.Output { + if len(m.cfg.ProviderAllowlists) == 0 { return nil } - // Fail closed: with an allowlist configured, a request whose model the - // upstream parser could not extract (absent or empty) must be denied rather - // than allowed. This is what enforces the allowlist for URL/path-routed - // providers (Bedrock, Vertex, ...) whose model lives outside the JSON body. + // Restrictions exist for this account but the resolved provider is unknown, + // so we cannot tell whether this request is bound for a restricted provider. + // Fail closed rather than wave it through. In practice llm_router always + // stamps the resolved provider before the guardrail runs (or denies first), + // so this is a defensive guard, not the common path. + if providerID == "" { + return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown) + } + allowlist, restricted := m.cfg.ProviderAllowlists[providerID] + if !restricted { + // This provider has no allowlist (some authorising policy left it + // unrestricted); management owns any per-policy/group decision. + return nil + } + // Fail closed: with an allowlist in effect for this provider, a request whose + // model the upstream parser could not extract (absent or empty) must be + // denied rather than allowed. This is what enforces the allowlist for + // URL/path-routed providers (Bedrock, Vertex, ...) whose model lives outside + // the JSON body. if !modelPresent || normaliseModel(model) == "" { return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown) } - if m.modelInAllowlist(model) { + if modelInAllowlist(allowlist, model) { return nil } return denyModel(model, denyCodeModel, denyMessageModel, denyReasonModel) @@ -151,14 +170,15 @@ func denyModel(model, code, message, reason string) *middleware.Output { } } -// modelInAllowlist reports whether the model matches any allowlist -// entry under the case-insensitive, trim-tolerant comparison rule. -func (m *Middleware) modelInAllowlist(model string) bool { +// modelInAllowlist reports whether the model matches any entry in the supplied +// (already-normalised) allowlist under the case-insensitive, trim-tolerant +// comparison rule. +func modelInAllowlist(allowlist []string, model string) bool { normalised := normaliseModel(model) if normalised == "" { return false } - for _, allowed := range m.cfg.ModelAllowlist { + for _, allowed := range allowlist { if allowed == normalised { return true } diff --git a/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go b/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go index cd7e256dd..b27974b50 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go @@ -26,6 +26,25 @@ func newInput(meta ...middleware.KV) *middleware.Input { return &middleware.Input{Slot: middleware.SlotOnRequest, Metadata: meta} } +const ( + testProvider = "prov-1" + otherProvider = "prov-2" +) + +// providerCfg builds a Config restricting testProvider to the given models. +func providerCfg(models ...string) Config { + return Config{ProviderAllowlists: map[string][]string{testProvider: models}} +} + +// newInputProvider builds an input that carries a resolved provider id (as +// llm_router would stamp) plus any extra metadata. +func newInputProvider(provider string, meta ...middleware.KV) *middleware.Input { + all := make([]middleware.KV, 0, len(meta)+1) + all = append(all, middleware.KV{Key: middleware.KeyLLMResolvedProviderID, Value: provider}) + all = append(all, meta...) + return &middleware.Input{Slot: middleware.SlotOnRequest, Metadata: all} +} + func TestMiddlewareIdentity(t *testing.T) { mw := New(Config{}) assert.Equal(t, ID, mw.ID(), "middleware ID must be llm_guardrail") @@ -47,12 +66,12 @@ func TestMiddlewareIdentity(t *testing.T) { func TestAllowlistEmptyAllowsAnyModel(t *testing.T) { mw := New(Config{}) - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionAllow, out.Decision, "empty allowlist must allow any model") + assert.Equal(t, middleware.DecisionAllow, out.Decision, "no provider allowlists must allow any model") v, ok := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision) require.True(t, ok, "decision metadata must be emitted") assert.Equal(t, "allow", v, "decision must be allow") @@ -62,8 +81,8 @@ func TestAllowlistEmptyAllowsAnyModel(t *testing.T) { } func TestAllowlistMatchAllows(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{"gpt-4o", "claude-opus-4"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o", "claude-opus-4")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) @@ -71,8 +90,8 @@ func TestAllowlistMatchAllows(t *testing.T) { } func TestAllowlistMissDenies(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, )) require.NoError(t, err) @@ -91,10 +110,10 @@ func TestAllowlistMissDenies(t *testing.T) { } func TestAllowlistCaseInsensitive(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{" GPT-4o ", "Claude-OPUS-4"}}) + mw := New(providerCfg(" GPT-4o ", "Claude-OPUS-4")) cases := []string{"gpt-4o", "GPT-4O", " claude-opus-4 "} for _, model := range cases { - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: model}, )) require.NoError(t, err) @@ -103,14 +122,15 @@ func TestAllowlistCaseInsensitive(t *testing.T) { } func TestAllowlistMissingModelKeyDenies(t *testing.T) { - // Fail closed: with an allowlist configured, a request whose model the - // parser could not extract (URL/path-routed providers such as Bedrock or - // Vertex whose shape wasn't recognised) must be denied, not allowed. - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput()) + // Fail closed: with an allowlist in effect for the resolved provider, a + // request whose model the parser could not extract (URL/path-routed + // providers such as Bedrock or Vertex whose shape wasn't recognised) must be + // denied, not allowed. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider)) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionDeny, out.Decision, "absent model must be denied when an allowlist is set") + assert.Equal(t, middleware.DecisionDeny, out.Decision, "absent model must be denied when the provider is restricted") assert.Equal(t, 403, out.DenyStatus, "deny status must be 403") require.NotNil(t, out.DenyReason, "deny reason must be populated") assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") @@ -122,26 +142,72 @@ func TestAllowlistMissingModelKeyDenies(t *testing.T) { func TestAllowlistEmptyModelValueDenies(t *testing.T) { // A present-but-empty model is as undeterminable as an absent one. - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: " "}, )) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionDeny, out.Decision, "empty model must be denied when an allowlist is set") + assert.Equal(t, middleware.DecisionDeny, out.Decision, "empty model must be denied when the provider is restricted") require.NotNil(t, out.DenyReason, "deny reason must be populated") assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") } func TestAllowlistEmptyListAllowsMissingModel(t *testing.T) { - // Without an allowlist there is nothing to enforce, so a missing model is - // still allowed — the fail-closed rule only applies when a list is set. + // Without any provider allowlists there is nothing to enforce, so a missing + // model is still allowed — the fail-closed rule only applies when a + // restriction is in effect. mw := New(Config{}) out, err := mw.Invoke(context.Background(), newInput()) require.NoError(t, err) assert.Equal(t, middleware.DecisionAllow, out.Decision, "no allowlist must allow even without a model") } +func TestUnrestrictedProviderAllowsAnyModel(t *testing.T) { + // testProvider is restricted, but the request resolved to otherProvider, + // which carries no allowlist. Its traffic must not be caught by the other + // provider's restriction — this is the cross-provider-leak / false-deny + // guard. Management owns any per-policy/group decision for otherProvider. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(otherProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, + )) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionAllow, out.Decision, "an unrestricted provider must not inherit another provider's allowlist") +} + +func TestPerProviderAllowlistsAreIsolated(t *testing.T) { + // gpt-4o is allowed only on testProvider; claude-opus-4 only on + // otherProvider. A model allowlisted for one provider must not be usable on + // the other — the fail-closed layer never unions allowlists across providers. + mw := New(Config{ProviderAllowlists: map[string][]string{ + testProvider: {"gpt-4o"}, + otherProvider: {"claude-opus-4"}, + }}) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, + )) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "claude-opus-4 is allowed only on otherProvider, not testProvider") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_blocked", out.DenyReason.Code, "cross-provider model must be blocked, not model_unknown") +} + +func TestRestrictionsButNoResolvedProviderFailsClosed(t *testing.T) { + // Restrictions exist for the account but the resolved provider id is absent, + // so the request cannot be scoped to a provider. Fail closed rather than + // wave it through. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInput( + middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + )) + require.NoError(t, err) + require.NotNil(t, out) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "missing resolved provider must fail closed when restrictions exist") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") +} + func TestPromptCaptureDisabledEmitsNoPrompt(t *testing.T) { mw := New(Config{}) out, err := mw.Invoke(context.Background(), newInput( @@ -217,8 +283,8 @@ func TestFactoryAcceptsZeroConfigs(t *testing.T) { func TestFactoryDecodesValidConfig(t *testing.T) { cfg := Config{ - ModelAllowlist: []string{"gpt-4o"}, - PromptCapture: PromptCapture{Enabled: true, RedactPii: true}, + ProviderAllowlists: map[string][]string{testProvider: {"gpt-4o"}}, + PromptCapture: PromptCapture{Enabled: true, RedactPii: true}, } raw, err := json.Marshal(cfg) require.NoError(t, err, "marshalling test config must succeed") @@ -234,15 +300,15 @@ func TestFactoryRejectsMalformedJSON(t *testing.T) { } func TestFactoryNormalisesAllowlist(t *testing.T) { - raw := []byte(`{"model_allowlist":[" GPT-4o ","",""," Claude-3 "]}`) + raw := []byte(`{"provider_allowlists":{"prov-1":[" GPT-4o ","",""," Claude-3 "]}}`) mw, err := Factory{}.New(raw) require.NoError(t, err) - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) assert.Equal(t, middleware.DecisionAllow, out.Decision, "factory must lowercase + trim allowlist entries") - out2, err := mw.Invoke(context.Background(), newInput( + out2, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-3"}, )) require.NoError(t, err) diff --git a/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go b/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go index 0074411cc..1eae26d73 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go @@ -25,10 +25,17 @@ func runParserGuardrail(t *testing.T, url string, body []byte, allowlist []strin }) require.NoError(t, err, "parser must not error") - guard := llm_guardrail.New(llm_guardrail.Config{ModelAllowlist: allowlist}) + const providerID = "prov-under-test" + guard := llm_guardrail.New(llm_guardrail.Config{ + ProviderAllowlists: map[string][]string{providerID: allowlist}, + }) + // The real chain has llm_router stamp the resolved provider id before the + // guardrail runs; the parser doesn't, so add it here so the guardrail can + // scope the allowlist to this provider. + meta := append([]middleware.KV{{Key: middleware.KeyLLMResolvedProviderID, Value: providerID}}, parsed.Metadata...) out, err := guard.Invoke(context.Background(), &middleware.Input{ Slot: middleware.SlotOnRequest, - Metadata: parsed.Metadata, + Metadata: meta, }) require.NoError(t, err, "guardrail must not error") require.NotNil(t, out, "guardrail must return an output")