mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-03 20:19:07 +02:00
agentnetwork: scope model allowlist per policy/group and provider
The model-allowlist guardrail was merged into one account-wide union
and enforced flat on every request, independent of which policy,
group, or provider actually authorised it. With multiple policies —
especially a mix of guardrailed and un-guardrailed ones — this
produced two defects:
- false-allow: a model allowlisted for one group/provider leaked to
any caller, because the flat union ignored which policy authorised
the request; and
- false-deny: a policy with no guardrail (intended unrestricted) had
its traffic blocked by an unrelated policy's allowlist.
Make enforcement policy/group-aware, mirroring how llm_limit_check
already resolves caps:
- Management (SelectPolicyForRequest) is authoritative. It now uses
the request model (already carried on the wire in
CheckLLMPolicyLimitsRequest.model, previously ignored) to keep only
the applicable policies whose guardrails permit the model; a policy
with no allowlist-enabled guardrail is unrestricted. When policies
govern the (provider, groups) but none permits the model it denies
with llm_policy.model_blocked. No proto change — the field already
existed and the proxy already sends it.
- The proxy llm_guardrail becomes a per-provider fail-closed backstop.
The synthesiser emits a per-provider allowlist only when every
policy authorising that provider restricts models; a provider
reachable by any un-guardrailed policy is omitted (unrestricted).
The middleware keys the allowlist by the resolved provider id and
keeps the unknown/undetermined-model fail-closed behaviour so
path-routed providers stay enforced when management is unreachable.
Unit tests cover the selector gate (block, allow, cross-group no-leak,
un-guardrailed policy unrestricted, undetermined-model fail-closed) and
the per-provider synthesis; an e2e proves all three properties end to
end over the tunnel. Docs updated for the new config shape and the
two-layer enforcement.
This commit is contained in:
@@ -109,7 +109,7 @@ sequenceDiagram
|
|||||||
Chk->>Inj: continue
|
Chk->>Inj: continue
|
||||||
Inj->>Inj: inject NetBird identity headers per provider config
|
Inj->>Inj: inject NetBird identity headers per provider config
|
||||||
Inj->>Grd: continue
|
Inj->>Grd: continue
|
||||||
Grd->>Grd: enforce model allowlist
|
Grd->>Grd: enforce per-provider allowlist (fail-closed backstop)
|
||||||
Grd->>Up: forward (over WireGuard)
|
Grd->>Up: forward (over WireGuard)
|
||||||
Up-->>Resp: response (JSON or SSE stream)
|
Up-->>Resp: response (JSON or SSE stream)
|
||||||
Resp->>Resp: parse usage tokens, completion
|
Resp->>Resp: parse usage tokens, completion
|
||||||
@@ -135,6 +135,15 @@ sequenceDiagram
|
|||||||
(`redact_pii = settings.RedactPii`). Phones, emails, credit cards,
|
(`redact_pii = settings.RedactPii`). Phones, emails, credit cards,
|
||||||
PII names — see `redact.go` for the full set. See
|
PII names — see `redact.go` for the full set. See
|
||||||
[`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md).
|
[`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
|
- SSE streaming requires special handling on the response side; the
|
||||||
parser must handle partial chunks without buffering the whole
|
parser must handle partial chunks without buffering the whole
|
||||||
stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md).
|
stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.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 | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** |
|
||||||
| on_request | 2 | `llm_limit_check` | `{}` | – |
|
| on_request | 2 | `llm_limit_check` | `{}` | – |
|
||||||
| on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** |
|
| 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 | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | – |
|
||||||
| on_response | 6 | `cost_meter` | `{}` | – |
|
| on_response | 6 | `cost_meter` | `{}` | – |
|
||||||
| on_response | 7 | `llm_response_parser` | `{"capture_completion": <bool>, "redact_pii"?: true}` | – |
|
| on_response | 7 | `llm_response_parser` | `{"capture_completion": <bool>, "redact_pii"?: true}` | – |
|
||||||
|
|||||||
@@ -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_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_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` |
|
||||||
| `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` |
|
| `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}` |
|
| `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` |
|
||||||
| `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) |
|
| `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) |
|
||||||
| `llm_limit_record` | `{}` — same pattern as `llm_limit_check` |
|
| `llm_limit_record` | `{}` — same pattern as `llm_limit_check` |
|
||||||
|
|||||||
@@ -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()))
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -86,12 +86,20 @@ type Manager interface {
|
|||||||
|
|
||||||
// PolicySelectionInput is the per-request selection envelope. The
|
// PolicySelectionInput is the per-request selection envelope. The
|
||||||
// proxy populates it from CapturedData (account, user, groups) plus
|
// 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 {
|
type PolicySelectionInput struct {
|
||||||
AccountID string
|
AccountID string
|
||||||
UserID string
|
UserID string
|
||||||
GroupIDs []string
|
GroupIDs []string
|
||||||
ProviderID 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
|
// PolicySelectionResult names the policy that "pays" for this request
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||||
@@ -35,6 +36,11 @@ const (
|
|||||||
denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded"
|
denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded"
|
||||||
//nolint:gosec // account deny code label, not a credential
|
//nolint:gosec // account deny code label, not a credential
|
||||||
denyCodeAccountBudgetCapExceeded = "llm_account.budget_cap_exceeded"
|
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
|
// 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)
|
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
|
// Prefetch every consumption counter the ceiling + candidate policies will
|
||||||
// read, in a single store round-trip, then score against the cache.
|
// read, in a single store round-trip, then score against the cache.
|
||||||
cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now)
|
cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now)
|
||||||
@@ -250,6 +284,93 @@ func filterApplicablePolicies(policies []*types.Policy, in PolicySelectionInput)
|
|||||||
return out
|
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
|
// candidate is the per-policy intermediate the selector ranks. A
|
||||||
// policy that's been exhausted on any enabled cap never makes it
|
// policy that's been exhausted on any enabled cap never makes it
|
||||||
// into this slice; the selector's deny envelope carries the latest
|
// into this slice; the selector's deny envelope carries the latest
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -235,7 +235,14 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
|||||||
|
|
||||||
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
|
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
|
||||||
applyAccountCollectionControls(&mergedGuardrails, settings)
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -780,10 +787,12 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt
|
|||||||
|
|
||||||
// guardrailConfig is the JSON shape the proxy-side llm_guardrail
|
// guardrailConfig is the JSON shape the proxy-side llm_guardrail
|
||||||
// middleware expects. Mirrors the proxy registration documented in
|
// 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 {
|
type guardrailConfig struct {
|
||||||
ModelAllowlist []string `json:"model_allowlist,omitempty"`
|
ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"`
|
||||||
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
|
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type guardrailPromptCapture struct {
|
type guardrailPromptCapture struct {
|
||||||
@@ -828,12 +837,12 @@ func applyAccountCollectionControls(merged *MergedGuardrails, settings *types.Se
|
|||||||
merged.PromptCapture.RedactPii = settings.RedactPii || merged.PromptCapture.RedactPii
|
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{
|
cfg := guardrailConfig{
|
||||||
ModelAllowlist: merged.ModelAllowlist,
|
ProviderAllowlists: providerAllowlists,
|
||||||
PromptCapture: guardrailPromptCapture{
|
PromptCapture: guardrailPromptCapture{
|
||||||
Enabled: merged.PromptCapture.Enabled,
|
Enabled: capture.Enabled,
|
||||||
RedactPii: merged.PromptCapture.RedactPii,
|
RedactPii: capture.RedactPii,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
out, err := json.Marshal(cfg)
|
out, err := json.Marshal(cfg)
|
||||||
@@ -843,6 +852,81 @@ func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) {
|
|||||||
return out, nil
|
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
|
// buildAccountService composes the per-account gateway Service. The
|
||||||
// target carries the noop placeholder URL — the router middleware
|
// target carries the noop placeholder URL — the router middleware
|
||||||
// rewrites every request to the matched provider's upstream before the
|
// 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
|
// expectations and is intentionally distinct from
|
||||||
// types.GuardrailChecks so we can evolve either side independently.
|
// types.GuardrailChecks so we can evolve either side independently.
|
||||||
type MergedGuardrails struct {
|
type MergedGuardrails struct {
|
||||||
ModelAllowlist []string `json:"model_allowlist,omitempty"`
|
TokenLimits MergedTokenLimits `json:"token_limits"`
|
||||||
TokenLimits MergedTokenLimits `json:"token_limits"`
|
Budget MergedBudget `json:"budget"`
|
||||||
Budget MergedBudget `json:"budget"`
|
PromptCapture MergedPromptCapture `json:"prompt_capture"`
|
||||||
PromptCapture MergedPromptCapture `json:"prompt_capture"`
|
Retention MergedRetention `json:"retention"`
|
||||||
Retention MergedRetention `json:"retention"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type MergedTokenLimits struct {
|
type MergedTokenLimits struct {
|
||||||
@@ -1030,59 +1113,38 @@ type MergedRetention struct {
|
|||||||
Days int `json:"days"`
|
Days int `json:"days"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// mergeGuardrails computes the effective guardrail spec applied at the
|
// mergeGuardrails computes the prompt-capture portion of the effective
|
||||||
// proxy, given the referencing policies and the account's guardrail
|
// guardrail spec, given the referencing policies and the account's guardrail
|
||||||
// catalogue. Policy enabled-ness is the caller's responsibility — only
|
// catalogue. Policy enabled-ness is the caller's responsibility — only enabled
|
||||||
// enabled policies should be passed in.
|
// policies should be passed in.
|
||||||
//
|
//
|
||||||
// Merge rules:
|
// The model allowlist is NOT merged here: it is enforced per-policy/group in
|
||||||
// - Model allowlist: union of allowlists across policies that enable it.
|
// management (SelectPolicyForRequest) and shipped to the proxy per-provider via
|
||||||
// - Token / Budget: most-restrictive (min of non-zero caps) per window.
|
// buildProviderAllowlists, so a single account-wide union would be both
|
||||||
// - Prompt capture: enabled if any policy enables it; redact_pii sticks
|
// incorrect (leaks a model across groups/providers) and redundant. Token,
|
||||||
// if any enabling policy turns it on.
|
// budget, and retention likewise live off guardrails now (Policy.Limits and
|
||||||
// - Retention: enabled if any enables it; smallest non-zero days wins.
|
// 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 {
|
func mergeGuardrails(policies []*types.Policy, byID map[string]*types.Guardrail) MergedGuardrails {
|
||||||
merged := MergedGuardrails{}
|
merged := MergedGuardrails{}
|
||||||
allowlist := make(map[string]struct{})
|
|
||||||
allowlistEnabled := false
|
|
||||||
|
|
||||||
for _, policy := range policies {
|
for _, policy := range policies {
|
||||||
for _, gID := range policy.GuardrailIDs {
|
for _, gID := range policy.GuardrailIDs {
|
||||||
g, ok := byID[gID]
|
g, ok := byID[gID]
|
||||||
if !ok || g == nil {
|
if !ok || g == nil {
|
||||||
continue
|
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
|
return merged
|
||||||
}
|
}
|
||||||
|
|
||||||
// mergeGuardrail folds a single guardrail's enabled checks into the
|
// mergeGuardrail folds a single guardrail's prompt-capture settings into the
|
||||||
// running merge: model-allowlist models join the shared set (and flip
|
// running merge: enabled / redact-pii stick once any enabling guardrail turns
|
||||||
// allowlistEnabled), and prompt-capture / redact-pii stick once any
|
// them on.
|
||||||
// enabling guardrail turns them on.
|
func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) {
|
||||||
//
|
|
||||||
// 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{}{}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if g.Checks.PromptCapture.Enabled {
|
if g.Checks.PromptCapture.Enabled {
|
||||||
merged.PromptCapture.Enabled = true
|
merged.PromptCapture.Enabled = true
|
||||||
if g.Checks.PromptCapture.RedactPii {
|
if g.Checks.PromptCapture.RedactPii {
|
||||||
|
|||||||
@@ -81,8 +81,8 @@ func TestSynthesizeServices_RealStore_PromptCaptureAccountIsSoleControl(t *testi
|
|||||||
require.Len(t, services, 1)
|
require.Len(t, services, 1)
|
||||||
|
|
||||||
cfg := decodeServiceGuardrailConfig(t, services[0])
|
cfg := decodeServiceGuardrailConfig(t, services[0])
|
||||||
assert.Equal(t, []string{"gpt-5.4"}, cfg.ModelAllowlist,
|
assert.Equal(t, map[string][]string{"prov-1": {"gpt-5.4"}}, cfg.ProviderAllowlists,
|
||||||
"model allowlist is a pure policy guardrail and must always reach the config")
|
"model allowlist is a pure policy guardrail and must reach the per-provider config")
|
||||||
assert.False(t, cfg.PromptCapture.Enabled,
|
assert.False(t, cfg.PromptCapture.Enabled,
|
||||||
"prompt capture must be off when the account toggle is off, even with a capture-enabled guardrail")
|
"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")
|
require.Len(t, services, 1, "exactly one synth service expected")
|
||||||
|
|
||||||
cfg := decodeServiceGuardrailConfig(t, services[0])
|
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.Enabled, "no guardrail → prompt capture off by default")
|
||||||
assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default")
|
assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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"])
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -1031,8 +1031,14 @@ func TestSynthesizeServices_GuardrailMerge_AllowlistUnion_LimitsRestrictive(t *t
|
|||||||
|
|
||||||
var cfg guardrailConfig
|
var cfg guardrailConfig
|
||||||
require.NoError(t, json.Unmarshal(guardrailJSON, &cfg), "guardrail config must unmarshal cleanly")
|
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,
|
// Both policies restrict the same provider, so the proxy-side per-provider
|
||||||
"model allowlist union must keep both models")
|
// 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) {
|
func TestSynthesizeServices_BackfillsMissingSessionKeys(t *testing.T) {
|
||||||
|
|||||||
@@ -285,6 +285,7 @@ func (s *ProxyServiceServer) CheckLLMPolicyLimits(ctx context.Context, req *prot
|
|||||||
UserID: req.GetUserId(),
|
UserID: req.GetUserId(),
|
||||||
GroupIDs: req.GetGroupIds(),
|
GroupIDs: req.GetGroupIds(),
|
||||||
ProviderID: req.GetProviderId(),
|
ProviderID: req.GetProviderId(),
|
||||||
|
Model: req.GetModel(),
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.WithContext(ctx).Errorf("select policy for request: %v", err)
|
log.WithContext(ctx).Errorf("select policy for request: %v", err)
|
||||||
|
|||||||
@@ -10,11 +10,22 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// Config is the JSON-decoded shape accepted by the factory. The
|
// 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.
|
// retained beyond construction.
|
||||||
type Config struct {
|
type Config struct {
|
||||||
ModelAllowlist []string `json:"model_allowlist"`
|
// ProviderAllowlists maps a resolved provider id (the value llm_router
|
||||||
PromptCapture PromptCapture `json:"prompt_capture"`
|
// 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
|
// 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
|
// 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 {
|
func normaliseConfig(cfg Config) Config {
|
||||||
if len(cfg.ModelAllowlist) == 0 {
|
if len(cfg.ProviderAllowlists) == 0 {
|
||||||
|
cfg.ProviderAllowlists = nil
|
||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
cleaned := make([]string, 0, len(cfg.ModelAllowlist))
|
cleaned := make(map[string][]string, len(cfg.ProviderAllowlists))
|
||||||
for _, entry := range cfg.ModelAllowlist {
|
for provider, models := range cfg.ProviderAllowlists {
|
||||||
n := normaliseModel(entry)
|
list := make([]string, 0, len(models))
|
||||||
if n == "" {
|
for _, entry := range models {
|
||||||
continue
|
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
|
return cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -83,8 +83,9 @@ func (m *Middleware) MutationsSupported() bool { return false }
|
|||||||
// prompt capture only affects the metadata emitted alongside an allow.
|
// prompt capture only affects the metadata emitted alongside an allow.
|
||||||
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
|
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
|
||||||
model, modelPresent := lookupMetadata(in.Metadata, middleware.KeyLLMModel)
|
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
|
return denial, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -110,20 +111,38 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
|||||||
// is a no-op.
|
// is a no-op.
|
||||||
func (m *Middleware) Close() error { return nil }
|
func (m *Middleware) Close() error { return nil }
|
||||||
|
|
||||||
// evaluateAllowlist returns a deny Output when the configured allowlist
|
// evaluateAllowlist returns a deny Output when the allowlist for the request's
|
||||||
// rejects the model. A nil return means the request should proceed.
|
// resolved provider rejects the model. A nil return means the request should
|
||||||
func (m *Middleware) evaluateAllowlist(model string, modelPresent bool) *middleware.Output {
|
// proceed. The allowlist is scoped to the provider llm_router resolved: a
|
||||||
if len(m.cfg.ModelAllowlist) == 0 {
|
// 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
|
return nil
|
||||||
}
|
}
|
||||||
// Fail closed: with an allowlist configured, a request whose model the
|
// Restrictions exist for this account but the resolved provider is unknown,
|
||||||
// upstream parser could not extract (absent or empty) must be denied rather
|
// so we cannot tell whether this request is bound for a restricted provider.
|
||||||
// than allowed. This is what enforces the allowlist for URL/path-routed
|
// Fail closed rather than wave it through. In practice llm_router always
|
||||||
// providers (Bedrock, Vertex, ...) whose model lives outside the JSON body.
|
// 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) == "" {
|
if !modelPresent || normaliseModel(model) == "" {
|
||||||
return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown)
|
return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown)
|
||||||
}
|
}
|
||||||
if m.modelInAllowlist(model) {
|
if modelInAllowlist(allowlist, model) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return denyModel(model, denyCodeModel, denyMessageModel, denyReasonModel)
|
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
|
// modelInAllowlist reports whether the model matches any entry in the supplied
|
||||||
// entry under the case-insensitive, trim-tolerant comparison rule.
|
// (already-normalised) allowlist under the case-insensitive, trim-tolerant
|
||||||
func (m *Middleware) modelInAllowlist(model string) bool {
|
// comparison rule.
|
||||||
|
func modelInAllowlist(allowlist []string, model string) bool {
|
||||||
normalised := normaliseModel(model)
|
normalised := normaliseModel(model)
|
||||||
if normalised == "" {
|
if normalised == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
for _, allowed := range m.cfg.ModelAllowlist {
|
for _, allowed := range allowlist {
|
||||||
if allowed == normalised {
|
if allowed == normalised {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -26,6 +26,25 @@ func newInput(meta ...middleware.KV) *middleware.Input {
|
|||||||
return &middleware.Input{Slot: middleware.SlotOnRequest, Metadata: meta}
|
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) {
|
func TestMiddlewareIdentity(t *testing.T) {
|
||||||
mw := New(Config{})
|
mw := New(Config{})
|
||||||
assert.Equal(t, ID, mw.ID(), "middleware ID must be llm_guardrail")
|
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) {
|
func TestAllowlistEmptyAllowsAnyModel(t *testing.T) {
|
||||||
mw := New(Config{})
|
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"},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NotNil(t, out)
|
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)
|
v, ok := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision)
|
||||||
require.True(t, ok, "decision metadata must be emitted")
|
require.True(t, ok, "decision metadata must be emitted")
|
||||||
assert.Equal(t, "allow", v, "decision must be allow")
|
assert.Equal(t, "allow", v, "decision must be allow")
|
||||||
@@ -62,8 +81,8 @@ func TestAllowlistEmptyAllowsAnyModel(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowlistMatchAllows(t *testing.T) {
|
func TestAllowlistMatchAllows(t *testing.T) {
|
||||||
mw := New(Config{ModelAllowlist: []string{"gpt-4o", "claude-opus-4"}})
|
mw := New(providerCfg("gpt-4o", "claude-opus-4"))
|
||||||
out, err := mw.Invoke(context.Background(), newInput(
|
out, err := mw.Invoke(context.Background(), newInputProvider(testProvider,
|
||||||
middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -71,8 +90,8 @@ func TestAllowlistMatchAllows(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowlistMissDenies(t *testing.T) {
|
func TestAllowlistMissDenies(t *testing.T) {
|
||||||
mw := New(Config{ModelAllowlist: []string{"gpt-4o"}})
|
mw := New(providerCfg("gpt-4o"))
|
||||||
out, err := mw.Invoke(context.Background(), newInput(
|
out, err := mw.Invoke(context.Background(), newInputProvider(testProvider,
|
||||||
middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -91,10 +110,10 @@ func TestAllowlistMissDenies(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowlistCaseInsensitive(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 "}
|
cases := []string{"gpt-4o", "GPT-4O", " claude-opus-4 "}
|
||||||
for _, model := range cases {
|
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},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: model},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -103,14 +122,15 @@ func TestAllowlistCaseInsensitive(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowlistMissingModelKeyDenies(t *testing.T) {
|
func TestAllowlistMissingModelKeyDenies(t *testing.T) {
|
||||||
// Fail closed: with an allowlist configured, a request whose model the
|
// Fail closed: with an allowlist in effect for the resolved provider, a
|
||||||
// parser could not extract (URL/path-routed providers such as Bedrock or
|
// request whose model the parser could not extract (URL/path-routed
|
||||||
// Vertex whose shape wasn't recognised) must be denied, not allowed.
|
// providers such as Bedrock or Vertex whose shape wasn't recognised) must be
|
||||||
mw := New(Config{ModelAllowlist: []string{"gpt-4o"}})
|
// denied, not allowed.
|
||||||
out, err := mw.Invoke(context.Background(), newInput())
|
mw := New(providerCfg("gpt-4o"))
|
||||||
|
out, err := mw.Invoke(context.Background(), newInputProvider(testProvider))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NotNil(t, out)
|
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")
|
assert.Equal(t, 403, out.DenyStatus, "deny status must be 403")
|
||||||
require.NotNil(t, out.DenyReason, "deny reason must be populated")
|
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")
|
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) {
|
func TestAllowlistEmptyModelValueDenies(t *testing.T) {
|
||||||
// A present-but-empty model is as undeterminable as an absent one.
|
// A present-but-empty model is as undeterminable as an absent one.
|
||||||
mw := New(Config{ModelAllowlist: []string{"gpt-4o"}})
|
mw := New(providerCfg("gpt-4o"))
|
||||||
out, err := mw.Invoke(context.Background(), newInput(
|
out, err := mw.Invoke(context.Background(), newInputProvider(testProvider,
|
||||||
middleware.KV{Key: middleware.KeyLLMModel, Value: " "},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: " "},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NotNil(t, out)
|
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")
|
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")
|
assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowlistEmptyListAllowsMissingModel(t *testing.T) {
|
func TestAllowlistEmptyListAllowsMissingModel(t *testing.T) {
|
||||||
// Without an allowlist there is nothing to enforce, so a missing model is
|
// Without any provider allowlists there is nothing to enforce, so a missing
|
||||||
// still allowed — the fail-closed rule only applies when a list is set.
|
// model is still allowed — the fail-closed rule only applies when a
|
||||||
|
// restriction is in effect.
|
||||||
mw := New(Config{})
|
mw := New(Config{})
|
||||||
out, err := mw.Invoke(context.Background(), newInput())
|
out, err := mw.Invoke(context.Background(), newInput())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "no allowlist must allow even without a model")
|
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) {
|
func TestPromptCaptureDisabledEmitsNoPrompt(t *testing.T) {
|
||||||
mw := New(Config{})
|
mw := New(Config{})
|
||||||
out, err := mw.Invoke(context.Background(), newInput(
|
out, err := mw.Invoke(context.Background(), newInput(
|
||||||
@@ -217,8 +283,8 @@ func TestFactoryAcceptsZeroConfigs(t *testing.T) {
|
|||||||
|
|
||||||
func TestFactoryDecodesValidConfig(t *testing.T) {
|
func TestFactoryDecodesValidConfig(t *testing.T) {
|
||||||
cfg := Config{
|
cfg := Config{
|
||||||
ModelAllowlist: []string{"gpt-4o"},
|
ProviderAllowlists: map[string][]string{testProvider: {"gpt-4o"}},
|
||||||
PromptCapture: PromptCapture{Enabled: true, RedactPii: true},
|
PromptCapture: PromptCapture{Enabled: true, RedactPii: true},
|
||||||
}
|
}
|
||||||
raw, err := json.Marshal(cfg)
|
raw, err := json.Marshal(cfg)
|
||||||
require.NoError(t, err, "marshalling test config must succeed")
|
require.NoError(t, err, "marshalling test config must succeed")
|
||||||
@@ -234,15 +300,15 @@ func TestFactoryRejectsMalformedJSON(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFactoryNormalisesAllowlist(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)
|
mw, err := Factory{}.New(raw)
|
||||||
require.NoError(t, err)
|
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"},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "factory must lowercase + trim allowlist entries")
|
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"},
|
middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-3"},
|
||||||
))
|
))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|||||||
@@ -25,10 +25,17 @@ func runParserGuardrail(t *testing.T, url string, body []byte, allowlist []strin
|
|||||||
})
|
})
|
||||||
require.NoError(t, err, "parser must not error")
|
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{
|
out, err := guard.Invoke(context.Background(), &middleware.Input{
|
||||||
Slot: middleware.SlotOnRequest,
|
Slot: middleware.SlotOnRequest,
|
||||||
Metadata: parsed.Metadata,
|
Metadata: meta,
|
||||||
})
|
})
|
||||||
require.NoError(t, err, "guardrail must not error")
|
require.NoError(t, err, "guardrail must not error")
|
||||||
require.NotNil(t, out, "guardrail must return an output")
|
require.NotNil(t, out, "guardrail must return an output")
|
||||||
|
|||||||
Reference in New Issue
Block a user