mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 06:59:08 +02:00
[management,proxy] Agent network: per-account LLM gateway (policy, metering, multi-provider) (#6555)
* [agent-network] Shared proto, OpenAPI schema, and generated types * [agent-network] Management: store, manager, synthesizer, policy engine, provider catalog, HTTP/gRPC API Adds the account-scoped agent-network module: provider/policy/budget CRUD and store, the reverse-proxy service synthesizer, policy selection + limit enforcement, the provider catalog (incl. Vertex AI and AWS Bedrock entries), and the management HTTP + proxy gRPC surfaces. * [management] Fix agent-network proxy-peer fan-out on affected-peer recompute The affected-peers resolver loaded only persisted reverse-proxy services, but agent-network services are synthesized on demand and never persisted. As a result the embedded proxy peer was never folded into the affected set when a client's group changed, so the proxy received no network-map update for a newly authorised client and rejected its handshake until a full resync (restart). loadProxyServices now merges the synthesized agent-network services (injected via a registration hook to avoid an import cycle), so proxy peers learn newly authorised clients immediately. * [proxy] Reverse-proxy middleware framework, chain, and request plumbing The per-target middleware chain (slots, dispatcher, mutation gate, metadata merger), body capture, access-log terminal sink, and the proxy wiring that builds + runs chains for synthesized agent-network services. * [proxy] LLM parsers, pricing, and builtin middlewares (OpenAI, Anthropic, Vertex AI, AWS Bedrock) Request/response parsers and SSE/event-stream metering, the embedded pricing table, and the builtin middleware set: request parser, router, policy limit-check/record, cost meter, guardrail, identity inject, response parser. Includes the path-routed providers — Google Vertex AI (keyfile:: service-account OAuth minting) and AWS Bedrock (bearer auth, invoke/converse/streaming, optional /bedrock prefix) — plus the Models allowlist and unmeterable-publisher deny. * [proxy] IPv6 in-place apply and TCP accept-loop hardening on netstack listeners * [agent-network] End-to-end test suite, module docs, and deployment preset * [agent-network] Fix codespell typos and exclude false positives - labelgen word pool: vermillion -> vermilion, racoon -> raccoon. - codespell ignore list: add flate (Go compress/flate package), recordin (a test-local identifier), and unparseable (a valid alternative spelling used consistently across identifiers + a metadata-value constant). * [management] Set LastSeen on injected proxy peer in realstack test (MySQL strict-mode) The injected embedded proxy peer had a PeerStatus with a zero LastSeen, which serializes to '0000-00-00' and is rejected by MySQL in strict mode (SQLite tolerates it). Set LastSeen to a valid time so SaveAccount succeeds on both engines. * [agent-network] Remove e2e shell-script suite from this branch The end-to-end shell scripts under scripts/e2e/ are maintained in a separate testing suite and are not part of this change set. * [agent-network] Polish module docs: remove internal review scaffolding, fix links, verify diagrams Strip PR-review framing, commit references, absolute paths, and stale internal references from the agent-network module docs; fix broken relative links; verify all diagrams against the current architecture. Remove the internal AI-reviewer prompt file. * [management] Refine session expiration handling to support 3-state encoding for SSO deadlines * [agent-network] Relocate agentnetwork package to internals/modules Move management/server/agentnetwork (and its catalog/, labelgen/, types/ subpackages) to management/internals/modules/agentnetwork, alongside the reverse-proxy module, and rewrite all importers. Pure relocation: package names, the synthesizer + affectedpeers registration hook, and store access (shared store.Store) are unchanged, so no import cycle is introduced (affectedpeers still depends only on the agentnetwork/types leaf). * [agent-network] Co-locate HTTP handlers in the module (RegisterEndpoints) Move the agent-network HTTP handlers from server/http/handlers/agentnetwork into the module at internals/modules/agentnetwork/handlers (package handlers) and rename the entrypoint AddEndpoints -> RegisterEndpoints, matching the reverse-proxy module convention. Wiring in http/handler.go updated accordingly.
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
package llm_router
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
)
|
||||
|
||||
// ProviderRoute describes one upstream LLM provider the router can
|
||||
// hand a request to. Models lists the model identifiers the provider
|
||||
// claims; UpstreamScheme + UpstreamHost replace the synth target's
|
||||
// placeholder URL on a match. UpstreamPath is the path component of
|
||||
// the configured upstream URL — the router uses it to disambiguate
|
||||
// providers that claim the same model: when more than one provider
|
||||
// matches the model, the route whose UpstreamPath is a prefix of the
|
||||
// incoming request path is preferred (longest match wins, empty path
|
||||
// is the catchall). AuthHeaderName + AuthHeaderValue are the
|
||||
// per-provider credential the router injects after stripping the
|
||||
// vendor auth headers from the inbound request.
|
||||
//
|
||||
// AllowedGroupIDs is the union of source-group IDs across every
|
||||
// enabled policy that authorises this provider. The router treats it
|
||||
// as a hard filter: a route whose AllowedGroupIDs has no intersection
|
||||
// with the caller's UserGroups is removed from the candidate list
|
||||
// before the path-prefix tiebreak. A route with empty AllowedGroupIDs
|
||||
// is unreachable; the synthesiser only emits policy-bound routes.
|
||||
type ProviderRoute struct {
|
||||
ID string `json:"id"`
|
||||
// Vendor is the parser surface this provider speaks ("openai",
|
||||
// "anthropic", …), matching the llm.provider value llm_request_parser
|
||||
// emits from the request. When set, the router keeps a vendor-tagged
|
||||
// request on a same-vendor route so catch-all gateways of a different
|
||||
// vendor can't swallow it. Empty disables vendor filtering for this
|
||||
// route.
|
||||
Vendor string `json:"vendor,omitempty"`
|
||||
Models []string `json:"models"`
|
||||
UpstreamScheme string `json:"upstream_scheme"`
|
||||
UpstreamHost string `json:"upstream_host"`
|
||||
UpstreamPath string `json:"upstream_path,omitempty"`
|
||||
AuthHeaderName string `json:"auth_header_name"`
|
||||
AuthHeaderValue string `json:"auth_header_value"`
|
||||
AllowedGroupIDs []string `json:"allowed_group_ids"`
|
||||
// Vertex marks a Google Vertex AI provider. Vertex requests carry the
|
||||
// model in the URL path, so the router selects this route by path
|
||||
// (isVertexPath) and bypasses the model/vendor table entirely.
|
||||
Vertex bool `json:"vertex,omitempty"`
|
||||
// Bedrock marks an AWS Bedrock provider. Bedrock requests carry the model
|
||||
// in the URL path (/model/{id}/{action}), so the router selects this route
|
||||
// by path (isBedrockPath) and bypasses the model/vendor table; auth is the
|
||||
// static AuthHeaderValue bearer token (no token minting).
|
||||
Bedrock bool `json:"bedrock,omitempty"`
|
||||
// GCPServiceAccountKeyB64 is a base64-encoded GCP service-account JSON
|
||||
// key. When set, the router mints + refreshes a short-lived OAuth2 access
|
||||
// token from it at request time and injects it as the auth header value
|
||||
// (instead of the static AuthHeaderValue) — so the gateway holds a durable
|
||||
// Vertex credential rather than a 1-hour token.
|
||||
GCPServiceAccountKeyB64 string `json:"gcp_sa_key_b64,omitempty"`
|
||||
}
|
||||
|
||||
// Config is the on-wire configuration accepted by the factory. An
|
||||
// empty Providers slice yields a router that denies every request as
|
||||
// not-routable; the synthesiser is responsible for stamping the
|
||||
// account's enabled providers into this slice.
|
||||
type Config struct {
|
||||
Providers []ProviderRoute `json:"providers"`
|
||||
}
|
||||
|
||||
// Factory builds llm_router instances from raw config bytes.
|
||||
type Factory struct{}
|
||||
|
||||
// ID returns the registry identifier.
|
||||
func (Factory) ID() string { return ID }
|
||||
|
||||
// New constructs a middleware instance. Empty, null, and {} configs
|
||||
// yield a router with an empty Providers slice — every request denies
|
||||
// with model_not_routable. Non-empty payloads must parse cleanly so
|
||||
// misconfigurations surface at chain build time.
|
||||
func (Factory) New(rawConfig []byte) (middleware.Middleware, error) {
|
||||
cfg := Config{}
|
||||
if !isEmptyJSON(rawConfig) {
|
||||
if err := json.Unmarshal(rawConfig, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("decode config: %w", err)
|
||||
}
|
||||
}
|
||||
return New(cfg), nil
|
||||
}
|
||||
|
||||
// isEmptyJSON reports whether the payload is whitespace, null, or an
|
||||
// empty object/array. The caller skips Unmarshal in that case so the
|
||||
// zero-value Config flows through unchanged.
|
||||
func isEmptyJSON(raw []byte) bool {
|
||||
trimmed := strings.TrimSpace(string(bytes.TrimSpace(raw)))
|
||||
switch trimmed {
|
||||
case "", "null", "{}", "[]":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func init() {
|
||||
builtin.Register(Factory{})
|
||||
}
|
||||
@@ -0,0 +1,793 @@
|
||||
// Package llm_router implements the SlotOnRequest middleware that
|
||||
// routes a request to an upstream LLM provider based on the model name
|
||||
// emitted upstream by llm_request_parser. The router rewrites the
|
||||
// request's outbound target (scheme + host), strips known LLM-vendor
|
||||
// auth headers, and injects the per-provider auth header from the
|
||||
// matched route. Unknown or unconfigured models deny with a 403 and
|
||||
// the canonical llm_policy.model_not_routable code.
|
||||
package llm_router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
"golang.org/x/oauth2/google"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
)
|
||||
|
||||
// gcpScope is the OAuth2 scope minted for Vertex AI service-account auth.
|
||||
const gcpScope = "https://www.googleapis.com/auth/cloud-platform"
|
||||
|
||||
// gcpTokenTimeout bounds each GCP token mint/refresh HTTP call so a slow or
|
||||
// unreachable token endpoint can't block the request indefinitely.
|
||||
const gcpTokenTimeout = 10 * time.Second
|
||||
|
||||
// ID is the registry key for this middleware.
|
||||
const ID = "llm_router"
|
||||
|
||||
// Version is reported via Middleware.Version().
|
||||
const Version = "1.0.0"
|
||||
|
||||
const (
|
||||
denyCodeNotRoutable = "llm_policy.model_not_routable"
|
||||
denyReasonNotRoutable = "model_not_routable"
|
||||
denyCodeNoAuthorisedRoute = "llm_policy.no_authorised_provider"
|
||||
denyReasonNoAuthorisedRoute = "no_authorised_provider"
|
||||
//nolint:gosec // deny code label, not a credential
|
||||
denyCodeUpstreamAuth = "llm_policy.upstream_auth_failed"
|
||||
denyCodeUnmeterable = "llm_policy.unmeterable_publisher"
|
||||
denyReasonUnmeterable = "unmeterable_publisher"
|
||||
)
|
||||
|
||||
// strippedAuthHeaders is the closed list of vendor authentication
|
||||
// credentials the router clears before injecting the provider-specific
|
||||
// credential. Strictly auth headers — vendor-specific metadata
|
||||
// (anthropic-version, openai-organization, openai-project, etc.) is
|
||||
// NOT stripped because the client SDK sets those and the upstream
|
||||
// requires them (e.g. Anthropic returns 400 without
|
||||
// anthropic-version). Each entry is canonicalised by Go's
|
||||
// http.Header.Del/Set, so listing the canonical shapes here is
|
||||
// sufficient.
|
||||
var strippedAuthHeaders = []string{
|
||||
"Authorization", // OpenAI, OpenAI-compatible, most vendors, Bedrock bearer
|
||||
"Proxy-Authorization", // upstream proxy auth (defense-in-depth)
|
||||
"x-api-key", // Anthropic
|
||||
"api-key", // Azure OpenAI
|
||||
"X-Amz-Date", // AWS SigV4 — strip client-supplied AWS signing material
|
||||
"X-Amz-Security-Token",
|
||||
"X-Amz-Content-Sha256",
|
||||
}
|
||||
|
||||
// Middleware routes requests to upstream LLM providers based on the
|
||||
// llm.model metadata emitted by llm_request_parser.
|
||||
type Middleware struct {
|
||||
cfg Config
|
||||
// tokenSrc caches one auto-refreshing OAuth2 TokenSource per GCP
|
||||
// service-account key (keyed by a hash of the key material), so Vertex
|
||||
// token minting happens once and refreshes are amortised across requests.
|
||||
tokenMu sync.Mutex
|
||||
tokenSrc map[string]oauth2.TokenSource
|
||||
}
|
||||
|
||||
// New constructs a Middleware with the supplied configuration. Empty
|
||||
// or nil Providers slice yields a router that denies every request as
|
||||
// not-routable.
|
||||
func New(cfg Config) *Middleware {
|
||||
return &Middleware{cfg: cfg, tokenSrc: map[string]oauth2.TokenSource{}}
|
||||
}
|
||||
|
||||
// ID returns the registry identifier.
|
||||
func (m *Middleware) ID() string { return ID }
|
||||
|
||||
// Version returns the implementation version.
|
||||
func (m *Middleware) Version() string { return Version }
|
||||
|
||||
// Slot reports the chain slot the middleware lives in.
|
||||
func (m *Middleware) Slot() middleware.Slot { return middleware.SlotOnRequest }
|
||||
|
||||
// AcceptedContentTypes returns nil because the router only consults
|
||||
// the metadata emitted by llm_request_parser.
|
||||
func (m *Middleware) AcceptedContentTypes() []string { return nil }
|
||||
|
||||
// MetadataKeys is the closed set of metadata keys this middleware may
|
||||
// emit. The accumulator drops anything outside this allowlist.
|
||||
func (m *Middleware) MetadataKeys() []string {
|
||||
return []string{
|
||||
middleware.KeyLLMResolvedProviderID,
|
||||
middleware.KeyLLMAuthorisingGroups,
|
||||
middleware.KeyLLMPolicyDecision,
|
||||
middleware.KeyLLMPolicyReason,
|
||||
}
|
||||
}
|
||||
|
||||
// MutationsSupported reports that the middleware emits header and
|
||||
// upstream-rewrite mutations.
|
||||
func (m *Middleware) MutationsSupported() bool { return true }
|
||||
|
||||
// Close releases resources owned by the middleware. The router is
|
||||
// stateless, so this is a no-op.
|
||||
func (m *Middleware) Close() error { return nil }
|
||||
|
||||
// matchOutcome captures why matchRoute returned what it did so the
|
||||
// caller can distinguish "no provider knows this model" from "providers
|
||||
// know it but none authorise this peer's groups".
|
||||
type matchOutcome int
|
||||
|
||||
const (
|
||||
matchOutcomeFound matchOutcome = iota
|
||||
matchOutcomeUnknownModel
|
||||
matchOutcomeUnauthorised
|
||||
)
|
||||
|
||||
// Invoke resolves the model to a provider authorised for the caller's
|
||||
// groups, strips known vendor auth headers, and injects the route's
|
||||
// auth header. Unknown models deny with model_not_routable; models
|
||||
// known to a provider that no policy authorises for the caller deny
|
||||
// with no_authorised_provider.
|
||||
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
|
||||
// Vertex AI carries the model in the URL path, not the body, and is
|
||||
// selected by path rather than by the model/vendor table. Route it before
|
||||
// the model lookup so a model the parser extracted from the path can't be
|
||||
// claimed by a same-vendor direct provider (e.g. claude-* on api.anthropic.com).
|
||||
reqPath := requestPath(in.URL)
|
||||
if isVertexPath(reqPath) {
|
||||
model, _ := lookupMetadata(in.Metadata, middleware.KeyLLMModel)
|
||||
// The request parser emits no llm.provider for a Vertex publisher it
|
||||
// can't parse (e.g. google/gemini). Forwarding such a request would
|
||||
// bypass token/budget metering, so deny it rather than serve it
|
||||
// unmetered.
|
||||
if vendor, _ := lookupMetadata(in.Metadata, middleware.KeyLLMProvider); vendor == "" {
|
||||
return denyUnmeterable(), nil
|
||||
}
|
||||
route, outcome := m.matchVertex(reqPath, model, in.UserGroups)
|
||||
switch outcome {
|
||||
case matchOutcomeFound:
|
||||
return m.allowWithRoute(route, in.UserGroups), nil
|
||||
case matchOutcomeUnauthorised:
|
||||
return denyNoAuthorisedRoute(model), nil
|
||||
default:
|
||||
return denyUnknownModel(model), nil
|
||||
}
|
||||
}
|
||||
|
||||
// Bedrock likewise carries the model in the URL path (/model/{id}/{action}),
|
||||
// optionally behind a "/bedrock" gateway-namespace prefix. Route it by path
|
||||
// before the model lookup; when the prefix is present, strip it from the
|
||||
// forwarded path so the real Bedrock endpoint receives its native path.
|
||||
if isBedrockPath(reqPath) {
|
||||
model, _ := lookupMetadata(in.Metadata, middleware.KeyLLMModel)
|
||||
native, hadPrefix := splitBedrockNamespace(reqPath)
|
||||
route, outcome := m.matchBedrock(native, model, in.UserGroups)
|
||||
switch outcome {
|
||||
case matchOutcomeFound:
|
||||
out := m.allowWithRoute(route, in.UserGroups)
|
||||
if hadPrefix && out.Mutations != nil && out.Mutations.RewriteUpstream != nil {
|
||||
out.Mutations.RewriteUpstream.StripPathPrefix = bedrockNamespacePrefix
|
||||
}
|
||||
return out, nil
|
||||
case matchOutcomeUnauthorised:
|
||||
return denyNoAuthorisedRoute(model), nil
|
||||
default:
|
||||
return denyUnknownModel(model), nil
|
||||
}
|
||||
}
|
||||
|
||||
model, ok := lookupMetadata(in.Metadata, middleware.KeyLLMModel)
|
||||
if !ok || model == "" {
|
||||
// Non-inference endpoints (model listing) carry no model but still
|
||||
// need rewriting from the synth placeholder to a real upstream;
|
||||
// clients such as Codex call GET /v1/models at startup to enumerate
|
||||
// availability and read a 403 as "model unavailable".
|
||||
route, outcome := m.matchModelless(requestPath(in.URL), in.UserGroups)
|
||||
switch outcome {
|
||||
case matchOutcomeFound:
|
||||
return m.allowWithRoute(route, in.UserGroups), nil
|
||||
case matchOutcomeUnauthorised:
|
||||
// A recognised model-less endpoint exists but no provider
|
||||
// authorises the caller — deny as an authorisation failure
|
||||
// rather than masking it as a missing model.
|
||||
return denyNoAuthorisedRoute(model), nil
|
||||
default:
|
||||
return denyMissingModel(), nil
|
||||
}
|
||||
}
|
||||
|
||||
vendor, _ := lookupMetadata(in.Metadata, middleware.KeyLLMProvider)
|
||||
route, outcome := m.matchRoute(model, vendor, requestPath(in.URL), in.UserGroups)
|
||||
switch outcome {
|
||||
case matchOutcomeFound:
|
||||
return m.allowWithRoute(route, in.UserGroups), nil
|
||||
case matchOutcomeUnauthorised:
|
||||
return denyNoAuthorisedRoute(model), nil
|
||||
default:
|
||||
return denyUnknownModel(model), nil
|
||||
}
|
||||
}
|
||||
|
||||
// matchRoute returns the ProviderRoute that should serve the given
|
||||
// model + request path for a caller in the given user-groups. Selection
|
||||
// is:
|
||||
//
|
||||
// 1. Filter the configured providers to those whose Models list
|
||||
// contains the model.
|
||||
// 2. Filter the model-matched candidates to those whose
|
||||
// AllowedGroupIDs intersect the caller's UserGroups. A route with
|
||||
// no AllowedGroupIDs is the catch-all: it stays in the list. If
|
||||
// the model was known but no candidate is authorised for this
|
||||
// peer, return matchOutcomeUnauthorised so the caller can emit
|
||||
// the dedicated no_authorised_provider deny code.
|
||||
// 3. Vendor precedence: when the request carries a detected vendor
|
||||
// (llm.provider) and at least one candidate is the same vendor,
|
||||
// drop the rest — a vendor-tagged request must never cross to
|
||||
// another vendor's route (e.g. an Anthropic call landing on an
|
||||
// OpenAI-compatible gateway that also claims the model).
|
||||
// 4. Model precedence over path: a route that explicitly lists the
|
||||
// model beats a catch-all (empty Models) gateway.
|
||||
// 5. Disambiguate the survivors by URL path prefix: longest
|
||||
// UpstreamPath that prefix-matches the request path wins; an empty
|
||||
// UpstreamPath is the catchall. If none prefix-matches, fall back
|
||||
// to declaration order so the model stays routable.
|
||||
func (m *Middleware) matchRoute(model, vendor, reqPath string, userGroups []string) (ProviderRoute, matchOutcome) {
|
||||
var modelMatched []ProviderRoute
|
||||
for _, route := range m.cfg.Providers {
|
||||
if routeClaimsModel(route, model) {
|
||||
modelMatched = append(modelMatched, route)
|
||||
}
|
||||
}
|
||||
if len(modelMatched) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnknownModel
|
||||
}
|
||||
|
||||
// Vendor pinning runs BEFORE the group filter so a request the parser
|
||||
// tagged with a vendor can never cross to another vendor's route — not
|
||||
// even an authorised one. Narrow to same-vendor routes when any
|
||||
// model-matched route declares that vendor; setups with no vendor tag on
|
||||
// any route fall through unchanged. After narrowing, if no same-vendor
|
||||
// route authorises the caller, that's matchOutcomeUnauthorised (no
|
||||
// cross-vendor fallback).
|
||||
if vendor != "" {
|
||||
if vendorMatched := matchingVendor(modelMatched, vendor); len(vendorMatched) > 0 {
|
||||
modelMatched = vendorMatched
|
||||
}
|
||||
}
|
||||
|
||||
var candidates []ProviderRoute
|
||||
for _, route := range modelMatched {
|
||||
if routeAuthorisesGroups(route, userGroups) {
|
||||
candidates = append(candidates, route)
|
||||
}
|
||||
}
|
||||
if len(candidates) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnauthorised
|
||||
}
|
||||
|
||||
// Model routing takes precedence over path. A route that explicitly
|
||||
// lists the model must beat a catch-all (empty Models) gateway that
|
||||
// claims every model — otherwise an Anthropic request can fall through
|
||||
// to an OpenAI-compatible gateway declared earlier. Only when no
|
||||
// candidate explicitly claims the model do the catch-alls compete, and
|
||||
// the path-prefix tiebreak applies within whichever tier wins.
|
||||
if explicit := explicitlyClaiming(candidates, model); len(explicit) > 0 {
|
||||
candidates = explicit
|
||||
}
|
||||
if len(candidates) == 1 {
|
||||
return candidates[0], matchOutcomeFound
|
||||
}
|
||||
|
||||
best := candidates[0]
|
||||
bestLen := -1
|
||||
for _, c := range candidates {
|
||||
if !pathPrefixMatches(c.UpstreamPath, reqPath) {
|
||||
continue
|
||||
}
|
||||
if len(c.UpstreamPath) > bestLen {
|
||||
best = c
|
||||
bestLen = len(c.UpstreamPath)
|
||||
}
|
||||
}
|
||||
return best, matchOutcomeFound
|
||||
}
|
||||
|
||||
// isModelLessPath reports whether reqPath is a known OpenAI-shaped
|
||||
// non-inference endpoint that legitimately carries no model in its
|
||||
// request (the model-listing endpoints). These must route to an upstream
|
||||
// rather than deny, so model enumeration works end to end.
|
||||
func isModelLessPath(reqPath string) bool {
|
||||
return reqPath == "/v1/models" || strings.HasPrefix(reqPath, "/v1/models/")
|
||||
}
|
||||
|
||||
// isVertexPath reports whether reqPath is a Google Vertex AI publisher
|
||||
// endpoint: /v1/projects/{project}/locations/{region}/publishers/{publisher}/
|
||||
// models/{model}:{action}. The model + vendor live in the path, so these
|
||||
// requests are routed by path to the Vertex provider rather than by model.
|
||||
func isVertexPath(reqPath string) bool {
|
||||
return strings.HasPrefix(reqPath, "/v1/projects/") &&
|
||||
strings.Contains(reqPath, "/publishers/") &&
|
||||
strings.Contains(reqPath, "/models/")
|
||||
}
|
||||
|
||||
// bedrockNamespacePrefix is an optional gateway-namespace prefix some clients
|
||||
// place before the native Bedrock path to disambiguate it from other providers
|
||||
// that also use "/model/...". It is stripped before forwarding upstream.
|
||||
const bedrockNamespacePrefix = "/bedrock"
|
||||
|
||||
// splitBedrockNamespace removes an optional "/bedrock" namespace prefix,
|
||||
// returning the native Bedrock path and whether the prefix was present.
|
||||
func splitBedrockNamespace(reqPath string) (string, bool) {
|
||||
if strings.HasPrefix(reqPath, bedrockNamespacePrefix+"/") {
|
||||
return strings.TrimPrefix(reqPath, bedrockNamespacePrefix), true
|
||||
}
|
||||
return reqPath, false
|
||||
}
|
||||
|
||||
// isBedrockPath reports whether reqPath is an AWS Bedrock runtime model
|
||||
// endpoint: /model/{modelId}/{action} where action is invoke,
|
||||
// invoke-with-response-stream, converse, or converse-stream — optionally behind
|
||||
// a "/bedrock" gateway-namespace prefix. The model lives in the path, so these
|
||||
// requests are routed by path to the Bedrock provider.
|
||||
func isBedrockPath(reqPath string) bool {
|
||||
native, _ := splitBedrockNamespace(reqPath)
|
||||
if !strings.HasPrefix(native, "/model/") {
|
||||
return false
|
||||
}
|
||||
return strings.HasSuffix(native, "/invoke") ||
|
||||
strings.HasSuffix(native, "/invoke-with-response-stream") ||
|
||||
strings.HasSuffix(native, "/converse") ||
|
||||
strings.HasSuffix(native, "/converse-stream")
|
||||
}
|
||||
|
||||
// matchVertex selects the Vertex provider authorised for the caller's groups
|
||||
// and claiming the requested model.
|
||||
func (m *Middleware) matchVertex(reqPath, model string, userGroups []string) (ProviderRoute, matchOutcome) {
|
||||
return m.matchPathRoute(reqPath, model, userGroups, func(r ProviderRoute) bool { return r.Vertex })
|
||||
}
|
||||
|
||||
// matchBedrock selects the Bedrock provider authorised for the caller's groups
|
||||
// and claiming the requested model.
|
||||
func (m *Middleware) matchBedrock(reqPath, model string, userGroups []string) (ProviderRoute, matchOutcome) {
|
||||
return m.matchPathRoute(reqPath, model, userGroups, func(r ProviderRoute) bool { return r.Bedrock })
|
||||
}
|
||||
|
||||
// matchPathRoute selects a path-routed provider (Vertex/Bedrock). These carry
|
||||
// the model in the URL, so the model/vendor table is bypassed — but the route's
|
||||
// configured Models allowlist is still enforced (empty Models = catch-all) so a
|
||||
// provider credential can't be used for models the operator didn't authorise.
|
||||
// Returns matchOutcomeUnauthorised when no style route authorises the caller's
|
||||
// groups, matchOutcomeUnknownModel when an authorised route exists but none
|
||||
// claims the model (or no style route exists at all), else the chosen route
|
||||
// (longest UpstreamPath prefix-match wins among multiple).
|
||||
func (m *Middleware) matchPathRoute(reqPath, model string, userGroups []string, isStyle func(ProviderRoute) bool) (ProviderRoute, matchOutcome) {
|
||||
var styled []ProviderRoute
|
||||
for _, route := range m.cfg.Providers {
|
||||
if isStyle(route) {
|
||||
styled = append(styled, route)
|
||||
}
|
||||
}
|
||||
if len(styled) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnknownModel
|
||||
}
|
||||
|
||||
var authorised []ProviderRoute
|
||||
for _, route := range styled {
|
||||
if routeAuthorisesGroups(route, userGroups) {
|
||||
authorised = append(authorised, route)
|
||||
}
|
||||
}
|
||||
if len(authorised) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnauthorised
|
||||
}
|
||||
|
||||
var candidates []ProviderRoute
|
||||
for _, route := range authorised {
|
||||
if routeClaimsModel(route, model) {
|
||||
candidates = append(candidates, route)
|
||||
}
|
||||
}
|
||||
if len(candidates) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnknownModel
|
||||
}
|
||||
if len(candidates) == 1 {
|
||||
return candidates[0], matchOutcomeFound
|
||||
}
|
||||
|
||||
best := candidates[0]
|
||||
bestLen := -1
|
||||
for _, c := range candidates {
|
||||
if !pathPrefixMatches(c.UpstreamPath, reqPath) {
|
||||
continue
|
||||
}
|
||||
if len(c.UpstreamPath) > bestLen {
|
||||
best = c
|
||||
bestLen = len(c.UpstreamPath)
|
||||
}
|
||||
}
|
||||
return best, matchOutcomeFound
|
||||
}
|
||||
|
||||
// matchModelless selects a route for a non-inference, model-less request.
|
||||
// It mirrors matchRoute's group-authorisation filter and path-prefix
|
||||
// tiebreak but skips the per-model filter, since any provider the caller's
|
||||
// groups authorise can serve a model-listing request. Returns
|
||||
// matchOutcomeFound with the chosen route (single authorised provider wins
|
||||
// outright; multiple fall to the longest UpstreamPath prefix-match, then
|
||||
// declaration order), matchOutcomeUnauthorised when no provider authorises
|
||||
// the caller, or matchOutcomeUnknownModel when the path isn't a recognised
|
||||
// model-less endpoint.
|
||||
func (m *Middleware) matchModelless(reqPath string, userGroups []string) (ProviderRoute, matchOutcome) {
|
||||
if !isModelLessPath(reqPath) {
|
||||
return ProviderRoute{}, matchOutcomeUnknownModel
|
||||
}
|
||||
var candidates []ProviderRoute
|
||||
for _, route := range m.cfg.Providers {
|
||||
// Vertex/Bedrock are path-routed and don't serve OpenAI-style
|
||||
// model-listing endpoints; including them here could rewrite a
|
||||
// GET /v1/models to an upstream that 404s it.
|
||||
if route.Vertex || route.Bedrock {
|
||||
continue
|
||||
}
|
||||
if routeAuthorisesGroups(route, userGroups) {
|
||||
candidates = append(candidates, route)
|
||||
}
|
||||
}
|
||||
if len(candidates) == 0 {
|
||||
return ProviderRoute{}, matchOutcomeUnauthorised
|
||||
}
|
||||
if len(candidates) == 1 {
|
||||
return candidates[0], matchOutcomeFound
|
||||
}
|
||||
|
||||
best := candidates[0]
|
||||
bestLen := -1
|
||||
for _, c := range candidates {
|
||||
if !pathPrefixMatches(c.UpstreamPath, reqPath) {
|
||||
continue
|
||||
}
|
||||
if len(c.UpstreamPath) > bestLen {
|
||||
best = c
|
||||
bestLen = len(c.UpstreamPath)
|
||||
}
|
||||
}
|
||||
return best, matchOutcomeFound
|
||||
}
|
||||
|
||||
// routeAuthorisesGroups reports whether the route's AllowedGroupIDs
|
||||
// intersect the caller's userGroups. A route with empty AllowedGroupIDs
|
||||
// is unreachable: the synthesiser only emits routes bound to at least
|
||||
// one enabled policy, so an empty list signals a misconfiguration that
|
||||
// must not be allowed to fall through.
|
||||
func routeAuthorisesGroups(r ProviderRoute, userGroups []string) bool {
|
||||
for _, ug := range userGroups {
|
||||
for _, ag := range r.AllowedGroupIDs {
|
||||
if ug == ag {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// authorisingGroupsCSV returns the sorted, deduplicated comma-separated
|
||||
// intersection of routeGroups and userGroups — i.e. the groups that
|
||||
// actually authorise the resolved route for this caller. Returns the
|
||||
// empty string when the intersection is empty (shouldn't happen on the
|
||||
// allow path, but defensive).
|
||||
func authorisingGroupsCSV(routeGroups, userGroups []string) string {
|
||||
if len(routeGroups) == 0 || len(userGroups) == 0 {
|
||||
return ""
|
||||
}
|
||||
allowed := make(map[string]struct{}, len(routeGroups))
|
||||
for _, g := range routeGroups {
|
||||
allowed[g] = struct{}{}
|
||||
}
|
||||
seen := make(map[string]struct{}, len(userGroups))
|
||||
out := make([]string, 0, len(userGroups))
|
||||
for _, ug := range userGroups {
|
||||
if _, ok := allowed[ug]; !ok {
|
||||
continue
|
||||
}
|
||||
if _, dup := seen[ug]; dup {
|
||||
continue
|
||||
}
|
||||
seen[ug] = struct{}{}
|
||||
out = append(out, ug)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return ""
|
||||
}
|
||||
sort.Strings(out)
|
||||
return strings.Join(out, ",")
|
||||
}
|
||||
|
||||
// matchingVendor returns the subset of routes whose Vendor equals the
|
||||
// request's detected vendor. Routes with an empty Vendor never match — an
|
||||
// untagged route can't be asserted to speak the request's surface, so it
|
||||
// stays out of the vendor-filtered set (but remains eligible via the
|
||||
// fall-through when no route matches the vendor at all).
|
||||
func matchingVendor(routes []ProviderRoute, vendor string) []ProviderRoute {
|
||||
var out []ProviderRoute
|
||||
for _, r := range routes {
|
||||
if r.Vendor == vendor {
|
||||
out = append(out, r)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// explicitlyClaiming returns the subset of routes whose Models list
|
||||
// names the model exactly. Catch-all routes (empty Models) are excluded,
|
||||
// so callers can prefer a provider that genuinely declares the model over
|
||||
// a gateway that claims everything.
|
||||
func explicitlyClaiming(routes []ProviderRoute, model string) []ProviderRoute {
|
||||
var out []ProviderRoute
|
||||
for _, r := range routes {
|
||||
for _, candidate := range r.Models {
|
||||
if candidate == model {
|
||||
out = append(out, r)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// routeClaimsModel reports whether the route's Models list contains
|
||||
// the given model identifier. An empty Models list is treated as
|
||||
// "claim every model" — used by gateway-style providers (LiteLLM,
|
||||
// custom OpenAI-compatible endpoints) that proxy an open-ended set of
|
||||
// upstream models the operator can't enumerate in NetBird's provider
|
||||
// config.
|
||||
func routeClaimsModel(route ProviderRoute, model string) bool {
|
||||
if len(route.Models) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, candidate := range route.Models {
|
||||
if candidate == model {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// pathPrefixMatches reports whether upstreamPath matches reqPath on a path-
|
||||
// segment boundary: an exact match, or reqPath continuing after
|
||||
// upstreamPath at a "/" separator. This avoids a sibling base like
|
||||
// "/openai" spuriously matching "/openai-test". An empty (or "/")
|
||||
// upstreamPath always matches (catchall).
|
||||
func pathPrefixMatches(upstreamPath, reqPath string) bool {
|
||||
if upstreamPath == "" || upstreamPath == "/" {
|
||||
return true
|
||||
}
|
||||
upstreamPath = strings.TrimRight(upstreamPath, "/")
|
||||
return reqPath == upstreamPath || strings.HasPrefix(reqPath, upstreamPath+"/")
|
||||
}
|
||||
|
||||
// requestPath extracts the path component from an Input.URL string
|
||||
// (which is r.URL.String() — typically "/path?query"). Returns the
|
||||
// raw input on parse failure so the prefix check can still operate on
|
||||
// the unparsed value.
|
||||
func requestPath(raw string) string {
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return raw
|
||||
}
|
||||
return parsed.Path
|
||||
}
|
||||
|
||||
// allowWithRoute builds the Output for a successful route match. The
|
||||
// returned Mutations carry the upstream rewrite plus — riding on it —
|
||||
// the StripHeaders list and the AuthHeader to inject.
|
||||
//
|
||||
// The strip + inject MUST go through UpstreamRewrite (not HeadersAdd /
|
||||
// HeadersRemove) because the framework's mutation gate runs every
|
||||
// header change through a denylist that blocks Authorization,
|
||||
// Cookie, etc. — exactly the headers the router is replacing. The
|
||||
// proxy's upstream-build path applies AuthHeader / StripHeaders
|
||||
// directly, bypassing the denylist by virtue of being a trusted
|
||||
// proxy operation rather than an arbitrary middleware mutation.
|
||||
//
|
||||
// Emits the authorising-groups intersection alongside the resolved
|
||||
// provider id so identity-stamping middlewares (llm_identity_inject)
|
||||
// tag the request with ONLY the groups that authorised this specific
|
||||
// route — not every group the peer happens to be in.
|
||||
func (m *Middleware) allowWithRoute(route ProviderRoute, userGroups []string) *middleware.Output {
|
||||
rewrite := &middleware.UpstreamRewrite{
|
||||
Scheme: route.UpstreamScheme,
|
||||
Host: route.UpstreamHost,
|
||||
// UpstreamPath is the path component the operator pasted on
|
||||
// the provider record (e.g. "/v1/{account}/{gateway}/compat"
|
||||
// for Cloudflare AI Gateway). Carrying it on the rewrite so
|
||||
// the proxy's URL composer joins it with the agent's request
|
||||
// path — without this, the operator's configured upstream
|
||||
// path is silently dropped and the gateway returns a 4xx for
|
||||
// the malformed URL. Empty value leaves the original
|
||||
// target's path untouched.
|
||||
Path: route.UpstreamPath,
|
||||
StripHeaders: append([]string(nil), strippedAuthHeaders...),
|
||||
}
|
||||
authValue := route.AuthHeaderValue
|
||||
if route.GCPServiceAccountKeyB64 != "" {
|
||||
// Mint a short-lived OAuth2 token from the service-account key at
|
||||
// request time (cached + auto-refreshed) instead of a static value.
|
||||
bearer, err := m.gcpBearer(route.GCPServiceAccountKeyB64)
|
||||
if err != nil {
|
||||
return denyUpstreamAuth()
|
||||
}
|
||||
authValue = bearer
|
||||
}
|
||||
if route.AuthHeaderName != "" && authValue != "" {
|
||||
rewrite.AuthHeader = &middleware.AuthHeader{
|
||||
Name: route.AuthHeaderName,
|
||||
Value: authValue,
|
||||
}
|
||||
}
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionAllow,
|
||||
Mutations: &middleware.Mutations{RewriteUpstream: rewrite},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: route.ID},
|
||||
{Key: middleware.KeyLLMAuthorisingGroups, Value: authorisingGroupsCSV(route.AllowedGroupIDs, userGroups)},
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "allow"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// gcpBearer returns a "Bearer <token>" value minted from a base64-encoded GCP
|
||||
// service-account key, using a cached, auto-refreshing token source.
|
||||
func (m *Middleware) gcpBearer(saKeyB64 string) (string, error) {
|
||||
ts, err := m.gcpTokenSource(saKeyB64)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
tok, err := ts.Token()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("mint gcp token: %w", err)
|
||||
}
|
||||
return "Bearer " + tok.AccessToken, nil
|
||||
}
|
||||
|
||||
// gcpTokenSource returns the cached TokenSource for the given service-account
|
||||
// key, building it (decode base64 → parse JSON → cloud-platform scope) on first
|
||||
// use. The returned source caches the token and refreshes it before expiry.
|
||||
func (m *Middleware) gcpTokenSource(saKeyB64 string) (oauth2.TokenSource, error) {
|
||||
sum := sha256.Sum256([]byte(saKeyB64))
|
||||
key := hex.EncodeToString(sum[:])
|
||||
|
||||
m.tokenMu.Lock()
|
||||
defer m.tokenMu.Unlock()
|
||||
if m.tokenSrc == nil {
|
||||
m.tokenSrc = map[string]oauth2.TokenSource{}
|
||||
}
|
||||
if ts, ok := m.tokenSrc[key]; ok {
|
||||
return ts, nil
|
||||
}
|
||||
jsonKey, err := base64.StdEncoding.DecodeString(strings.TrimSpace(saKeyB64))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode gcp service-account key: %w", err)
|
||||
}
|
||||
conf, err := google.JWTConfigFromJSON(jsonKey, gcpScope)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse gcp service-account key: %w", err)
|
||||
}
|
||||
// Bound mint/refresh with a timeout HTTP client so a slow token endpoint
|
||||
// can't hang the request. The oauth2 library uses this client for the
|
||||
// lifetime of the (auto-refreshing) source.
|
||||
ctx := context.WithValue(context.Background(), oauth2.HTTPClient, &http.Client{Timeout: gcpTokenTimeout})
|
||||
ts := conf.TokenSource(ctx)
|
||||
m.tokenSrc[key] = ts
|
||||
return ts, nil
|
||||
}
|
||||
|
||||
// denyUpstreamAuth is returned when the router cannot obtain the upstream
|
||||
// credential (e.g. a malformed service-account key or an unreachable token
|
||||
// endpoint). It surfaces as a 502 — an upstream problem, not a policy denial.
|
||||
func denyUpstreamAuth() *middleware.Output {
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionDeny,
|
||||
DenyStatus: 502,
|
||||
DenyReason: &middleware.DenyReason{
|
||||
Code: denyCodeUpstreamAuth,
|
||||
Message: "could not obtain upstream credential",
|
||||
},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "deny"},
|
||||
{Key: middleware.KeyLLMPolicyReason, Value: "upstream_auth_failed"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// denyUnmeterable returns the deny envelope for a path-routed request whose
|
||||
// publisher has no parser surface, so its usage can't be metered. Serving it
|
||||
// would bypass token/budget caps, so it is rejected with a 403.
|
||||
func denyUnmeterable() *middleware.Output {
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionDeny,
|
||||
DenyStatus: 403,
|
||||
DenyReason: &middleware.DenyReason{
|
||||
Code: denyCodeUnmeterable,
|
||||
Message: "request publisher is not supported for metering",
|
||||
},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "deny"},
|
||||
{Key: middleware.KeyLLMPolicyReason, Value: denyReasonUnmeterable},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// denyMissingModel returns the deny envelope for a request whose
|
||||
// envelope has no llm.model metadata.
|
||||
func denyMissingModel() *middleware.Output {
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionDeny,
|
||||
DenyStatus: 403,
|
||||
DenyReason: &middleware.DenyReason{
|
||||
Code: denyCodeNotRoutable,
|
||||
Message: "missing llm.model on request envelope",
|
||||
},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "deny"},
|
||||
{Key: middleware.KeyLLMPolicyReason, Value: denyReasonNotRoutable},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// denyUnknownModel returns the deny envelope for a model that no
|
||||
// configured provider claims.
|
||||
func denyUnknownModel(model string) *middleware.Output {
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionDeny,
|
||||
DenyStatus: 403,
|
||||
DenyReason: &middleware.DenyReason{
|
||||
Code: denyCodeNotRoutable,
|
||||
Message: fmt.Sprintf("no provider configured for model %s", model),
|
||||
Details: map[string]string{"model": model},
|
||||
},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "deny"},
|
||||
{Key: middleware.KeyLLMPolicyReason, Value: denyReasonNotRoutable},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// denyNoAuthorisedRoute returns the deny envelope for a model that one
|
||||
// or more providers claim, but where no policy authorises the caller's
|
||||
// groups for any of those providers.
|
||||
func denyNoAuthorisedRoute(model string) *middleware.Output {
|
||||
return &middleware.Output{
|
||||
Decision: middleware.DecisionDeny,
|
||||
DenyStatus: 403,
|
||||
DenyReason: &middleware.DenyReason{
|
||||
Code: denyCodeNoAuthorisedRoute,
|
||||
Message: fmt.Sprintf("no policy authorises model %s for the caller's groups", model),
|
||||
Details: map[string]string{"model": model},
|
||||
},
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMPolicyDecision, Value: "deny"},
|
||||
{Key: middleware.KeyLLMPolicyReason, Value: denyReasonNoAuthorisedRoute},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// lookupMetadata returns the value for key plus a presence flag so
|
||||
// callers can distinguish absent from empty.
|
||||
func lookupMetadata(meta []middleware.KV, key string) (string, bool) {
|
||||
for _, kv := range meta {
|
||||
if kv.Key == key {
|
||||
return kv.Value, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
@@ -0,0 +1,840 @@
|
||||
package llm_router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
)
|
||||
|
||||
// metaValue returns the value for the first KV with the given key.
|
||||
func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) {
|
||||
t.Helper()
|
||||
for _, kv := range kvs {
|
||||
if kv.Key == key {
|
||||
return kv.Value, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// defaultTestGroup is the group id used by routes and inputs in tests
|
||||
// that don't specifically exercise the group-filter logic. Pairing it
|
||||
// with the same id on every test route keeps the legacy assertions
|
||||
// focused on routing/path behaviour without each one having to bake in
|
||||
// its own ACL.
|
||||
const defaultTestGroup = "grp-test"
|
||||
|
||||
// newInputWithModel returns an Input carrying llm.model in its metadata
|
||||
// bag, mimicking the post-llm_request_parser state the router observes
|
||||
// in production. UserGroups is populated with defaultTestGroup so the
|
||||
// router's group-filter pass authorises any test route whose
|
||||
// AllowedGroupIDs contains the same id.
|
||||
func newInputWithModel(model string) *middleware.Input {
|
||||
return &middleware.Input{
|
||||
Slot: middleware.SlotOnRequest,
|
||||
Metadata: []middleware.KV{{Key: middleware.KeyLLMModel, Value: model}},
|
||||
UserGroups: []string{defaultTestGroup},
|
||||
}
|
||||
}
|
||||
|
||||
// newInputWithModelAndURL returns an Input carrying both llm.model and
|
||||
// a request URL so router tests can exercise path-based disambiguation.
|
||||
func newInputWithModelAndURL(model, reqURL string) *middleware.Input {
|
||||
in := newInputWithModel(model)
|
||||
in.URL = reqURL
|
||||
return in
|
||||
}
|
||||
|
||||
func TestMiddlewareIdentity(t *testing.T) {
|
||||
mw := New(Config{})
|
||||
assert.Equal(t, ID, mw.ID(), "middleware ID must be llm_router")
|
||||
assert.Equal(t, Version, mw.Version(), "version must match the constant")
|
||||
assert.Equal(t, middleware.SlotOnRequest, mw.Slot(), "router must run in SlotOnRequest")
|
||||
assert.True(t, mw.MutationsSupported(), "router must declare mutations support")
|
||||
assert.Nil(t, mw.AcceptedContentTypes(), "router does not inspect bodies")
|
||||
assert.ElementsMatch(t,
|
||||
[]string{
|
||||
middleware.KeyLLMResolvedProviderID,
|
||||
middleware.KeyLLMAuthorisingGroups,
|
||||
middleware.KeyLLMPolicyDecision,
|
||||
middleware.KeyLLMPolicyReason,
|
||||
},
|
||||
mw.MetadataKeys(),
|
||||
"metadata key allowlist must match the spec",
|
||||
)
|
||||
require.NoError(t, mw.Close())
|
||||
}
|
||||
|
||||
func TestRouter_HappyPath(t *testing.T) {
|
||||
route := ProviderRoute{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer sk-test-123",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{route}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-4o"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "matched model must allow")
|
||||
|
||||
require.NotNil(t, out.Mutations, "matched route must emit mutations")
|
||||
rewrite := out.Mutations.RewriteUpstream
|
||||
require.NotNil(t, rewrite, "matched route must emit upstream rewrite")
|
||||
assert.Equal(t, "https", rewrite.Scheme, "rewrite scheme must come from the matched route")
|
||||
assert.Equal(t, "api.openai.com", rewrite.Host, "rewrite host must come from the matched route")
|
||||
|
||||
assert.ElementsMatch(t, strippedAuthHeaders, rewrite.StripHeaders,
|
||||
"strip list rides on UpstreamRewrite (bypasses framework denylist) and must cover every known vendor auth header")
|
||||
require.NotNil(t, rewrite.AuthHeader, "router must inject the auth header via the rewrite (not HeadersAdd) so the proxy bypasses the denylist")
|
||||
assert.Equal(t, "Authorization", rewrite.AuthHeader.Name, "injected header name must come from the route")
|
||||
assert.Equal(t, "Bearer sk-test-123", rewrite.AuthHeader.Value, "injected header value must come from the route")
|
||||
assert.Empty(t, out.Mutations.HeadersAdd, "router must not use HeadersAdd; auth flows through UpstreamRewrite.AuthHeader")
|
||||
assert.Empty(t, out.Mutations.HeadersRemove, "router must not use HeadersRemove; strip flows through UpstreamRewrite.StripHeaders")
|
||||
|
||||
resolved, ok := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
require.True(t, ok, "router must emit llm.resolved_provider_id on a match")
|
||||
assert.Equal(t, "openai-prod", resolved, "resolved provider id must be the matched route's ID")
|
||||
dec, _ := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision)
|
||||
assert.Equal(t, "allow", dec, "decision metadata must be allow on a match")
|
||||
}
|
||||
|
||||
func TestRouter_MissingModel(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{Slot: middleware.SlotOnRequest})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "missing llm.model must deny")
|
||||
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_not_routable", out.DenyReason.Code, "deny code must be model_not_routable")
|
||||
assert.Equal(t, "missing llm.model on request envelope", out.DenyReason.Message, "deny message must match spec")
|
||||
|
||||
dec, _ := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision)
|
||||
assert.Equal(t, "deny", dec, "decision metadata must be deny")
|
||||
reason, _ := metaValue(t, out.Metadata, middleware.KeyLLMPolicyReason)
|
||||
assert.Equal(t, "model_not_routable", reason, "reason metadata must be model_not_routable")
|
||||
}
|
||||
|
||||
// newModellessInput returns an Input with no llm.model and the given
|
||||
// request path, mimicking a GET /v1/models call (which carries no body
|
||||
// from which a model could be parsed). UserGroups matches defaultTestGroup.
|
||||
func newModellessInput(reqURL string) *middleware.Input {
|
||||
return &middleware.Input{
|
||||
Slot: middleware.SlotOnRequest,
|
||||
URL: reqURL,
|
||||
UserGroups: []string{defaultTestGroup},
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_ModelLessPath_RoutesToAuthorisedProvider(t *testing.T) {
|
||||
route := ProviderRoute{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{route}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newModellessInput("/v1/models?client_version=1"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "GET /v1/models must pass through, not deny")
|
||||
require.NotNil(t, out.Mutations, "a pass-through must rewrite the upstream")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream, "model-less route must still rewrite to the real upstream")
|
||||
assert.Equal(t, "api.openai.com", out.Mutations.RewriteUpstream.Host, "must target the authorised provider's host")
|
||||
|
||||
provider, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "openai-prod", provider, "resolved provider must be the authorised route")
|
||||
}
|
||||
|
||||
func TestRouter_ModelLessPath_MultiProviderDeclarationOrder(t *testing.T) {
|
||||
first := ProviderRoute{
|
||||
ID: "openai-a",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "a.example.com",
|
||||
}
|
||||
second := ProviderRoute{
|
||||
ID: "openai-b",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "b.example.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{first, second}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newModellessInput("/v1/models"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "model-less path must pass through with multiple providers")
|
||||
provider, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "openai-a", provider, "no path-prefix match falls back to declaration order")
|
||||
}
|
||||
|
||||
func TestRouter_ModelLessPath_UnauthorisedDenies(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{"some-other-group"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newModellessInput("/v1/models"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "no provider authorising the caller must still deny")
|
||||
}
|
||||
|
||||
func TestRouter_NonModelLessBodilessStillDenies(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}}})
|
||||
|
||||
// A bodiless POST to an inference path has no model and is NOT a
|
||||
// model-less endpoint, so it must keep denying.
|
||||
out, err := mw.Invoke(context.Background(), newModellessInput("/v1/responses"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "bodiless inference request must still deny")
|
||||
assert.Equal(t, "llm_policy.model_not_routable", out.DenyReason.Code, "deny code stays model_not_routable")
|
||||
}
|
||||
|
||||
// TestRouter_ExplicitModelBeatsCatchallGateway is the regression guard
|
||||
// for multi-provider misrouting: a catch-all (empty Models) OpenAI-compat
|
||||
// gateway declared first must NOT swallow a model an explicit provider
|
||||
// claims. Anthropic's claude request must reach the Anthropic route even
|
||||
// though the gateway claims every model and wins declaration order.
|
||||
func TestRouter_ExplicitModelBeatsCatchallGateway(t *testing.T) {
|
||||
gateway := ProviderRoute{
|
||||
ID: "openai-gateway",
|
||||
Models: nil, // catch-all: claims every model
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}
|
||||
anthropic := ProviderRoute{
|
||||
ID: "anthropic-prod",
|
||||
Models: []string{"claude-opus-4"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.anthropic.com",
|
||||
}
|
||||
// Gateway declared first to prove explicit claim beats declaration order.
|
||||
mw := New(Config{Providers: []ProviderRoute{gateway, anthropic}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("claude-opus-4", "/v1/messages"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "explicit-model request must route, not deny")
|
||||
require.NotNil(t, out.Mutations)
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "api.anthropic.com", out.Mutations.RewriteUpstream.Host, "claude must reach the explicit Anthropic route, not the catch-all gateway")
|
||||
|
||||
provider, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "anthropic-prod", provider, "resolved provider must be the explicit Anthropic route")
|
||||
}
|
||||
|
||||
// TestRouter_CatchallStillServesUnlistedModel confirms the catch-all
|
||||
// gateway still wins models no explicit provider claims (its whole point).
|
||||
func TestRouter_CatchallStillServesUnlistedModel(t *testing.T) {
|
||||
gateway := ProviderRoute{
|
||||
ID: "openai-gateway",
|
||||
Models: nil,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "gateway.example.com",
|
||||
}
|
||||
anthropic := ProviderRoute{
|
||||
ID: "anthropic-prod",
|
||||
Models: []string{"claude-opus-4"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.anthropic.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{gateway, anthropic}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("some-exotic-model", "/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "unlisted model must still route via the catch-all")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "gateway.example.com", out.Mutations.RewriteUpstream.Host, "unlisted model falls to the catch-all gateway")
|
||||
}
|
||||
|
||||
// newInputVendorModelURL returns an Input carrying both the detected
|
||||
// vendor (llm.provider) and the model, plus a request URL — mimicking the
|
||||
// post-llm_request_parser state for a real inference call.
|
||||
func newInputVendorModelURL(vendor, model, reqURL string) *middleware.Input {
|
||||
return &middleware.Input{
|
||||
Slot: middleware.SlotOnRequest,
|
||||
URL: reqURL,
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: vendor},
|
||||
{Key: middleware.KeyLLMModel, Value: model},
|
||||
},
|
||||
UserGroups: []string{defaultTestGroup},
|
||||
}
|
||||
}
|
||||
|
||||
// TestRouter_VendorKeepsAnthropicOffOpenAIGateway is the regression guard
|
||||
// for the reported multi-provider break: two catch-all providers (neither
|
||||
// enumerates models), the OpenAI one declared first. Without vendor
|
||||
// awareness, a claude request matches both, no path prefixes, and
|
||||
// declaration order sends it to OpenAI → 502. The detected vendor must
|
||||
// pin it to the Anthropic route.
|
||||
func TestRouter_VendorKeepsAnthropicOffOpenAIGateway(t *testing.T) {
|
||||
openai := ProviderRoute{
|
||||
ID: "openai-gw",
|
||||
Vendor: "openai",
|
||||
Models: nil, // catch-all
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}
|
||||
anthropic := ProviderRoute{
|
||||
ID: "anthropic-gw",
|
||||
Vendor: "anthropic",
|
||||
Models: nil, // catch-all
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.anthropic.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{openai, anthropic}}) // openai first
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputVendorModelURL("anthropic", "claude-opus-4-8", "/v1/messages"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "claude request must route, not deny")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "api.anthropic.com", out.Mutations.RewriteUpstream.Host, "anthropic vendor must pin to the anthropic route despite openai being declared first")
|
||||
|
||||
provider, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "anthropic-gw", provider)
|
||||
}
|
||||
|
||||
// TestRouter_VendorKeepsOpenAIOffAnthropic is the reciprocal: an OpenAI
|
||||
// request must stay on the OpenAI route even when the Anthropic catch-all
|
||||
// is declared first.
|
||||
func TestRouter_VendorKeepsOpenAIOffAnthropic(t *testing.T) {
|
||||
anthropic := ProviderRoute{
|
||||
ID: "anthropic-gw",
|
||||
Vendor: "anthropic",
|
||||
Models: nil,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.anthropic.com",
|
||||
}
|
||||
openai := ProviderRoute{
|
||||
ID: "openai-gw",
|
||||
Vendor: "openai",
|
||||
Models: nil,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{anthropic, openai}}) // anthropic first
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputVendorModelURL("openai", "gpt-5.5", "/v1/responses"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "api.openai.com", out.Mutations.RewriteUpstream.Host, "openai vendor must pin to the openai route despite anthropic being declared first")
|
||||
}
|
||||
|
||||
// TestRouter_VendorAbsentFallsBackToModelPath confirms vendor filtering is
|
||||
// inert when the request carries no detected vendor: routing then relies on
|
||||
// model/path as before.
|
||||
func TestRouter_VendorAbsentFallsBackToModelPath(t *testing.T) {
|
||||
openai := ProviderRoute{
|
||||
ID: "openai-gw",
|
||||
Vendor: "openai",
|
||||
Models: []string{"gpt-5.5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{openai}})
|
||||
|
||||
// No llm.provider in metadata — only the model.
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-5.5"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "explicit-model match must still route with no vendor present")
|
||||
}
|
||||
|
||||
func TestRouter_UnknownModel(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("claude-opus-4"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "unrouted model must deny")
|
||||
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_not_routable", out.DenyReason.Code, "deny code must be model_not_routable")
|
||||
assert.Equal(t, "no provider configured for model claude-opus-4", out.DenyReason.Message, "deny message must reference the offending model")
|
||||
assert.Equal(t, "claude-opus-4", out.DenyReason.Details["model"], "deny details must include the offending model")
|
||||
}
|
||||
|
||||
func TestRouter_HeaderStripList(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-prod",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer sk-test-123",
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-4o"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
require.NotNil(t, out.Mutations, "matched route must emit mutations")
|
||||
|
||||
expected := []string{
|
||||
"Authorization",
|
||||
"Proxy-Authorization",
|
||||
"x-api-key",
|
||||
"api-key",
|
||||
}
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream, "matched route must emit upstream rewrite")
|
||||
for _, header := range expected {
|
||||
assert.Contains(t, out.Mutations.RewriteUpstream.StripHeaders, header,
|
||||
"strip list (on UpstreamRewrite) must include the well-known vendor auth header %s", header)
|
||||
}
|
||||
|
||||
// Vendor metadata headers MUST NOT be stripped: the client SDK sets them
|
||||
// and the upstream requires them. Anthropic returns 400 "anthropic-version:
|
||||
// header is required" if we drop it. Lock the regression.
|
||||
preserved := []string{"anthropic-version", "openai-organization", "openai-project"}
|
||||
for _, header := range preserved {
|
||||
assert.NotContains(t, out.Mutations.RewriteUpstream.StripHeaders, header,
|
||||
"vendor metadata header %s must NOT be stripped — upstreams require it", header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_FirstMatchWins(t *testing.T) {
|
||||
first := ProviderRoute{
|
||||
ID: "first",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "first.test",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer first",
|
||||
}
|
||||
second := ProviderRoute{
|
||||
ID: "second",
|
||||
Models: []string{"gpt-4o"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "second.test",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer second",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{first, second}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-4o"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "duplicate-model match must still allow")
|
||||
require.NotNil(t, out.Mutations, "matched route must emit mutations")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream, "matched route must emit upstream rewrite")
|
||||
assert.Equal(t, "first.test", out.Mutations.RewriteUpstream.Host, "first-match-wins must pick the earlier route")
|
||||
|
||||
resolved, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "first", resolved, "resolved provider id must be the earlier route's ID")
|
||||
}
|
||||
|
||||
// TestRouter_PathDisambiguation_PrefixWinsOverCatchall locks in the
|
||||
// rule the user nailed down: two providers claim the same model, one
|
||||
// has an UpstreamPath that prefixes the incoming URL, the other has
|
||||
// no path. The path-prefixed provider wins because the path is a
|
||||
// strictly more specific match than the empty catchall.
|
||||
func TestRouter_PathDisambiguation_PrefixWinsOverCatchall(t *testing.T) {
|
||||
corp := ProviderRoute{
|
||||
ID: "corp-openai-compat",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "corp.example.com",
|
||||
UpstreamPath: "/openai",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer corp",
|
||||
}
|
||||
openai := ProviderRoute{
|
||||
ID: "openai",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer openai",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{openai, corp}}) // openai listed first to prove path beats declaration order
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("gpt-5", "/openai/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision, "path-prefix match must allow")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "corp.example.com", out.Mutations.RewriteUpstream.Host,
|
||||
"path-prefixed provider must beat the catchall when its UpstreamPath is a prefix of the request path")
|
||||
resolved, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "corp-openai-compat", resolved, "resolved provider id must reflect the path-prefix winner, not the first declared")
|
||||
}
|
||||
|
||||
// TestRouter_PathDisambiguation_CatchallWhenNoPrefixMatches is the
|
||||
// inverse: the path-prefixed provider does NOT match the incoming
|
||||
// path, so the empty-path catchall takes the request.
|
||||
func TestRouter_PathDisambiguation_CatchallWhenNoPrefixMatches(t *testing.T) {
|
||||
corp := ProviderRoute{
|
||||
ID: "corp-openai-compat",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "corp.example.com",
|
||||
UpstreamPath: "/openai",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer corp",
|
||||
}
|
||||
openai := ProviderRoute{
|
||||
ID: "openai",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer openai",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{corp, openai}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("gpt-5", "/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision, "catchall must allow when no path prefix matches")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "api.openai.com", out.Mutations.RewriteUpstream.Host,
|
||||
"empty-path catchall must win when the path-prefixed provider's UpstreamPath does not match the request")
|
||||
resolved, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "openai", resolved, "resolved provider id must be the catchall")
|
||||
}
|
||||
|
||||
// TestRouter_PathDisambiguation_LongestPrefixWins covers the case
|
||||
// where multiple providers have non-empty UpstreamPath values that
|
||||
// both prefix the request — the longer (more specific) one wins.
|
||||
func TestRouter_PathDisambiguation_LongestPrefixWins(t *testing.T) {
|
||||
short := ProviderRoute{
|
||||
ID: "short-prefix",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "short.example.com",
|
||||
UpstreamPath: "/openai",
|
||||
}
|
||||
long := ProviderRoute{
|
||||
ID: "long-prefix",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "long.example.com",
|
||||
UpstreamPath: "/openai/v1",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{short, long}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("gpt-5", "/openai/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "long.example.com", out.Mutations.RewriteUpstream.Host,
|
||||
"longest matching UpstreamPath must win — most specific match")
|
||||
}
|
||||
|
||||
// TestRouter_SingleMatchIgnoresPath proves the path-prefix rule is a
|
||||
// disambiguation pass, not a gate: when only one provider claims the
|
||||
// model, it wins regardless of UpstreamPath. Otherwise a path-scoped
|
||||
// provider would 403 every request whose URL doesn't include the
|
||||
// path, which would break SDKs configured to hit the gateway root.
|
||||
func TestRouter_SingleMatchIgnoresPath(t *testing.T) {
|
||||
only := ProviderRoute{
|
||||
ID: "only",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "only.example.com",
|
||||
UpstreamPath: "/openai",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer only",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{only}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("gpt-5", "/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision,
|
||||
"single model-matching provider must serve the request even when UpstreamPath doesn't prefix the URL — path is a tiebreaker, not a gate")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "only.example.com", out.Mutations.RewriteUpstream.Host, "the only model-matching provider should be selected")
|
||||
}
|
||||
|
||||
// TestRouter_PathDisambiguation_FallbackWhenNoPrefixMatches covers
|
||||
// the multi-candidate edge case where every candidate has a
|
||||
// non-matching non-empty UpstreamPath. The router falls back to
|
||||
// declaration order so the model is still routable rather than 403'd.
|
||||
func TestRouter_PathDisambiguation_FallbackWhenNoPrefixMatches(t *testing.T) {
|
||||
first := ProviderRoute{
|
||||
ID: "first",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "first.example.com",
|
||||
UpstreamPath: "/openai",
|
||||
}
|
||||
second := ProviderRoute{
|
||||
ID: "second",
|
||||
Models: []string{"gpt-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "second.example.com",
|
||||
UpstreamPath: "/anthropic",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{first, second}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModelAndURL("gpt-5", "/v1/chat/completions"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "no path match among multi-candidates must still allow")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "first.example.com", out.Mutations.RewriteUpstream.Host,
|
||||
"when no candidate's UpstreamPath prefix-matches the request, fall back to declaration order")
|
||||
}
|
||||
|
||||
func TestRouter_FactoryRejectsBadJSON(t *testing.T) {
|
||||
_, err := Factory{}.New([]byte("{not json"))
|
||||
require.Error(t, err, "malformed JSON config must be rejected at chain build time")
|
||||
}
|
||||
|
||||
func TestRouter_FactoryAcceptsEmptyShapes(t *testing.T) {
|
||||
cases := [][]byte{nil, []byte(""), []byte(" "), []byte("null"), []byte("{}"), []byte("[]")}
|
||||
for _, raw := range cases {
|
||||
mw, err := Factory{}.New(raw)
|
||||
require.NoError(t, err, "empty-shaped config must yield a router with an empty Providers slice")
|
||||
require.NotNil(t, mw, "factory must return a non-nil middleware on empty config")
|
||||
|
||||
out, invErr := mw.Invoke(context.Background(), newInputWithModel("gpt-4o"))
|
||||
require.NoError(t, invErr)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision,
|
||||
"router with no providers must deny every model as not-routable")
|
||||
}
|
||||
}
|
||||
|
||||
// newInputWithModelAndGroups returns an Input carrying llm.model + the
|
||||
// caller's UserGroups, mimicking the post-auth, post-llm_request_parser
|
||||
// state the router observes.
|
||||
func newInputWithModelAndGroups(model string, groups []string) *middleware.Input {
|
||||
in := newInputWithModel(model)
|
||||
in.UserGroups = append([]string(nil), groups...)
|
||||
return in
|
||||
}
|
||||
|
||||
// TestRouter_GroupFilter_PicksAuthorisedAmongDuplicates pins the Fix A
|
||||
// behaviour: when two providers claim the same model but each
|
||||
// authorises a different group, the router must pick the route the
|
||||
// caller's groups intersect, regardless of declaration order.
|
||||
func TestRouter_GroupFilter_PicksAuthorisedAmongDuplicates(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{
|
||||
{
|
||||
ID: "openai-marketing",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "mkt-openai.example.com",
|
||||
AllowedGroupIDs: []string{"grp-mkt"},
|
||||
},
|
||||
{
|
||||
ID: "openai-engineering",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "eng-openai.example.com",
|
||||
AllowedGroupIDs: []string{"grp-eng"},
|
||||
},
|
||||
}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(),
|
||||
newInputWithModelAndGroups("gpt-4o-mini", []string{"grp-eng"}))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision,
|
||||
"authorised candidate exists; must allow")
|
||||
|
||||
resolved, ok := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "openai-engineering", resolved,
|
||||
"router must pick the route whose AllowedGroupIDs intersects the caller's groups, ignoring declaration order")
|
||||
}
|
||||
|
||||
// TestRouter_GroupFilter_NoIntersection_DeniesNoAuthorisedRoute pins
|
||||
// the dedicated deny code that fires when the model is known to a
|
||||
// provider but no candidate is authorised for the caller's groups.
|
||||
func TestRouter_GroupFilter_NoIntersection_DeniesNoAuthorisedRoute(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-marketing",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "mkt-openai.example.com",
|
||||
AllowedGroupIDs: []string{"grp-mkt"},
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(),
|
||||
newInputWithModelAndGroups("gpt-4o-mini", []string{"grp-eng"}))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision,
|
||||
"model exists but no route authorises grp-eng; must deny")
|
||||
require.NotNil(t, out.DenyReason)
|
||||
assert.Equal(t, "llm_policy.no_authorised_provider", out.DenyReason.Code,
|
||||
"deny code must be no_authorised_provider, not model_not_routable")
|
||||
assert.Equal(t, "gpt-4o-mini", out.DenyReason.Details["model"],
|
||||
"deny details must reference the offending model")
|
||||
|
||||
dec, _ := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision)
|
||||
assert.Equal(t, "deny", dec)
|
||||
reason, _ := metaValue(t, out.Metadata, middleware.KeyLLMPolicyReason)
|
||||
assert.Equal(t, "no_authorised_provider", reason)
|
||||
}
|
||||
|
||||
// TestRouter_GroupFilter_EmptyAllowedGroupsIsUnreachable pins the
|
||||
// strict semantics: a route with no AllowedGroupIDs is unreachable.
|
||||
// The synthesiser only emits policy-bound routes, so an empty ACL
|
||||
// signals a misconfiguration that must not silently fall through.
|
||||
func TestRouter_GroupFilter_EmptyAllowedGroupsIsUnreachable(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-shared",
|
||||
Models: []string{"gpt-4o"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "api.openai.com",
|
||||
// AllowedGroupIDs intentionally left empty.
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-4o"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision,
|
||||
"empty AllowedGroupIDs must deny — there is no catch-all for routes without an authorising policy")
|
||||
require.NotNil(t, out.DenyReason)
|
||||
assert.Equal(t, "llm_policy.no_authorised_provider", out.DenyReason.Code,
|
||||
"empty ACL fails the group-filter pass; deny code must reflect that")
|
||||
}
|
||||
|
||||
// TestRouter_GroupFilter_OverlapTiebreakUnchanged pins that when more
|
||||
// than one route is authorised for the caller's groups, the existing
|
||||
// path-prefix tiebreak still decides. Group filtering is a hard gate
|
||||
// before the tiebreak; it does not change the tiebreak semantics.
|
||||
func TestRouter_GroupFilter_OverlapTiebreakUnchanged(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{
|
||||
{
|
||||
ID: "openai-a",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "a.example.com",
|
||||
UpstreamPath: "",
|
||||
AllowedGroupIDs: []string{"grp-eng"},
|
||||
},
|
||||
{
|
||||
ID: "openai-b",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "b.example.com",
|
||||
UpstreamPath: "/v1/chat",
|
||||
AllowedGroupIDs: []string{"grp-eng"},
|
||||
},
|
||||
}})
|
||||
|
||||
in := newInputWithModelAndURL("gpt-4o-mini", "/v1/chat/completions")
|
||||
in.UserGroups = []string{"grp-eng"}
|
||||
out, err := mw.Invoke(context.Background(), in)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision)
|
||||
|
||||
resolved, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "openai-b", resolved,
|
||||
"longest-prefix path tiebreak still wins among group-authorised candidates")
|
||||
}
|
||||
|
||||
// TestRouter_AuthorisingGroups_EmitsIntersection pins that the router
|
||||
// emits llm.authorising_groups containing only the intersection of the
|
||||
// caller's UserGroups with the resolved route's AllowedGroupIDs — not
|
||||
// every group the peer happens to be in.
|
||||
func TestRouter_AuthorisingGroups_EmitsIntersection(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "openai-eng",
|
||||
Models: []string{"gpt-4o-mini"},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "eng-openai.example.com",
|
||||
AllowedGroupIDs: []string{"grp-eng", "grp-shared"},
|
||||
}}})
|
||||
|
||||
in := newInputWithModelAndGroups("gpt-4o-mini",
|
||||
[]string{"grp-eng", "grp-it", "grp-shared", "grp-oncall"})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), in)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision)
|
||||
|
||||
csv, ok := metaValue(t, out.Metadata, middleware.KeyLLMAuthorisingGroups)
|
||||
require.True(t, ok, "router must emit llm.authorising_groups on a match")
|
||||
assert.Equal(t, "grp-eng,grp-shared", csv,
|
||||
"only groups in BOTH UserGroups AND AllowedGroupIDs may appear; result must be sorted and unique")
|
||||
}
|
||||
|
||||
// TestRouter_EmptyModelsClaimsAnyModel pins that a route with no
|
||||
// configured Models matches every model — used by gateway-style
|
||||
// providers (LiteLLM, custom OpenAI-compatible endpoints) where the
|
||||
// operator can't enumerate the upstream's model catalog in NetBird.
|
||||
func TestRouter_EmptyModelsClaimsAnyModel(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{{
|
||||
ID: "litellm",
|
||||
Models: nil, // catch-all
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "litellm.example.com",
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
}}})
|
||||
|
||||
out, err := mw.Invoke(context.Background(), newInputWithModel("gpt-5.5"))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision,
|
||||
"a route with empty Models must claim any model so gateway-style providers can route open-ended sets")
|
||||
resolved, _ := metaValue(t, out.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
assert.Equal(t, "litellm", resolved)
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package llm_router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
)
|
||||
|
||||
// pathRoutedInput builds an Input mimicking the post-llm_request_parser state
|
||||
// for a path-routed (Vertex/Bedrock) request: a request URL plus the model and
|
||||
// (optionally) provider/vendor metadata the parser emits.
|
||||
func pathRoutedInput(url, provider, model string) *middleware.Input {
|
||||
md := []middleware.KV{{Key: middleware.KeyLLMModel, Value: model}}
|
||||
if provider != "" {
|
||||
md = append(md, middleware.KV{Key: middleware.KeyLLMProvider, Value: provider})
|
||||
}
|
||||
return &middleware.Input{
|
||||
Slot: middleware.SlotOnRequest,
|
||||
URL: url,
|
||||
Metadata: md,
|
||||
UserGroups: []string{defaultTestGroup},
|
||||
}
|
||||
}
|
||||
|
||||
func vertexRoute() ProviderRoute {
|
||||
return ProviderRoute{
|
||||
ID: "vertex-prod", Vertex: true,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "europe-west1-aiplatform.googleapis.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer x",
|
||||
}
|
||||
}
|
||||
|
||||
// A Vertex publisher with no parser surface (google/gemini emits no
|
||||
// llm.provider) must be denied, not forwarded unmetered.
|
||||
func TestRouter_VertexUnmeterablePublisherDenied(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{vertexRoute()}})
|
||||
in := pathRoutedInput(
|
||||
"/v1/projects/p/locations/global/publishers/google/models/gemini-2.5-pro:generateContent",
|
||||
"", // google -> request parser emits NO llm.provider
|
||||
"gemini-2.5-pro",
|
||||
)
|
||||
out, err := mw.Invoke(context.Background(), in)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "unmeterable Vertex publisher must deny")
|
||||
assert.Equal(t, 403, out.DenyStatus, "unmeterable deny is a 403")
|
||||
require.NotNil(t, out.DenyReason)
|
||||
assert.Equal(t, denyCodeUnmeterable, out.DenyReason.Code, "deny code must flag the unmeterable publisher")
|
||||
}
|
||||
|
||||
// A Vertex publisher with a parser surface (anthropic) is allowed.
|
||||
func TestRouter_VertexMeterablePublisherAllowed(t *testing.T) {
|
||||
mw := New(Config{Providers: []ProviderRoute{vertexRoute()}})
|
||||
in := pathRoutedInput(
|
||||
"/v1/projects/p/locations/global/publishers/anthropic/models/claude-sonnet-4-5:rawPredict",
|
||||
"anthropic",
|
||||
"claude-sonnet-4-5",
|
||||
)
|
||||
out, err := mw.Invoke(context.Background(), in)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "meterable Vertex publisher must allow")
|
||||
}
|
||||
|
||||
// A path-routed provider with an explicit Models list must reject models not in
|
||||
// the list (the provider credential can't be used for unauthorised models).
|
||||
func TestRouter_PathRoutedModelAllowlistEnforced(t *testing.T) {
|
||||
route := ProviderRoute{
|
||||
ID: "bedrock-prod", Bedrock: true,
|
||||
Models: []string{"anthropic.claude-sonnet-4-5"},
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer x",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{route}})
|
||||
|
||||
allowed := pathRoutedInput(
|
||||
"/model/eu.anthropic.claude-sonnet-4-5-20250929-v1:0/invoke",
|
||||
"bedrock", "anthropic.claude-sonnet-4-5",
|
||||
)
|
||||
out, err := mw.Invoke(context.Background(), allowed)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "model in the allowlist must be served")
|
||||
|
||||
denied := pathRoutedInput(
|
||||
"/model/amazon.nova-pro-v1:0/invoke",
|
||||
"bedrock", "amazon.nova-pro",
|
||||
)
|
||||
out, err = mw.Invoke(context.Background(), denied)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionDeny, out.Decision, "model outside the allowlist must deny")
|
||||
require.NotNil(t, out.DenyReason)
|
||||
assert.Equal(t, denyCodeNotRoutable, out.DenyReason.Code, "unlisted model denies as not-routable")
|
||||
}
|
||||
|
||||
// A "/bedrock" gateway-namespace prefix routes the same as the native path and
|
||||
// records the prefix on the rewrite so the proxy strips it before forwarding.
|
||||
func TestRouter_BedrockNamespacePrefixStripped(t *testing.T) {
|
||||
route := ProviderRoute{
|
||||
ID: "bedrock-prod", Bedrock: true,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer x",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{route}})
|
||||
|
||||
prefixed := pathRoutedInput(
|
||||
"/bedrock/model/eu.anthropic.claude-sonnet-4-5-20250929-v1:0/invoke-with-response-stream",
|
||||
"bedrock", "anthropic.claude-sonnet-4-5",
|
||||
)
|
||||
out, err := mw.Invoke(context.Background(), prefixed)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision, "prefixed Bedrock path must route")
|
||||
require.NotNil(t, out.Mutations)
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Equal(t, "/bedrock", out.Mutations.RewriteUpstream.StripPathPrefix,
|
||||
"namespace prefix must be recorded so the proxy strips it before forwarding")
|
||||
|
||||
native := pathRoutedInput(
|
||||
"/model/eu.anthropic.claude-sonnet-4-5-20250929-v1:0/invoke",
|
||||
"bedrock", "anthropic.claude-sonnet-4-5",
|
||||
)
|
||||
out, err = mw.Invoke(context.Background(), native)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision, "native Bedrock path must route")
|
||||
require.NotNil(t, out.Mutations.RewriteUpstream)
|
||||
assert.Empty(t, out.Mutations.RewriteUpstream.StripPathPrefix,
|
||||
"native path carries no namespace prefix to strip")
|
||||
}
|
||||
|
||||
// A path-routed provider with no configured Models is catch-all: any model the
|
||||
// credential can reach is served (preserves the zero-config behaviour).
|
||||
func TestRouter_PathRoutedCatchAllServesAnyModel(t *testing.T) {
|
||||
route := ProviderRoute{
|
||||
ID: "bedrock-catchall", Bedrock: true,
|
||||
AllowedGroupIDs: []string{defaultTestGroup},
|
||||
UpstreamScheme: "https",
|
||||
UpstreamHost: "bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer x",
|
||||
}
|
||||
mw := New(Config{Providers: []ProviderRoute{route}})
|
||||
in := pathRoutedInput(
|
||||
"/model/amazon.nova-pro-v1:0/invoke",
|
||||
"bedrock", "amazon.nova-pro",
|
||||
)
|
||||
out, err := mw.Invoke(context.Background(), in)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "catch-all path-routed provider serves any model")
|
||||
}
|
||||
Reference in New Issue
Block a user