Merge remote-tracking branch 'origin/main' into HEAD

# Conflicts:
#	shared/management/proto/proxy_service.pb.go
This commit is contained in:
Viktor Liu
2026-09-23 07:46:09 +02:00
1091 changed files with 82476 additions and 19116 deletions
@@ -0,0 +1,300 @@
package agentnetwork
import (
"context"
"fmt"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/store"
)
// GetAgentConfigForUser returns the Agent Network setup the calling user's
// groups authorize. It deliberately performs no role permission check:
// the result is scoped to the caller's own groups, which is strictly
// tighter than any role gate, so every authenticated user (any role) may
// read it. The group source matches enforcement: the proxy authorizes
// each Agent Network request against the calling user's groups as well —
// session validation resolves them from the same user record's
// auto-groups — so this answer and the proxy's verdict are computed from
// the same memberships.
func (m *managerImpl) GetAgentConfigForUser(ctx context.Context, accountID, userID string) (*types.AgentConfig, error) {
user, err := m.store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
if err != nil {
return nil, fmt.Errorf("get user: %w", err)
}
return m.agentConfigForGroups(ctx, accountID, user.AutoGroups)
}
// agentConfigForGroups computes the effective Agent Network setup for
// a set of caller groups: the account endpoint plus, per authorized
// provider, the effective model set. It mirrors what the proxy enforces
// at request time — the policy filter matches filterApplicablePolicies,
// the model logic matches policyPermitsModel, and orphan providers
// (enabled but referenced by no applicable policy) are omitted just like
// the router synthesizer omits them — so the answer never advertises
// anything the proxy would refuse.
//
// Configured tracks the account, not the caller: once the account has an
// endpoint every member gets it, with Providers empty for those no policy
// covers yet. The dashboard shows each user the same connection config
// regardless of role, and an empty provider list tells them to ask for
// access. Only the account having no Agent Network at all reads as not
// configured. Providers stays caller-scoped either way — the endpoint on
// its own authorizes nothing, and the proxy still refuses every request
// no policy permits.
func (m *managerImpl) agentConfigForGroups(ctx context.Context, accountID string, groupIDs []string) (*types.AgentConfig, error) {
notConfigured := &types.AgentConfig{Providers: []types.AgentConfigProvider{}}
settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
switch {
case err == nil:
case isNotFound(err):
return notConfigured, nil
default:
return nil, fmt.Errorf("get agent network settings: %w", err)
}
if settings.Endpoint() == "" {
return notConfigured, nil
}
authorized, applicable, err := m.authorizedProvidersForGroups(ctx, accountID, groupIDs)
if err != nil {
return nil, err
}
out := &types.AgentConfig{
Configured: true,
Endpoint: "https://" + settings.Endpoint(),
Providers: make([]types.AgentConfigProvider, 0, len(authorized)),
}
if len(authorized) == 0 {
return out, nil
}
var guardrailsByID map[string]*types.Guardrail
if anyPolicyHasGuardrails(applicable) {
guardrailsByID, err = m.loadGuardrailsByID(ctx, accountID)
if err != nil {
return nil, err
}
}
for _, p := range authorized {
allAllowed, models := effectiveModelsForProvider(p, policiesForProvider(applicable, p.ID), guardrailsByID)
flavor := ""
if entry, ok := catalog.Lookup(p.ProviderID); ok {
flavor = entry.ParserID
}
out.Providers = append(out.Providers, types.AgentConfigProvider{
Name: p.Name,
CatalogID: p.ProviderID,
APIFlavor: flavor,
AllModelsAllowed: allAllowed,
Models: models,
})
}
return out, nil
}
// authorizedProvidersForGroups returns the enabled providers referenced
// by at least one enabled policy whose source groups intersect groupIDs —
// the providers the caller's own policies authorize — in created_at order
// with ID tiebreak, the same deterministic order the router synthesizer
// presents. The applicable policies come back alongside so callers that
// need per-provider policy context (the setup's model computation) don't
// re-filter. Both the self-service setup answer and the caller-scoped
// provider list are built from this selection, so what the dashboard
// offers and what the proxy enforces never diverge.
func (m *managerImpl) authorizedProvidersForGroups(ctx context.Context, accountID string, groupIDs []string) ([]*types.Provider, []*types.Policy, error) {
policies, err := m.store.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID)
if err != nil {
return nil, nil, fmt.Errorf("list account policies: %w", err)
}
applicable := filterPoliciesByGroups(policies, groupIDs)
if len(applicable) == 0 {
return nil, nil, nil
}
providers, err := m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID)
if err != nil {
return nil, nil, fmt.Errorf("list account providers: %w", err)
}
// filterEnabledProviders carries the enabled filter and the
// created_at/ID order shared with the router synthesizer.
enabled := filterEnabledProviders(providers)
authorized := make([]*types.Provider, 0, len(enabled))
for _, p := range enabled {
if len(policiesForProvider(applicable, p.ID)) == 0 {
continue
}
authorized = append(authorized, p)
}
return authorized, applicable, nil
}
// filterPoliciesByGroups returns the enabled policies whose SourceGroups
// intersect the caller's groups. Same group matching as
// filterApplicablePolicies, without the per-provider filter — the setup
// answer spans every provider the caller can reach.
func filterPoliciesByGroups(policies []*types.Policy, groupIDs []string) []*types.Policy {
groupSet := make(map[string]struct{}, len(groupIDs))
for _, g := range groupIDs {
if g != "" {
groupSet[g] = struct{}{}
}
}
out := make([]*types.Policy, 0, len(policies))
for _, p := range policies {
if p == nil || !p.Enabled {
continue
}
if !anyGroupMatches(p.SourceGroups, groupSet) {
continue
}
out = append(out, p)
}
return out
}
// policiesForProvider returns the subset of policies targeting the
// provider, order preserved.
func policiesForProvider(policies []*types.Policy, providerID string) []*types.Policy {
out := make([]*types.Policy, 0, len(policies))
for _, p := range policies {
if sliceContains(p.DestinationProviderIDs, providerID) {
out = append(out, p)
}
}
return out
}
// effectiveModelsForProvider derives the caller's effective model set for
// one provider from the applicable policies that target it, mirroring
// policyPermitsModel: a policy with no allowlist-enabled guardrail is
// unrestricted, and one unrestricted policy makes the whole provider
// unrestricted (the proxy would admit any model through it). Otherwise
// the union of the policies' allowlists applies, intersected with the
// provider's declared models when the operator declared any — the router
// only claims declared models, so an allowlisted-but-undeclared model is
// unreachable and must not be advertised. With no declared models the
// router claims every model, so the allowlist union stands alone.
// Allowlist entries and declared ids both compare through the canonical
// id the proxy's parser emits, so an allowlist may hold either form: the
// raw declared id the dashboard's picker copies from the provider, or
// the stripped id the parser matches at request time.
func effectiveModelsForProvider(provider *types.Provider, policies []*types.Policy, guardrailsByID map[string]*types.Guardrail) (bool, []string) {
restricted := true
union := make([]string, 0)
seen := make(map[string]struct{})
for _, p := range policies {
policyRestricted := false
for _, gID := range p.GuardrailIDs {
g, ok := guardrailsByID[gID]
if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled {
continue
}
policyRestricted = true
for _, model := range g.Checks.ModelAllowlist.Models {
key := canonicalModelKey(provider.ProviderID, model)
if key == "" {
continue
}
if _, dup := seen[key]; dup {
continue
}
seen[key] = struct{}{}
union = append(union, key)
}
}
if !policyRestricted {
restricted = false
}
}
declared := declaredModelIDs(provider)
if !restricted {
return true, declared
}
if len(provider.Models) == 0 {
// No operator declaration: the router claims every model, so the
// allowlist union is the effective set as-is.
return false, union
}
out := make([]string, 0, len(declared))
for _, id := range declared {
// Compare through the canonical id the proxy's parser emits — a
// Bedrock declaration may carry the region/version form
// ("eu.anthropic.claude-...-v1:0") that the parser strips at
// request time, and the raw forms would never intersect. The
// declared id itself is what gets advertised, matching the
// router's route claim.
if _, ok := seen[canonicalModelKey(provider.ProviderID, id)]; ok {
out = append(out, id)
}
}
return false, out
}
// canonicalModelKey builds the compare key for a model id: lowercased,
// trimmed, and canonicalized through the provider-aware normalization the
// proxy's parser applies. Lowercase/trim comes FIRST — the path-style
// strippers anchor on a lowercase id's tail, so a trailing space or a
// case-variant geography/version would otherwise survive into the key.
func canonicalModelKey(catalogProviderID, id string) string {
return normaliseModelID(normalizePricingModelID(catalogProviderID, normaliseModelID(id)))
}
// providerModelsByID maps effective model ids (as effectiveModelsForProvider
// returns them) back onto the operator's declared entries, keeping the
// declared casing and prices. With no operator declaration the ids are the
// allowlist union and have no declared entry to map to, so bare entries are
// synthesized — the router claims every model in that case, so those ids are
// reachable and belong in the answer.
func providerModelsByID(provider *types.Provider, ids []string) []types.ProviderModel {
if len(provider.Models) == 0 {
out := make([]types.ProviderModel, 0, len(ids))
for _, id := range ids {
out = append(out, types.ProviderModel{ID: id})
}
return out
}
keep := make(map[string]struct{}, len(ids))
for _, id := range ids {
keep[normaliseModelID(id)] = struct{}{}
}
out := make([]types.ProviderModel, 0, len(ids))
for _, m := range provider.Models {
if _, ok := keep[normaliseModelID(m.ID)]; ok {
out = append(out, m)
}
}
return out
}
// declaredModelIDs returns the models a provider exposes: the operator's
// curated list when present, otherwise the catalog entry's models (an
// empty operator list means "all catalog models"). Gateway/custom catalog
// entries declare no models, so the result may be empty.
func declaredModelIDs(provider *types.Provider) []string {
if ids := providerModelIDs(provider); len(ids) > 0 {
return ids
}
entry, ok := catalog.Lookup(provider.ProviderID)
if !ok {
return []string{}
}
out := make([]string, 0, len(entry.Models))
for _, m := range entry.Models {
if m.ID != "" {
out = append(out, m.ID)
}
}
return out
}
// GetAgentConfigForUser on the mock manager reports "not configured" so tests
// that don't care about setup still compile.
func (*mockManager) GetAgentConfigForUser(_ context.Context, _, _ string) (*types.AgentConfig, error) {
return &types.AgentConfig{Providers: []types.AgentConfigProvider{}}, nil
}
@@ -0,0 +1,406 @@
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/store"
nbtypes "github.com/netbirdio/netbird/management/server/types"
)
// These tests drive the effective-setup computation through the real
// sqlite store, mirroring the policyselect realstore suite: assert on
// observable answers (configured / providers / models), not on which
// store methods get called. The computation must agree with what the
// proxy enforces — policy filtering matches filterApplicablePolicies,
// model logic matches policyPermitsModel, and orphan providers are
// omitted like the router synthesizer omits them.
func newAgentConfigTestMgr(t *testing.T) (*managerImpl, store.Store) {
t.Helper()
ctx := context.Background()
s, cleanup, err := store.NewTestStoreFromSQL(ctx, "", t.TempDir())
require.NoError(t, err, "real sqlite test store must come up")
t.Cleanup(cleanup)
return &managerImpl{store: s}, s
}
// newSetupTestGuardrail returns an allowlist-enabled guardrail.
func newSetupTestGuardrail(id string, models ...string) *types.Guardrail {
return &types.Guardrail{
ID: id,
AccountID: testAccountID,
Name: "allowlist " + id,
Checks: types.GuardrailChecks{
ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true, Models: models},
},
}
}
func TestAgentConfig_RealStore_NoSettingsRow(t *testing.T) {
mgr, _ := newAgentConfigTestMgr(t)
setup, err := mgr.agentConfigForGroups(context.Background(), testAccountID, []string{"grp-eng"})
require.NoError(t, err)
assert.False(t, setup.Configured, "account without settings must read as not configured")
assert.Empty(t, setup.Endpoint)
assert.Empty(t, setup.Providers)
}
func TestAgentConfig_RealStore_NoApplicablePolicy(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-other"})
require.NoError(t, err)
assert.True(t, setup.Configured, "the account is set up, so every member reads as configured")
assert.Equal(t, "https://"+testEndpoint, setup.Endpoint, "every member gets the same connection config")
assert.Empty(t, setup.Providers, "a caller no policy covers is authorized for nothing")
}
func TestAgentConfig_RealStore_UnrestrictedPolicyListsDeclaredModels(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
assert.True(t, setup.Configured)
assert.Equal(t, "https://"+testEndpoint, setup.Endpoint)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.Equal(t, "OpenAI", p.Name)
assert.Equal(t, "openai_api", p.CatalogID)
assert.Equal(t, "openai", p.APIFlavor)
assert.True(t, p.AllModelsAllowed, "policy without allowlist guardrail is unrestricted")
assert.Equal(t, []string{"gpt-5.4"}, p.Models, "declared models listed as a courtesy")
}
func TestAgentConfig_RealStore_AllowlistIntersectsDeclaredModels(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
provider.Models = []types.ProviderModel{{ID: "gpt-5.4"}, {ID: "gpt-4o"}}
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
// Allowlist admits gpt-5.4 (declared, odd casing/spacing) and gpt-4.1
// (NOT declared — the router would never route it, so it must not be
// advertised).
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", " GPT-5.4 ", "gpt-4.1")))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.False(t, p.AllModelsAllowed)
assert.Equal(t, []string{"gpt-5.4"}, p.Models, "allowlist ∩ declared, in declared order and casing")
}
func TestAgentConfig_RealStore_AllowlistMatchesBedrockDeclaredIDsCanonically(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
// A Bedrock operator typically declares the region/version form the
// vendor lists, while the allowlist holds the canonical id the proxy's
// parser emits at request time. The intersection must compare through
// the same normalization the parser applies, and the declared (raw)
// id is what gets advertised — it is what the router claims.
provider := newSynthTestProvider()
provider.ProviderID = "bedrock_api"
provider.Name = "Bedrock"
provider.Models = []types.ProviderModel{
{ID: "eu.anthropic.claude-sonnet-4-5-20250929-v1:0"},
{ID: "eu.amazon.nova-pro-v1:0"},
}
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", "anthropic.claude-sonnet-4-5")))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.False(t, p.AllModelsAllowed)
assert.Equal(t, []string{"eu.anthropic.claude-sonnet-4-5-20250929-v1:0"}, p.Models,
"the allowlisted canonical id must admit the declared region/version form, and only it")
}
func TestAgentConfig_RealStore_AllowlistHoldsRawDeclaredIDs(t *testing.T) {
// The dashboard's allowlist picker copies the provider's declared ids
// verbatim, so for path-style providers the allowlist carries the
// region/version form rather than the canonical id the parser emits.
// Both forms must admit the declared model.
cases := []struct {
name string
catalogID string
declared string
allowlist string
}{
{"bedrock", "bedrock_api", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", ""},
{"vertex", "vertex_ai_api", "claude-sonnet-4-5@20250929", ""},
// The geography/version strippers anchor on a lowercase tail, so a
// case-variant entry must be lowercased before canonicalization or
// the prefix and suffix survive into the compare key.
{"bedrock-case-variant", "bedrock_api", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
" EU.Anthropic.Claude-Sonnet-4-5-20250929-V1:0 "},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
allowlisted := tc.allowlist
if allowlisted == "" {
allowlisted = tc.declared
}
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
provider.ProviderID = tc.catalogID
provider.Name = tc.name
provider.Models = []types.ProviderModel{{ID: tc.declared}}
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", allowlisted)))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.False(t, p.AllModelsAllowed)
assert.Equal(t, []string{tc.declared}, p.Models,
"an allowlist holding the raw declared id must admit that declared model")
})
}
}
func TestAgentConfig_RealStore_UnrestrictedPolicyWinsOverRestricted(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", "gpt-5.4")))
restricted := newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, restricted))
open := newSynthTestPolicy(provider.ID, "grp-eng", "")
open.ID = "pol-2"
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, open))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
assert.True(t, setup.Providers[0].AllModelsAllowed,
"one applicable policy without an allowlist makes the provider unrestricted — the proxy would admit any model through it")
}
func TestAgentConfig_RealStore_AllowlistUnionAcrossPolicies(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
provider.Models = []types.ProviderModel{{ID: "gpt-5.4"}, {ID: "gpt-4o"}, {ID: "o4-mini"}}
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", "gpt-5.4")))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-2", "gpt-4o")))
p1 := newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, p1))
p2 := newSynthTestPolicy(provider.ID, "grp-eng", "guard-2")
p2.ID = "pol-2"
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, p2))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.False(t, p.AllModelsAllowed)
assert.ElementsMatch(t, []string{"gpt-5.4", "gpt-4o"}, p.Models, "union of allowlists across applicable policies")
}
func TestAgentConfig_RealStore_OrphanAndDisabledProvidersOmitted(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
// Orphan: enabled but referenced by no policy.
orphan := newSynthTestProvider()
orphan.ID = "prov-orphan"
require.NoError(t, s.SaveAgentNetworkProvider(ctx, orphan))
// Disabled but referenced by an applicable policy.
disabled := newSynthTestProvider()
disabled.ID = "prov-disabled"
disabled.Enabled = false
require.NoError(t, s.SaveAgentNetworkProvider(ctx, disabled))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(disabled.ID, "grp-eng", "")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
assert.True(t, setup.Configured)
assert.Empty(t, setup.Providers, "neither an orphan nor a disabled provider is reachable for the caller")
}
func TestAgentConfig_RealStore_DisabledPolicyIgnored(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
policy := newSynthTestPolicy(provider.ID, "grp-eng", "")
policy.Enabled = false
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
assert.True(t, setup.Configured)
assert.Empty(t, setup.Providers, "a disabled policy authorizes nothing")
}
func TestAgentConfig_RealStore_UndeclaredModelsUseAllowlistAsIs(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
// Gateway-style provider: no declared models — the router claims every
// model, so the allowlist union is the effective set on its own.
provider := newSynthTestProvider()
provider.ProviderID = "litellm_proxy"
provider.Name = "LiteLLM"
provider.Models = nil
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-1", "claude-sonnet-4-5")))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "guard-1")))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 1)
p := setup.Providers[0]
assert.False(t, p.AllModelsAllowed)
assert.Equal(t, []string{"claude-sonnet-4-5"}, p.Models)
}
func TestAgentConfig_RealStore_ProvidersInCreatedAtOrder(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
newer := newSynthTestProvider()
newer.ID = "prov-newer"
newer.Name = "Newer"
newer.CreatedAt = time.Date(2026, 2, 1, 0, 0, 0, 0, time.UTC)
require.NoError(t, s.SaveAgentNetworkProvider(ctx, newer))
older := newSynthTestProvider()
older.ID = "prov-older"
older.Name = "Older"
older.CreatedAt = time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)
require.NoError(t, s.SaveAgentNetworkProvider(ctx, older))
policy := newSynthTestPolicy(newer.ID, "grp-eng", "")
policy.DestinationProviderIDs = []string{newer.ID, older.ID}
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
setup, err := mgr.agentConfigForGroups(ctx, testAccountID, []string{"grp-eng"})
require.NoError(t, err)
require.Len(t, setup.Providers, 2)
assert.Equal(t, "Older", setup.Providers[0].Name)
assert.Equal(t, "Newer", setup.Providers[1].Name)
}
// TestGetAgentConfigForUser_RealStore pins the self-service entry point: the
// user's group memberships (AutoGroups — the same groups the user's peers
// carry) scope the providers, while the account's endpoint reaches every
// member — a user outside every policy gets the config with nothing
// authorized in it.
func TestGetAgentConfigForUser_RealStore(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
provider := newSynthTestProvider()
require.NoError(t, s.SaveAgentNetworkProvider(ctx, provider))
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, newSynthTestPolicy(provider.ID, "grp-eng", "")))
// users.account_id is a foreign key into accounts, enforced on
// MySQL/Postgres, so the account row must exist before its users.
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: testAccountID}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "user-in", AccountID: testAccountID, Role: nbtypes.UserRoleUser, AutoGroups: []string{"grp-eng"},
}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "user-out", AccountID: testAccountID, Role: nbtypes.UserRoleUser, AutoGroups: []string{"grp-other"},
}))
setupIn, err := mgr.GetAgentConfigForUser(ctx, testAccountID, "user-in")
require.NoError(t, err)
assert.True(t, setupIn.Configured)
require.Len(t, setupIn.Providers, 1)
setupOut, err := mgr.GetAgentConfigForUser(ctx, testAccountID, "user-out")
require.NoError(t, err)
assert.True(t, setupOut.Configured, "the account is set up, so the user reads as configured")
assert.Equal(t, "https://"+testEndpoint, setupOut.Endpoint)
assert.Empty(t, setupOut.Providers, "user outside the policy's source groups is authorized for nothing")
}
// TestGetUsageOverview_RealStore_SelfScoped pins the self-scope fallback:
// a caller without the account-wide usage grant gets the same aggregation
// the admin overview serves, but only ever their own rows — a user_id
// filter for someone else must be overridden, not honored, and never
// denied. A caller holding the grant keeps the account-wide view.
func TestGetUsageOverview_RealStore_SelfScoped(t *testing.T) {
mgr, s := newAgentConfigTestMgr(t)
mgr.permissionsManager = permissions.NewManager(s)
ctx := context.Background()
require.NoError(t, s.SaveAgentNetworkSettings(ctx, newSynthTestSettings()))
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: testAccountID}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "user-a", AccountID: testAccountID, Role: nbtypes.UserRoleUser,
}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "admin", AccountID: testAccountID, Role: nbtypes.UserRoleAdmin,
}))
own1 := newIngestTestEntry()
own1.ID, own1.UserId = "log-own-1", "user-a"
own2 := newIngestTestEntry()
own2.ID, own2.UserId = "log-own-2", "user-a"
other := newIngestTestEntry()
other.ID, other.UserId = "log-other", "user-b"
for _, e := range []*accesslogs.AccessLogEntry{own1, own2, other} {
require.NoError(t, IngestAccessLog(ctx, s, e))
}
otherID := "user-b"
filter := types.AgentNetworkAccessLogFilter{UserID: &otherID}
buckets, err := mgr.GetUsageOverview(ctx, testAccountID, "user-a", filter, types.ParseUsageGranularity(""))
require.NoError(t, err)
require.Len(t, buckets, 1, "same-day rows aggregate into one daily bucket")
assert.Equal(t, int64(200), buckets[0].InputTokens, "only the caller's two rows count — the foreign user_id filter is overridden")
assert.Equal(t, int64(100), buckets[0].OutputTokens)
adminBuckets, err := mgr.GetUsageOverview(ctx, testAccountID, "admin", types.AgentNetworkAccessLogFilter{}, types.ParseUsageGranularity(""))
require.NoError(t, err)
require.Len(t, adminBuckets, 1)
assert.Equal(t, int64(300), adminBuckets[0].InputTokens, "the account-wide grant keeps the unscoped view")
}
@@ -81,6 +81,10 @@ type Provider struct {
// surface — the proxy middleware then falls back to URL sniffing
// or skips request-side enrichment.
ParserID string
// RouterVendors declares every parser surface a gateway route can serve.
// Leave empty for single-surface providers, where ParserID remains the
// router discriminator for backward compatibility.
RouterVendors []string
// PricingSurfaces names the cost-meter pricing surfaces this
// provider's Models are priced under ("openai", "anthropic",
// "bedrock" — the llm.Parser surface the request parser stamps as
@@ -113,8 +117,63 @@ type Provider struct {
// upstream provider + credentials on Portkey's hosted side).
ExtraHeaders []ExtraHeader
Models []Model
// Discovery, when non-nil, describes how to ask this vendor which
// models the operator's own credential can actually reach, so the
// provider form can offer a live list instead of only the hand-curated
// Models above. Nil entries keep free-text entry.
Discovery *Discovery
}
// ListingShape names the response envelope a vendor returns its model
// listing in. Every vendor invented its own, and none of them can be
// guessed from the request, so the catalog states it.
type ListingShape string
const (
// ShapeOpenAIData is {"data":[{"id":…}]} — OpenAI, and Anthropic, which
// adopted the same envelope.
ShapeOpenAIData ListingShape = "openai_data"
// ShapeBedrockInferenceProfiles is
// {"inferenceProfileSummaries":[{"inferenceProfileId":…}]}. The ids carry
// the region prefix that makes them invocable, which is exactly what an
// operator cannot reconstruct by hand.
ShapeBedrockInferenceProfiles ListingShape = "bedrock_inference_profiles"
// ShapeVertexPublisherModels is {"publisherModels":[{"name":…}]}, where
// name is a resource path and the invocable id is its last segment joined
// to a separate versionId field.
ShapeVertexPublisherModels ListingShape = "vertex_publisher_models"
)
// Discovery describes one vendor's model-listing endpoint.
//
// Host is deliberately separate from the provider record's upstream URL:
// Bedrock serves listings from the control plane (bedrock.<region>) while
// inference must go to the runtime host (bedrock-runtime.<region>), so the
// two cannot be the same value. Empty Host means "use the record's own
// upstream", which is right for every vendor that serves both from one host.
//
// The regionPlaceholder in Host is substituted from the provider record's
// region. Deriving the discovery host from the catalog rather than accepting
// one from the caller is also what keeps this from being an open proxy: the
// only hosts management will dial are the ones written here.
type Discovery struct {
Host string
Path string
Query string
Shape ListingShape
// ExactModelsOnly omits wildcard patterns from listings when NetBird's
// provider model rows cannot represent the vendor's matching semantics.
ExactModelsOnly bool
// Headers are static headers the vendor requires beyond the credential
// (Anthropic versions its API through one and rejects a request without
// it). The auth header itself comes from AuthHeaderName/Template.
Headers map[string]string
}
// RegionPlaceholder is replaced in Discovery.Host by the provider record's
// configured region.
const RegionPlaceholder = "<region>"
// ExtraHeader names a single optional per-provider routing/config
// header. Catalog declares N of these per provider type; the operator
// fills any subset on the provider record (see Provider.ExtraValues).
@@ -245,8 +304,12 @@ var providers = []Provider{
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#10A37F",
ParserID: "openai",
PricingSurfaces: []string{"openai"},
Discovery: &Discovery{
Path: "/v1/models",
Shape: ShapeOpenAIData,
},
ParserID: "openai",
PricingSurfaces: []string{"openai"},
// Pricing + context windows cross-checked against LiteLLM's
// model_prices_and_context_window.json. Notable corrections from
// earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40
@@ -284,8 +347,18 @@ var providers = []Provider{
AuthHeaderTemplate: "${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#D97757",
ParserID: "anthropic",
PricingSurfaces: []string{"anthropic"},
Discovery: &Discovery{
Path: "/v1/models",
// The default page is short and a picker wants the whole
// catalogue in one call.
Query: "limit=1000",
Shape: ShapeOpenAIData,
// Anthropic versions its API through a header and refuses a
// request that omits it, listing included.
Headers: map[string]string{"anthropic-version": "2023-06-01"},
},
ParserID: "anthropic",
PricingSurfaces: []string{"anthropic"},
// Per Anthropic's current model lineup. Pricing in USD per 1k
// tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at
// 200K. claude-3-7-sonnet and claude-3-5-haiku retired
@@ -296,6 +369,8 @@ var providers = []Provider{
// account to be on >= 30-day data retention or all requests
// 400.
Models: []Model{
{ID: "claude-opus-5", Label: "Claude Opus 5", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "claude-sonnet-5", Label: "Claude Sonnet 5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
@@ -343,6 +418,22 @@ var providers = []Provider{
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#FF9900",
// Listings come from the CONTROL PLANE, not the runtime host in
// DefaultHost above: ListInferenceProfiles is not an operation
// bedrock-runtime implements, and answers <UnknownOperationException/>
// there. Inference has to go to the runtime host, so the two hosts
// genuinely differ and Discovery.Host carries the difference.
//
// Inference profiles rather than foundation models because the profile
// id is the invocable one: it carries the region prefix (eu., us.,
// global.) that AWS requires and that cannot be derived from the
// configured region — an eu-central-1 account legitimately holds
// global.* profiles.
Discovery: &Discovery{
Host: "bedrock." + RegionPlaceholder + ".amazonaws.com",
Path: "/inference-profiles",
Shape: ShapeBedrockInferenceProfiles,
},
// ParserID stays empty (path-style dispatch via IsBedrockPathStyle);
// the request parser meters these under the "bedrock" surface.
PricingSurfaces: []string{"bedrock"},
@@ -355,6 +446,8 @@ var providers = []Provider{
// Llama 3.3 70B entry kept unchanged — LiteLLM tracks only
// per-region Llama 3 entries; standalone 3.3 not yet listed.
Models: []Model{
{ID: "anthropic.claude-opus-5", Label: "Claude Opus 5 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "anthropic.claude-sonnet-5", Label: "Claude Sonnet 5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
@@ -391,6 +484,15 @@ var providers = []Provider{
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#4285F4",
// Only the v1beta1 publisher listing answers: the v1 form and the
// project-scoped form under BOTH versions return 404. That means the
// list is publisher-global — it cannot say which models this project
// has enabled — so it is offered as a suggestion beside the catalog
// rather than replacing it. See the discovery e2e for the probes.
Discovery: &Discovery{
Path: "/v1beta1/publishers/anthropic/models",
Shape: ShapeVertexPublisherModels,
},
// ParserID stays empty (path-style dispatch via IsVertexPathStyle);
// Anthropic-on-Vertex requests are metered under the "anthropic"
// surface with the bare, unversioned model id.
@@ -406,6 +508,8 @@ var providers = []Provider{
// exists — the router denies unmeterable publishers rather than forward
// them uncounted.
Models: []Model{
{ID: "claude-opus-5", Label: "Claude Opus 5 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "claude-sonnet-5", Label: "Claude Sonnet 5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
@@ -537,6 +641,34 @@ var providers = []Provider{
},
Models: []Model{},
},
{
ID: "agentgateway",
Kind: KindGateway,
Name: "agentgateway",
Description: "Bring your own agentgateway with trusted NetBird identity stamped on every request",
DefaultHost: "",
AuthHeaderName: "Authorization",
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#8023C3",
// Agentgateway accepts both OpenAI and Anthropic request shapes.
// Leave ParserID empty so the proxy detects the shape from the URL.
ParserID: "",
RouterVendors: []string{"openai", "anthropic"},
PricingSurfaces: []string{"openai", "anthropic"},
Discovery: &Discovery{
Path: "/v1/models",
Shape: ShapeOpenAIData,
ExactModelsOnly: true,
},
IdentityInjection: &IdentityInjection{
HeaderPair: &HeaderPairInjection{
EndUserIDHeader: "x-netbird-user-id",
TagsHeader: "x-netbird-groups",
},
},
Models: []Model{},
},
{
ID: "portkey",
Kind: KindGateway,
@@ -0,0 +1,86 @@
package catalog
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// TestClaudeLineupSelectable pins the models Claude Code resolves to by
// default. A model absent from the lineup can't be ticked on a provider
// record, so llm_router denies it as not-routable and the operator has no
// way to authorise the client's own default.
func TestClaudeLineupSelectable(t *testing.T) {
for providerID, wanted := range map[string][]string{
"anthropic_api": {"claude-opus-5", "claude-sonnet-5", "claude-haiku-4-5"},
"bedrock_api": {"anthropic.claude-opus-5", "anthropic.claude-sonnet-5", "anthropic.claude-haiku-4-5"},
"vertex_ai_api": {"claude-opus-5", "claude-sonnet-5", "claude-haiku-4-5"},
} {
provider, ok := Lookup(providerID)
require.True(t, ok, "catalog must define %s", providerID)
selectable := make(map[string]Model, len(provider.Models))
for _, m := range provider.Models {
selectable[m.ID] = m
}
for _, id := range wanted {
model, found := selectable[id]
require.True(t, found, "%s must offer %s", providerID, id)
assert.NotEmpty(t, model.Label, "%s/%s needs a label for the picker", providerID, id)
assert.Positive(t, model.InputPer1k, "%s/%s needs an input rate", providerID, id)
assert.Positive(t, model.OutputPer1k, "%s/%s needs an output rate", providerID, id)
assert.Positive(t, model.ContextWindow, "%s/%s needs a context window", providerID, id)
}
}
}
func TestAgentgatewayCatalogEntry(t *testing.T) {
entry, ok := Lookup("agentgateway")
require.True(t, ok, "agentgateway must be available in the provider catalog")
assert.Equal(t, KindGateway, entry.Kind, "agentgateway must be grouped with AI gateways")
assert.Empty(t, entry.DefaultHost, "operators must provide their agentgateway proxy URL")
assert.Equal(t, "Authorization", entry.AuthHeaderName)
assert.Equal(t, "Bearer ${API_KEY}", entry.AuthHeaderTemplate)
assert.Equal(t, "application/json", entry.DefaultContentType)
assert.Empty(t, entry.ParserID, "URL detection must select the OpenAI or Anthropic parser")
assert.Equal(t, []string{"openai", "anthropic"}, entry.RouterVendors,
"agentgateway must accept both parser surfaces")
assert.Equal(t, []string{"openai", "anthropic"}, entry.PricingSurfaces,
"agentgateway models can use either pricing surface")
assert.Empty(t, entry.Models, "an empty model list makes agentgateway a catch-all route")
require.NotNil(t, entry.Discovery)
assert.Empty(t, entry.Discovery.Host, "discovery must use the configured proxy URL")
assert.Equal(t, "/v1/models", entry.Discovery.Path)
assert.Equal(t, ShapeOpenAIData, entry.Discovery.Shape)
assert.True(t, entry.Discovery.ExactModelsOnly,
"wildcard model semantics are not supported by NetBird")
require.NotNil(t, entry.IdentityInjection)
require.NotNil(t, entry.IdentityInjection.HeaderPair)
assert.Nil(t, entry.IdentityInjection.JSONMetadata)
assert.False(t, entry.IdentityInjection.HeaderPair.Customizable,
"NetBird identity header names are part of the integration contract")
assert.Equal(t, "x-netbird-user-id", entry.IdentityInjection.HeaderPair.EndUserIDHeader)
assert.Equal(t, "x-netbird-groups", entry.IdentityInjection.HeaderPair.TagsHeader)
assert.False(t, entry.IdentityInjection.HeaderPair.EndUserIDInBody)
assert.False(t, entry.IdentityInjection.HeaderPair.TagsInBody)
}
func TestAgentgatewayCatalogAPIResponse(t *testing.T) {
entry, ok := Lookup("agentgateway")
require.True(t, ok)
resp := entry.ToAPIResponse()
assert.Equal(t, "agentgateway", resp.Id)
assert.Equal(t, api.AgentNetworkCatalogProviderKindGateway, resp.Kind)
assert.Empty(t, resp.Models)
require.NotNil(t, resp.IdentityInjection)
require.NotNil(t, resp.IdentityInjection.HeaderPair)
assert.False(t, resp.IdentityInjection.HeaderPair.Customizable)
assert.Equal(t, "x-netbird-user-id", resp.IdentityInjection.HeaderPair.EndUserIdHeader)
assert.Equal(t, "x-netbird-groups", resp.IdentityInjection.HeaderPair.TagsHeader)
}
@@ -0,0 +1,135 @@
package agentnetwork
import (
"context"
"errors"
"net/http"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/shared/management/status"
)
// ModelLister is the vendor-facing half of the credential check.
// modeldiscovery.Client is the only production implementation; it is an
// interface because the check runs on a write path, so without a seam every
// test that saves a provider would reach a vendor to do it.
type ModelLister interface {
Fetch(ctx context.Context, req modeldiscovery.Request) ([]modeldiscovery.Model, error)
}
// checkProviderCredential refuses a record whose upstream or credential the
// vendor will not accept.
//
// It reuses the discovery Fetch rather than a lighter status probe so it
// exercises the path the model picker takes: a URL answering 200 with a login
// page fails here instead of producing an empty picker later.
func (m *managerImpl) checkProviderCredential(ctx context.Context, provider *types.Provider) error {
// A record that asks the proxy to skip certificate verification is one this
// check cannot speak for. Discovery verifies certificates, so a self-hosted
// endpoint behind a self-signed one would be refused for a reason the
// operator already told us to ignore — a lockout of exactly the setup the
// flag exists for. Sending the credential over a connection management
// declines to verify is the other way out, and a worse one.
if provider.SkipTLSVerification {
log.WithContext(ctx).Debugf("agent network provider %s not credential-checked: tls verification is disabled for it", provider.ProviderID)
return nil
}
_, err := m.modelDiscovery.Fetch(ctx, modeldiscovery.Request{
CatalogID: provider.ProviderID,
UpstreamURL: provider.UpstreamURL,
APIKey: provider.APIKey,
})
if err == nil {
return nil
}
message, blocking := credentialCheckFailure(err)
if !blocking {
log.WithContext(ctx).Debugf("agent network provider %s not credential-checked: %v", provider.ProviderID, err)
return nil
}
// WriteError logs only what we return, and that carries no status code,
// so the vendor's number is recorded here or nowhere.
log.WithContext(ctx).Infof("agent network provider %s failed its credential check: %v", provider.ProviderID, err)
return status.Errorf(status.InvalidArgument, "%s", message)
}
// discoveryFailure renders a failed model listing for the operator who pressed
// the button. Every outcome here is something they did or configured — a key
// the vendor refused, an upstream that does not answer — so it owes them the
// same sentence a refused save gives, not the generic 500 an unclassified
// error turns into.
//
// ErrNoDiscovery and ErrInvalidRequest pass through untouched: the handler
// already maps them, and "this provider has no listing endpoint" is a fact
// about the catalog rather than a failure to report as one.
func discoveryFailure(ctx context.Context, catalogID string, err error) error {
if errors.Is(err, modeldiscovery.ErrNoDiscovery) || errors.Is(err, modeldiscovery.ErrInvalidRequest) {
return err
}
message, _ := credentialCheckFailure(err)
if message == "" {
return err
}
// The operator's message carries no status code, so the vendor's number is
// recorded here or nowhere.
log.WithContext(ctx).Infof("agent network model discovery for %s failed: %v", catalogID, err)
return status.Errorf(status.InvalidArgument, "%s", message)
}
// credentialCheckFailure renders a discovery failure as the sentence the
// provider form shows, and reports whether it should block the write.
//
// The strings survive WriteError lowercasing them, and never echo the
// operator's URL: paths are case-sensitive, so an echoed URL comes back
// altered and describes something they did not type.
func credentialCheckFailure(err error) (message string, blocking bool) {
// Not checkable. The record may be perfectly good and we have no way to
// ask, so reporting a failure would be a guess.
switch {
case errors.Is(err, modeldiscovery.ErrNoDiscovery),
errors.Is(err, modeldiscovery.ErrNoDiscoveryHost),
errors.Is(err, modeldiscovery.ErrPrivateHost):
return "", false
}
var vendor *modeldiscovery.VendorStatusError
if errors.As(err, &vendor) {
switch vendor.Status {
case http.StatusUnauthorized, http.StatusForbidden:
return "the provider rejected the credential", true
case http.StatusNotFound, http.StatusMethodNotAllowed:
return "the upstream url did not answer a model listing", true
default:
// 5xx and 429 included: an outage still leaves the record
// unverified, which is what this refuses to save.
return "the provider returned an error", true
}
}
var unreachable *modeldiscovery.UnreachableError
if errors.As(err, &unreachable) {
if reason := unreachable.Reason(); reason != "" {
return "the upstream url could not be reached: " + reason, true
}
return "the upstream url could not be reached", true
}
if errors.Is(err, modeldiscovery.ErrUnparseableListing) {
return "the upstream url answered, but not with a model listing", true
}
// Ours rather than the vendor's — a request this code built badly, or a
// catalog entry that does not match its parser. Still unverified, so it
// still blocks.
return "the provider could not be checked", true
}
@@ -0,0 +1,605 @@
package agentnetwork
import (
"context"
"crypto/tls"
"errors"
"fmt"
"net"
"os"
"syscall"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/status"
)
// stubLister stands in for the vendor on the write path. It records what it
// was asked so a test can assert not only that the check ran, but that it ran
// against the right upstream and the right credential — and, for an edit that
// touches neither, that it did not run at all.
type stubLister struct {
err error
requests []modeldiscovery.Request
}
func (s *stubLister) Fetch(_ context.Context, req modeldiscovery.Request) ([]modeldiscovery.Model, error) {
s.requests = append(s.requests, req)
if s.err != nil {
return nil, s.err
}
return []modeldiscovery.Model{{ID: "a-model", PricingKnown: true}}, nil
}
func (s *stubLister) calls() int { return len(s.requests) }
func (s *stubLister) only(t *testing.T) modeldiscovery.Request {
t.Helper()
require.Len(t, s.requests, 1, "the vendor must be asked exactly once")
return s.requests[0]
}
// TestCredentialCheckFailure_SeparatesTheUrlFromTheCredential is the contract
// the provider form is written against: an operator gets told which of the two
// fields they have to look at, and the message says so without a status code
// and without echoing the URL back at them.
func TestCredentialCheckFailure_SeparatesTheUrlFromTheCredential(t *testing.T) {
cases := []struct {
name string
err error
want string
}{
{
name: "401 is the credential",
err: &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 401},
want: "the provider rejected the credential",
},
{
name: "403 is the credential",
err: &modeldiscovery.VendorStatusError{Provider: "Bedrock", Status: 403},
want: "the provider rejected the credential",
},
{
// The host authenticated us fine and then said it has no such
// endpoint, which is the URL being wrong rather than the key.
name: "404 is the url",
err: &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 404},
want: "the upstream url did not answer a model listing",
},
{
name: "405 is the url",
err: &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 405},
want: "the upstream url did not answer a model listing",
},
{
name: "500 is the vendor",
err: &modeldiscovery.VendorStatusError{Provider: "Anthropic", Status: 500},
want: "the provider returned an error",
},
{
name: "503 is the vendor",
err: &modeldiscovery.VendorStatusError{Provider: "Anthropic", Status: 503},
want: "the provider returned an error",
},
{
name: "429 is the vendor",
err: &modeldiscovery.VendorStatusError{Provider: "Anthropic", Status: 429},
want: "the provider returned an error",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
got, blocking := credentialCheckFailure(tc.err)
require.True(t, blocking, "a vendor refusal must block the write")
require.Equal(t, tc.want, got)
})
}
}
// TestCredentialCheckFailure_NamesTheTransportFault covers the failures that
// never reached the vendor. The distinction inside them is worth keeping: a
// refused connection is a wrong port and an unknown host is a wrong hostname,
// and an operator staring at a URL they believe in needs to be told which.
func TestCredentialCheckFailure_NamesTheTransportFault(t *testing.T) {
cases := []struct {
name string
err error
want string
}{
{
name: "unknown host",
err: &net.DNSError{Err: "no such host", Name: "api.example.com", IsNotFound: true},
want: "the upstream url could not be reached: no such host",
},
{
name: "dns failure that is not a missing name",
err: &net.DNSError{Err: "server misbehaving", Name: "api.example.com"},
want: "the upstream url could not be reached: dns lookup failed",
},
{
name: "connection refused",
err: &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED},
want: "the upstream url could not be reached: connection refused",
},
{
name: "host unreachable",
err: &net.OpError{Op: "dial", Net: "tcp", Err: syscall.EHOSTUNREACH},
want: "the upstream url could not be reached: host unreachable",
},
{
name: "timeout",
err: fmt.Errorf("dial: %w", os.ErrDeadlineExceeded),
want: "the upstream url could not be reached: connection timed out",
},
{
name: "context deadline",
err: fmt.Errorf("dial: %w", context.DeadlineExceeded),
want: "the upstream url could not be reached: connection timed out",
},
{
name: "untrusted certificate",
err: &tls.CertificateVerificationError{},
want: "the upstream url could not be reached: tls certificate not trusted",
},
{
name: "plaintext service on an https url",
err: tls.RecordHeaderError{Msg: "first record does not look like a TLS handshake"},
want: "the upstream url could not be reached: not a tls endpoint",
},
{
// Nothing we recognise. Better to say only that it could not be
// reached than to paste a Go error into the provider form.
name: "cause we do not recognise",
err: errors.New("something went sideways"),
want: "the upstream url could not be reached",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
wrapped := &modeldiscovery.UnreachableError{Provider: "OpenAI", Err: tc.err}
got, blocking := credentialCheckFailure(wrapped)
require.True(t, blocking, "an unreachable upstream must block the write")
require.Equal(t, tc.want, got)
})
}
}
// TestCredentialCheckFailure_AnAnsweringUrlThatIsNotTheApi covers the case a
// status probe would wave through: the host is up, the credential was accepted
// or not required, and the body is a login page. Reusing the discovery parser
// for the check is what catches it.
func TestCredentialCheckFailure_AnAnsweringUrlThatIsNotTheApi(t *testing.T) {
err := fmt.Errorf("%w: decode model listing: unexpected token", modeldiscovery.ErrUnparseableListing)
got, blocking := credentialCheckFailure(err)
require.True(t, blocking)
require.Equal(t, "the upstream url answered, but not with a model listing", got)
}
// TestCredentialCheckFailure_WhatCannotBeCheckedIsNotAFailure pins the
// difference between "this record is wrong" and "we have no way to ask". A
// gateway with no listing endpoint, a Bedrock record pointed at a proxy, and a
// self-hosted endpoint the proxy reaches through the tunnel are all legitimate
// providers. Blocking them would make the feature a lockout.
func TestCredentialCheckFailure_WhatCannotBeCheckedIsNotAFailure(t *testing.T) {
cases := map[string]error{
"no listing endpoint": modeldiscovery.ErrNoDiscovery,
"no derivable host": fmt.Errorf("%w: %w: bedrock", modeldiscovery.ErrInvalidRequest, modeldiscovery.ErrNoDiscoveryHost),
"private upstream": fmt.Errorf("%w: %w: 10.0.0.5", modeldiscovery.ErrInvalidRequest, modeldiscovery.ErrPrivateHost),
}
for name, err := range cases {
t.Run(name, func(t *testing.T) {
message, blocking := credentialCheckFailure(err)
require.False(t, blocking, "a provider we cannot check must still save")
require.Empty(t, message)
})
}
}
// TestCredentialCheckFailure_AnUnrecognisedFailureStillBlocks covers a fault of
// ours rather than the vendor's — a malformed request this code built, or a
// catalog entry whose parser does not match its endpoint. The record went
// unverified either way, and silently saving what we could not check is the
// thing this feature exists to prevent.
func TestCredentialCheckFailure_AnUnrecognisedFailureStillBlocks(t *testing.T) {
message, blocking := credentialCheckFailure(errors.New("no parser for listing shape \"\""))
require.True(t, blocking)
require.Equal(t, "the provider could not be checked", message)
}
// newCheckedProvider returns a record shaped the way the handler guarantees
// one: a known catalog id, a public upstream and a key.
func newCheckedProvider(accountID string) *types.Provider {
provider := types.NewProvider(accountID)
provider.ProviderID = "openai_api"
provider.Name = "openai"
provider.UpstreamURL = "https://api.openai.com"
provider.APIKey = "sk-good"
provider.Enabled = true
return provider
}
// TestCreateProvider_RefusesARecordTheVendorRejects is the whole point of the
// feature: a key with a character missing used to save cleanly and surface
// minutes later as a failed request with nothing pointing back at the record.
func TestCreateProvider_RefusesARecordTheVendorRejects(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.vendor.err = &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 401}
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
_, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.Error(t, err)
require.Contains(t, err.Error(), "the provider rejected the credential")
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
require.Equal(t, status.InvalidArgument, sErr.Type(), "the refusal must reach the caller as a 422")
stored, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1")
require.NoError(t, err)
require.Empty(t, stored, "a record that failed its check must not be written")
}
// TestCreateProvider_ChecksTheCredentialItWasGiven pins what the vendor is
// asked with, since a check run against the wrong upstream or a stale key
// would pass while proving nothing.
func TestCreateProvider_ChecksTheCredentialItWasGiven(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
_, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
asked := f.vendor.only(t)
require.Equal(t, "openai_api", asked.CatalogID)
require.Equal(t, "https://api.openai.com", asked.UpstreamURL)
require.Equal(t, "sk-good", asked.APIKey)
}
// TestUpdateProvider_AUrlOnlyChangeIsCheckedAgainstTheStoredKey covers the
// case that shaped where the check sits. The key never returns to the browser,
// so an operator editing only the URL has none to offer — the stored one is
// the only credential there is, and the new URL still has to be proven with
// it.
func TestUpdateProvider_AUrlOnlyChangeIsCheckedAgainstTheStoredKey(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
f.vendor.requests = nil
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Update, true)
edit := newCheckedProvider("account1")
edit.ID = created.ID
edit.UpstreamURL = "https://gateway.example.com"
edit.APIKey = "" // the form sends no key when it was not retyped
_, err = f.manager.UpdateProvider(ctx, "user1", edit)
require.NoError(t, err)
asked := f.vendor.only(t)
require.Equal(t, "https://gateway.example.com", asked.UpstreamURL, "the new url must be what gets tested")
require.Equal(t, "sk-good", asked.APIKey, "and the stored key must be what tests it")
}
// TestUpdateProvider_AFailedRotationLeavesTheWorkingKeyInPlace is the
// half-applied state the check must never produce: refusing the new key while
// having already replaced the old one would take the provider down.
func TestUpdateProvider_AFailedRotationLeavesTheWorkingKeyInPlace(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
f.vendor.err = &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 403}
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Update, true)
rotation := newCheckedProvider("account1")
rotation.ID = created.ID
rotation.APIKey = "sk-typo"
_, err = f.manager.UpdateProvider(ctx, "user1", rotation)
require.Error(t, err)
require.Contains(t, err.Error(), "the provider rejected the credential")
stored, err := f.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, "account1", created.ID)
require.NoError(t, err)
require.Equal(t, "sk-good", stored.APIKey, "the rejected key must not have replaced the working one")
}
// TestUpdateProvider_AnEditTouchingNeitherFieldAsksNoVendor keeps renames,
// model rows and price edits off the vendor's doorstep. They have nothing new
// to prove, and making them wait on a vendor — or fail because one is having a
// bad day — would be a tax on edits that carry no risk.
func TestUpdateProvider_AnEditTouchingNeitherFieldAsksNoVendor(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
f.vendor.requests = nil
// Any call at all now would fail the update, which is what makes the
// assertion below load-bearing rather than decorative.
f.vendor.err = &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 500}
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Update, true)
rename := newCheckedProvider("account1")
rename.ID = created.ID
rename.Name = "openai-renamed"
rename.APIKey = ""
_, err = f.manager.UpdateProvider(ctx, "user1", rename)
require.NoError(t, err, "an edit that changes neither url nor key must not be checked")
require.Zero(t, f.vendor.calls(), "and must not reach the vendor at all")
}
// TestCreateProvider_AProviderWeCannotCheckStillSaves covers the eleven
// catalog entries with no listing endpoint, a Bedrock record behind a proxy,
// and a self-hosted endpoint on a private network. None of those are evidence
// the record is wrong, and refusing them would make this a lockout.
func TestCreateProvider_AProviderWeCannotCheckStillSaves(t *testing.T) {
cases := map[string]error{
"gateway with no listing endpoint": modeldiscovery.ErrNoDiscovery,
"bedrock behind a proxy": fmt.Errorf("%w: %w", modeldiscovery.ErrInvalidRequest, modeldiscovery.ErrNoDiscoveryHost),
"self-hosted on a private network": fmt.Errorf("%w: %w", modeldiscovery.ErrInvalidRequest, modeldiscovery.ErrPrivateHost),
}
for name, vendorErr := range cases {
t.Run(name, func(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.vendor.err = vendorErr
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
require.NotNil(t, created)
stored, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1")
require.NoError(t, err)
require.Len(t, stored, 1, "a provider we cannot check must still be written")
})
}
}
// TestDiscoveryFailure_TellsTheOperatorWhatWentWrong covers the button, not the
// save. Pressing "Load models from provider" against a bad key used to answer
// "internal server error", which names neither the thing that failed nor
// anything the operator could act on — every outcome here is their key or their
// URL.
func TestDiscoveryFailure_TellsTheOperatorWhatWentWrong(t *testing.T) {
cases := map[string]struct {
err error
want string
}{
"refused credential": {
err: &modeldiscovery.VendorStatusError{Provider: "Bedrock", Status: 403},
want: "the provider rejected the credential",
},
"upstream that is not the api": {
err: &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 404},
want: "the upstream url did not answer a model listing",
},
"upstream that does not resolve": {
err: &modeldiscovery.UnreachableError{
Provider: "OpenAI",
Err: &net.DNSError{Err: "no such host", Name: "api.example.com", IsNotFound: true},
},
want: "the upstream url could not be reached: no such host",
},
"vendor having a bad day": {
err: &modeldiscovery.VendorStatusError{Provider: "Anthropic", Status: 503},
want: "the provider returned an error",
},
}
for name, tc := range cases {
t.Run(name, func(t *testing.T) {
err := discoveryFailure(context.Background(), "openai_api", tc.err)
require.EqualError(t, err, tc.want)
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
require.Equal(t, status.InvalidArgument, sErr.Type(),
"a failure the operator caused must not read as a server fault")
})
}
}
// TestDiscoveryFailure_LeavesTheCatalogFactsAlone keeps the two outcomes the
// handler already maps. A provider with no listing endpoint is a fact about the
// catalog entry, and the caller falls back to the catalog's own models rather
// than showing an error at all — rewriting it as a refusal would turn a normal
// path into one.
func TestDiscoveryFailure_LeavesTheCatalogFactsAlone(t *testing.T) {
for name, err := range map[string]error{
"no listing endpoint": modeldiscovery.ErrNoDiscovery,
"bad request": fmt.Errorf("%w: unknown catalog provider", modeldiscovery.ErrInvalidRequest),
} {
t.Run(name, func(t *testing.T) {
require.Equal(t, err, discoveryFailure(context.Background(), "openai_api", err),
"the handler's own mapping must still see the original error")
})
}
}
// TestDiscoverProviderModels_SurfacesTheVendorRefusal drives the manager rather
// than the classifier, so a future refactor that stops translating on this path
// fails here rather than silently going back to 500s.
func TestDiscoverProviderModels_SurfacesTheVendorRefusal(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.vendor.err = &modeldiscovery.VendorStatusError{Provider: "OpenAI", Status: 401}
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
_, err := f.manager.DiscoverProviderModels(ctx, "account1", "user1", modeldiscovery.Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
APIKey: "sk-wrong",
}, "")
require.EqualError(t, err, "the provider rejected the credential")
}
// TestDiscoverProviderModels_ListsAgainstTheUrlOnTheForm covers the edit the
// operator cannot otherwise make: the upstream has been retyped and the
// credential has not, because the API never returned it to be retyped. Naming
// the record supplies the key; the request supplies the URL under test.
func TestDiscoverProviderModels_ListsAgainstTheUrlOnTheForm(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
// Twice: the create, and the listing, which is gated on Create too.
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
f.vendor.requests = nil
_, err = f.manager.DiscoverProviderModels(ctx, "account1", "user1", modeldiscovery.Request{
CatalogID: "openai_api",
UpstreamURL: "https://gateway.example.com",
}, created.ID)
require.NoError(t, err)
asked := f.vendor.only(t)
require.Equal(t, "https://gateway.example.com", asked.UpstreamURL, "the typed url must be the one listed against")
require.Equal(t, "sk-good", asked.APIKey, "and the stored key must be what lists it")
}
// TestDiscoverProviderModels_FallsBackToTheStoredUrl keeps the plain refresh
// working: a request naming only the record still reaches the saved upstream.
func TestDiscoverProviderModels_FallsBackToTheStoredUrl(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
stored := f.vendor.only(t).UpstreamURL
f.vendor.requests = nil
_, err = f.manager.DiscoverProviderModels(ctx, "account1", "user1", modeldiscovery.Request{
CatalogID: "openai_api",
}, created.ID)
require.NoError(t, err)
require.Equal(t, stored, f.vendor.only(t).UpstreamURL)
}
// TestUpdateProvider_MovingARecordToAnotherVendorIsChecked covers the edit that
// changes neither field the vendor judges and still invalidates both. The
// catalog entry decides which vendor is asked and under which auth header, so
// the unchanged credential is now being offered somewhere it has never been
// accepted.
func TestUpdateProvider_MovingARecordToAnotherVendorIsChecked(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
created, err := f.manager.CreateProvider(ctx, "user1", newCheckedProvider("account1"))
require.NoError(t, err)
f.vendor.requests = nil
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Update, true)
edit := newCheckedProvider("account1")
edit.ID = created.ID
edit.ProviderID = "anthropic_api"
edit.APIKey = ""
_, err = f.manager.UpdateProvider(ctx, "user1", edit)
require.NoError(t, err)
require.Equal(t, "anthropic_api", f.vendor.only(t).CatalogID,
"the new vendor is the one that has to accept the key")
}
// TestCreateProvider_ASkipTlsRecordIsNotCheckedAgainstItsCertificate covers the
// lockout the check would otherwise be: the flag exists for a self-hosted
// endpoint behind a certificate nothing public can verify, and discovery
// verifies certificates. Refusing the save would reject the record for the one
// reason the operator already declared they accept.
func TestCreateProvider_ASkipTlsRecordIsNotCheckedAgainstItsCertificate(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.vendor.err = &modeldiscovery.UnreachableError{
Provider: "OpenAI",
Err: &tls.CertificateVerificationError{},
}
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
provider := newCheckedProvider("account1")
provider.SkipTLSVerification = true
created, err := f.manager.CreateProvider(ctx, "user1", provider)
require.NoError(t, err, "a record we were told not to verify must still save")
require.NotEmpty(t, created.ID)
require.Zero(t, f.vendor.calls(), "and the vendor must not be asked at all")
}
// TestCreateProvider_TheStoredKeyIsTheOneThatWasChecked pins the two halves to
// one value. The vendor call trims the credential before building its auth
// header; the synthesiser substitutes the stored one verbatim. A key pasted
// with surrounding whitespace would otherwise pass its check and then fail
// every request the provider serves.
func TestCreateProvider_TheStoredKeyIsTheOneThatWasChecked(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
provider := newCheckedProvider("account1")
provider.APIKey = " sk-good\n"
created, err := f.manager.CreateProvider(ctx, "user1", provider)
require.NoError(t, err)
require.Equal(t, "sk-good", f.vendor.only(t).APIKey, "the vendor is asked about the trimmed key")
stored, err := f.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, "account1", created.ID)
require.NoError(t, err)
require.Equal(t, "sk-good", stored.APIKey, "and that is the one the proxy will send")
}
// TestUpdateProvider_TurningTlsVerificationBackOnChecksTheRecord covers the
// hole the skip-TLS exemption opens on its own. Such a record is stored without
// ever being checked, so the moment verification is switched back on is the
// first moment it can be checked at all — and none of the three fields the
// re-check usually watches has to move for that to happen.
func TestUpdateProvider_TurningTlsVerificationBackOnChecksTheRecord(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
unchecked := newCheckedProvider("account1")
unchecked.SkipTLSVerification = true
created, err := f.manager.CreateProvider(ctx, "user1", unchecked)
require.NoError(t, err)
require.Zero(t, f.vendor.calls(), "the create was exempt")
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Update, true)
edit := newCheckedProvider("account1")
edit.ID = created.ID
edit.APIKey = ""
edit.SkipTLSVerification = false
_, err = f.manager.UpdateProvider(ctx, "user1", edit)
require.NoError(t, err)
require.Equal(t, 1, f.vendor.calls(), "switching verification on must check what was never checked")
}
@@ -0,0 +1,56 @@
package handlers
import (
"net/http"
"github.com/gorilla/mux"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/http/util"
)
// addAgentConfigEndpoints registers the self-service agent-config route.
// It is available to every authenticated user regardless of role: the
// providers in the response are scoped strictly to the caller, which is
// tighter than any role gate could be. The caller's own usage and requests are served by
// the regular usage/logs endpoints, which self-scope for callers without
// the account-wide grants.
func (h *handler) addAgentConfigEndpoints(router *mux.Router) {
router.HandleFunc("/agent-network/agent-config", h.getAgentConfig).Methods("GET", "OPTIONS")
}
func (h *handler) getAgentConfig(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
util.WriteError(r.Context(), err, w)
return
}
setup, err := h.manager.GetAgentConfigForUser(r.Context(), userAuth.AccountId, userAuth.UserId)
if err != nil {
util.WriteError(r.Context(), err, w)
return
}
util.WriteJSONObject(r.Context(), w, agentConfigToAPI(setup))
}
func agentConfigToAPI(setup *types.AgentConfig) api.AgentNetworkAgentConfig {
providers := make([]api.AgentNetworkAgentConfigProvider, 0, len(setup.Providers))
for _, p := range setup.Providers {
providers = append(providers, api.AgentNetworkAgentConfigProvider{
Name: p.Name,
CatalogId: p.CatalogID,
ApiFlavor: p.APIFlavor,
AllModelsAllowed: p.AllModelsAllowed,
Models: p.Models,
})
}
return api.AgentNetworkAgentConfig{
Configured: setup.Configured,
Endpoint: setup.Endpoint,
Providers: providers,
}
}
@@ -9,14 +9,16 @@ import (
"runtime"
"strings"
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/server/account"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/management/server/permissions"
@@ -29,6 +31,9 @@ import (
const (
testAccountID = "acc-1"
testUserID = "user-bob"
// testClusterAddress is the shared proxy cluster the settings tests pin
// their gateway to; the fixture seeds a connected private-capable proxy for it.
testClusterAddress = "eu.proxy.netbird.io"
)
// agentNetworkHandlerFixture builds a real agentnetwork.Manager with
@@ -75,6 +80,12 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture {
manager := agentnetwork.NewManager(st, perms, accounts, nil)
h := &handler{manager: manager}
// The labeled bootstrap validates its proxy_address against the live
// clusters, so seed the shared cluster these tests pin to as a real,
// private-capable one — the wire-shape assertions then run through the
// validated path rather than the "nothing connected yet" carve-out.
seedSharedPrivateCluster(t, st, testClusterAddress)
router := mux.NewRouter()
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST")
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET")
@@ -268,3 +279,21 @@ func TestConsumptionHandler_PopulatedAccountListsRows(t *testing.T) {
assert.Equal(t, groupRow.WindowStartUtc, userRow.WindowStartUtc,
"rows recorded in the same window must share the aligned window_start_utc")
}
// seedSharedPrivateCluster registers a connected, NetBird-operated proxy
// with private capabilities (the `private` capability) so
// clusterAddr is a cluster any account may pin its agent-network gateway to.
func seedSharedPrivateCluster(t *testing.T, st store.Store, clusterAddr string) {
t.Helper()
private := true
now := time.Now().UTC()
require.NoError(t, st.SaveProxy(context.Background(), &rpproxy.Proxy{
ID: "shared-proxy-" + clusterAddr,
SessionID: "shared-session",
ClusterAddress: clusterAddr,
LastSeen: now,
ConnectedAt: &now,
Status: rpproxy.StatusConnected,
Capabilities: rpproxy.Capabilities{Private: &private},
}), "seeding the shared proxy cluster must succeed")
}
@@ -0,0 +1,178 @@
package handlers
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// discoveryManagerStub records what the handler asked for and returns a canned
// answer. The Manager interface is embedded rather than implemented: only the
// one method is reachable from this handler, and a call to any other should
// fail loudly rather than silently return a zero value.
type discoveryManagerStub struct {
agentnetwork.Manager
gotReq modeldiscovery.Request
gotRecordID string
models []modeldiscovery.Model
err error
}
func (s *discoveryManagerStub) DiscoverProviderModels(
_ context.Context, _, _ string, req modeldiscovery.Request, recordID string,
) ([]modeldiscovery.Model, error) {
s.gotReq = req
s.gotRecordID = recordID
return s.models, s.err
}
// postDiscovery drives the handler with an authenticated request.
func postDiscovery(t *testing.T, stub *discoveryManagerStub, body string) *httptest.ResponseRecorder {
t.Helper()
h := &handler{manager: stub}
req := httptest.NewRequest(http.MethodPost, "/agent-network/catalog/providers/models", strings.NewReader(body))
req = req.WithContext(nbcontext.SetUserAuthInContext(req.Context(), auth.UserAuth{
AccountId: "acc-1",
UserId: "user-1",
}))
rec := httptest.NewRecorder()
h.discoverProviderModels(rec, req)
return rec
}
func TestDiscoverModelsReturnsTheVendorList(t *testing.T) {
stub := &discoveryManagerStub{models: []modeldiscovery.Model{
{ID: "eu.anthropic.claude-haiku-4-5-20251001-v1:0", Label: "EU Claude Haiku 4.5", PricingKnown: true},
{ID: "global.cohere.embed-v4:0", Label: "Global Cohere Embed v4"},
// A vendor that supplies no display name at all. Bedrock does for
// every profile, but the OpenAI listing carries none.
{ID: "gpt-4o-mini", PricingKnown: true},
}}
rec := postDiscovery(t, stub, `{
"catalog_provider_id":"bedrock_api",
"upstream_url":"https://bedrock-runtime.eu-central-1.amazonaws.com",
"api_key":"aws-bearer"
}`)
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
var out api.AgentNetworkModelDiscoveryResponse
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &out))
require.Len(t, out.Models, 3)
assert.Equal(t, "eu.anthropic.claude-haiku-4-5-20251001-v1:0", out.Models[0].Id)
assert.True(t, out.Models[0].PricingKnown)
// An unpriced model must say so rather than arriving indistinguishable
// from a priced one: registering it silently would meter at zero.
assert.False(t, out.Models[1].PricingKnown)
require.NotNil(t, out.Models[0].Label, "the vendor supplied a display name")
assert.Equal(t, "EU Claude Haiku 4.5", *out.Models[0].Label)
// A vendor that supplies no name must omit the key rather than send an
// empty string: the dashboard falls back to the id on absence, and would
// render a blank row for "".
assert.Nil(t, out.Models[2].Label, "an absent label must not serialize")
assert.NotContains(t, rec.Body.String(), `"label":""`)
assert.Equal(t, "bedrock_api", stub.gotReq.CatalogID)
assert.Equal(t, "aws-bearer", stub.gotReq.APIKey)
// The upstream is what the region is read back out of for Bedrock, so
// losing it here would break discovery for every regional provider.
assert.Equal(t, "https://bedrock-runtime.eu-central-1.amazonaws.com", stub.gotReq.UpstreamURL)
assert.Empty(t, stub.gotRecordID)
}
func TestDiscoverModelsUsesAStoredRecordWithoutAKey(t *testing.T) {
stub := &discoveryManagerStub{}
rec := postDiscovery(t, stub, `{"catalog_provider_id":"openai_api","provider_id":"prov-42"}`)
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
// The dashboard refreshes a saved provider's list without ever holding
// the credential, so the record id has to reach the manager.
assert.Equal(t, "prov-42", stub.gotRecordID)
assert.Empty(t, stub.gotReq.APIKey)
}
// TestDiscoverModelsRefusesMixedCredentials covers the case where a caller
// names a saved provider AND supplies a key. Accepting it would run an
// arbitrary credential under the identity of a record the caller may only be
// permitted to read.
func TestDiscoverModelsRefusesMixedCredentials(t *testing.T) {
stub := &discoveryManagerStub{}
rec := postDiscovery(t, stub, `{
"catalog_provider_id":"openai_api",
"provider_id":"prov-42",
"api_key":"sk-attacker"
}`)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Empty(t, stub.gotRecordID, "the request must be refused before it reaches the manager")
}
// TestDiscoverModelsReportsNoDiscoveryDistinctly matters because the caller
// falls back to the catalog's own model list on this outcome. Collapsing it
// into a generic 500 would turn "this provider has no listing endpoint" into
// "something went wrong", and the form would show an error instead of a list.
func TestDiscoverModelsReportsNoDiscoveryDistinctly(t *testing.T) {
stub := &discoveryManagerStub{err: modeldiscovery.ErrNoDiscovery}
rec := postDiscovery(t, stub, `{"catalog_provider_id":"litellm_proxy","upstream_url":"https://gw.example.com","api_key":"sk"}`)
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code)
}
// TestDiscoverModelsTrimsTheCatalogID pins that the id the emptiness check
// accepts is the id the manager receives. A padded value that clears the check
// but reaches the catalog untrimmed misses the lookup, and the operator is told
// their provider does not exist.
func TestDiscoverModelsTrimsTheCatalogID(t *testing.T) {
stub := &discoveryManagerStub{}
rec := postDiscovery(t, stub, `{"catalog_provider_id":" openai_api ","api_key":"sk"}`)
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
assert.Equal(t, "openai_api", stub.gotReq.CatalogID)
}
// TestDiscoverModelsReportsCallerInputAsBadRequest covers the other half of the
// error mapping. These failures are all reachable from a well-formed request
// with a bad field value, so answering 500 both misinforms the operator and
// puts their typo into the server's error rate.
func TestDiscoverModelsReportsCallerInputAsBadRequest(t *testing.T) {
stub := &discoveryManagerStub{
err: fmt.Errorf("%w: unknown catalog provider %q", modeldiscovery.ErrInvalidRequest, "nope"),
}
rec := postDiscovery(t, stub, `{"catalog_provider_id":"nope","api_key":"sk"}`)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Contains(t, rec.Body.String(), "unknown catalog provider")
}
func TestDiscoverModelsRejectsMalformedRequests(t *testing.T) {
for name, body := range map[string]string{
"not json": `{`,
"no catalog provider": `{"api_key":"sk"}`,
"blank catalog provider": `{"catalog_provider_id":" ","api_key":"sk"}`,
} {
t.Run(name, func(t *testing.T) {
stub := &discoveryManagerStub{}
rec := postDiscovery(t, stub, body)
assert.Equal(t, http.StatusBadRequest, rec.Code)
})
}
}
@@ -7,6 +7,7 @@ package handlers
import (
"encoding/json"
"errors"
"math"
"net/http"
"net/url"
@@ -16,6 +17,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
nbcontext "github.com/netbirdio/netbird/management/server/context"
@@ -32,6 +34,7 @@ type handler struct {
func RegisterEndpoints(manager agentnetwork.Manager, router *mux.Router) {
h := &handler{manager: manager}
router.HandleFunc("/agent-network/catalog/providers", h.getCatalogProviders).Methods("GET", "OPTIONS")
router.HandleFunc("/agent-network/catalog/providers/models", h.discoverProviderModels).Methods("POST", "OPTIONS")
router.HandleFunc("/agent-network/providers", h.getAllProviders).Methods("GET", "OPTIONS")
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST", "OPTIONS")
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET", "OPTIONS")
@@ -43,6 +46,7 @@ func RegisterEndpoints(manager agentnetwork.Manager, router *mux.Router) {
h.addConsumptionEndpoints(router)
h.addAccessLogEndpoints(router)
h.addBudgetRuleEndpoints(router)
h.addAgentConfigEndpoints(router)
}
func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) {
@@ -61,6 +65,98 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) {
util.WriteJSONObject(r.Context(), w, out)
}
// discoverProviderModels asks the vendor which models the operator's own
// credential can reach, so the provider form can offer a live list rather than
// only the static catalog.
func (h *handler) discoverProviderModels(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
util.WriteError(r.Context(), err, w)
return
}
var body api.AgentNetworkModelDiscoveryRequest
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
util.WriteErrorResponse("invalid json", http.StatusBadRequest, w)
return
}
// Trimmed once and carried, not trimmed for the emptiness test and then
// discarded: a padded " openai_api " would clear the check here and miss
// the catalog lookup, reporting the provider as unknown.
catalogID := strings.TrimSpace(body.CatalogProviderId)
if catalogID == "" {
util.WriteErrorResponse("catalog_provider_id is required", http.StatusBadRequest, w)
return
}
recordID := strValue(body.ProviderId)
req := modeldiscovery.Request{
CatalogID: catalogID,
UpstreamURL: strValue(body.UpstreamUrl),
APIKey: strValue(body.ApiKey),
}
// One source of credential or the other, never a mix: taking a key from
// the request while addressing a saved record would let a caller run an
// arbitrary credential against a provider they can only read.
if recordID != "" && req.APIKey != "" {
util.WriteErrorResponse("provide either provider_id or api_key, not both", http.StatusBadRequest, w)
return
}
models, err := h.manager.DiscoverProviderModels(r.Context(), userAuth.AccountId, userAuth.UserId, req, recordID)
if err != nil {
// A provider with no listing endpoint is a fact about the catalog
// entry, not a failure: the caller falls back to the catalog's own
// models, so it must be able to tell the two apart.
if errors.Is(err, modeldiscovery.ErrNoDiscovery) {
util.WriteErrorResponse(err.Error(), http.StatusUnprocessableEntity, w)
return
}
// An unknown provider, an unusable upstream, a missing region or a
// missing key are all things the caller sent, reachable from a
// well-formed request. Reporting them as 500 tells the operator the
// server broke and buries genuine faults in the error rate.
if errors.Is(err, modeldiscovery.ErrInvalidRequest) {
util.WriteErrorResponse(err.Error(), http.StatusBadRequest, w)
return
}
util.WriteError(r.Context(), err, w)
return
}
out := api.AgentNetworkModelDiscoveryResponse{Models: make([]api.AgentNetworkDiscoveredModel, 0, len(models))}
for _, m := range models {
entry := api.AgentNetworkDiscoveredModel{
Id: m.ID,
PricingKnown: m.PricingKnown,
// Sent even when zero: the form prefills every discovered model as
// an editable row, and an unpriced one is shown at zero and flagged
// rather than left out.
InputPer1k: m.InputPer1k,
OutputPer1k: m.OutputPer1k,
// Cache rates stay absent when unset, matching the catalog
// response — a zero would read as "free", not "not applicable".
CachedInputPer1k: positiveRatePtr(m.CachedInputPer1k),
CacheReadPer1k: positiveRatePtr(m.CacheReadPer1k),
CacheCreationPer1k: positiveRatePtr(m.CacheCreationPer1k),
}
if m.Label != "" {
label := m.Label
entry.Label = &label
}
out.Models = append(out.Models, entry)
}
util.WriteJSONObject(r.Context(), w, out)
}
// strValue reads an optional string field, treating absent as empty.
func strValue(v *string) string {
if v == nil {
return ""
}
return strings.TrimSpace(*v)
}
// applyDefaultPricing overwrites the catalog response's model rates with
// the LIVE default pricing table, which may differ from the compiled-in
// catalog rates when the operator provides a defaults_llm_pricing.yaml.
@@ -244,6 +340,14 @@ func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error {
if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") {
return status.Errorf(status.InvalidArgument, "api_key is required")
}
// An update omits api_key to keep the stored credential. A key that is
// present but blank is not that: Provider.FromAPIRequest drops it exactly
// as if it were absent, so a rotation the operator believes they performed
// would answer 200 having changed nothing. Refuse it here, where the
// request still carries the difference between absent and blank.
if req.ApiKey != nil && strings.TrimSpace(*req.ApiKey) == "" {
return status.Errorf(status.InvalidArgument, "api_key must be omitted to keep the stored credential rather than sent blank")
}
if req.Models != nil {
for i, m := range *req.Models {
if err := validateModel(i, m); err != nil {
@@ -54,6 +54,39 @@ func TestValidate_ModelRates(t *testing.T) {
}
}
// TestValidate_ABlankApiKeyIsNotTheSameAsAnOmittedOne covers the one shape the
// manager's own guard cannot see. Provider.FromAPIRequest assigns the key only
// when it trims to something, so a request carrying " " arrives at
// UpdateProvider indistinguishable from one that omitted it — the stored
// credential is kept and the write answers 200, telling an operator who thinks
// they just rotated a key that it worked.
//
// The request still knows the difference, so the refusal belongs here.
func TestValidate_ABlankApiKeyIsNotTheSameAsAnOmittedOne(t *testing.T) {
req := func(key *string) *api.AgentNetworkProviderRequest {
return &api.AgentNetworkProviderRequest{
ProviderId: "openai_api",
Name: "OpenAI",
UpstreamUrl: "https://api.openai.com",
ApiKey: key,
}
}
blank := " "
err := validate(req(&blank), false)
require.Error(t, err, "a blank api_key on update must not be read as 'keep what is stored'")
assert.Contains(t, err.Error(), "api_key")
require.NoError(t, validate(req(nil), false), "an omitted api_key is how an update keeps the stored credential")
// Create already refuses this, and keeps its own message: a caller who sent
// no usable key is told the field is required rather than being told how to
// preserve a credential that does not exist yet.
err = validate(req(&blank), true)
require.Error(t, err)
assert.Contains(t, err.Error(), "api_key is required")
}
// TestProviderHandler_UpdateReplacesFullState pins the update contract shared
// with the other PUT endpoints: the request replaces the provider's mutable
// state, so optional fields absent from the JSON land as their zero values.
@@ -64,10 +97,13 @@ func TestValidate_ModelRates(t *testing.T) {
func TestProviderHandler_UpdateReplacesFullState(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
// A private upstream: the save-time credential check leaves it unchecked
// rather than spending "sk-test" against the real api.openai.com, which
// the vendor refuses.
create := `{
"provider_id": "openai_api",
"name": "openai",
"upstream_url": "https://api.openai.com",
"upstream_url": "https://10.255.255.1",
"api_key": "sk-test",
"enabled": true,
"metadata_disabled": true,
@@ -84,7 +120,7 @@ func TestProviderHandler_UpdateReplacesFullState(t *testing.T) {
// Minimal update: only the required fields, no api_key. Everything
// optional must land as its zero value.
update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://api.openai.com", "enabled": true}`
update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://10.255.255.1", "enabled": true}`
rec = f.do(t, nethttp.MethodPut, "/agent-network/providers/"+created.Id, update)
require.Equal(t, nethttp.StatusOK, rec.Code, "update without api_key must succeed (key is preserved): %s", rec.Body.String())
@@ -3,9 +3,10 @@ package labelgen
import (
"fmt"
"math/rand"
"sort"
"sync"
"github.com/netbirdio/netbird/management/server/util"
)
// pickAttempts caps the random retries before falling back to the
@@ -40,16 +41,15 @@ func uniqueWords() []string {
// PickUnique selects a label not already in `taken`. It tries up to
// pickAttempts random picks; on exhaustion it scans the deduplicated
// wordlist for any remaining free entry, and if none is left appends
// `-<fallbackSuffix>` to a deterministic word and returns. The caller
// is responsible for seeding rng (math/rand).
func PickUnique(rng *rand.Rand, taken map[string]struct{}, fallbackSuffix string) string {
// `-<fallbackSuffix>` to a random word and returns.
func PickUnique(taken map[string]struct{}, fallbackSuffix string) string {
pool := uniqueWords()
if len(pool) == 0 {
return fallbackSuffix
}
for i := 0; i < pickAttempts; i++ {
w := pool[rng.Intn(len(pool))]
w := pool[util.RandIntn(len(pool))]
if _, ok := taken[w]; !ok {
return w
}
@@ -61,7 +61,7 @@ func PickUnique(rng *rand.Rand, taken map[string]struct{}, fallbackSuffix string
}
}
w := pool[rng.Intn(len(pool))]
w := pool[util.RandIntn(len(pool))]
return fmt.Sprintf("%s-%s", w, fallbackSuffix)
}
@@ -74,10 +74,10 @@ func PickUnique(rng *rand.Rand, taken map[string]struct{}, fallbackSuffix string
// a noun spans len(adjectives) * 857 instead. Uniqueness is enforced by a
// database constraint and retried by the caller, rather than guessed from a
// pre-read set that a concurrent allocation can invalidate.
func PickTuple(rng *rand.Rand) string {
func PickTuple() string {
nouns := uniqueWords()
if len(nouns) == 0 || len(adjectives) == 0 {
return ""
}
return adjectives[rng.Intn(len(adjectives))] + "-" + nouns[rng.Intn(len(nouns))]
return adjectives[util.RandIntn(len(adjectives))] + "-" + nouns[util.RandIntn(len(nouns))]
}
@@ -1,7 +1,7 @@
package labelgen
import (
"math/rand"
"slices"
"strings"
"testing"
@@ -9,19 +9,12 @@ import (
"github.com/stretchr/testify/require"
)
// TestPickUnique_DeterministicWithSeededRng locks the property the
// caller relies on: same seed + same taken set → same pick. Without
// that, the bootstrap flow can't reproduce a label across retries.
func TestPickUnique_DeterministicWithSeededRng(t *testing.T) {
taken := map[string]struct{}{}
// TestPickUnique_ReturnsWordFromPool confirms a pick against an empty
// taken set is always drawn verbatim from the wordlist.
func TestPickUnique_ReturnsWordFromPool(t *testing.T) {
got := PickUnique(map[string]struct{}{}, "abcd")
rngA := rand.New(rand.NewSource(42))
rngB := rand.New(rand.NewSource(42))
a := PickUnique(rngA, taken, "abcd")
b := PickUnique(rngB, taken, "abcd")
assert.Equal(t, a, b, "Same seed and taken set must produce identical pick")
assert.True(t, slices.Contains(uniqueWords(), got), "Pick %q must be drawn from the wordlist", got)
}
// TestPickUnique_AvoidsTakenWordsWhenMostAreReserved seeds taken with
@@ -46,8 +39,7 @@ func TestPickUnique_AvoidsTakenWordsWhenMostAreReserved(t *testing.T) {
taken[w] = struct{}{}
}
rng := rand.New(rand.NewSource(7))
got := PickUnique(rng, taken, "abcd")
got := PickUnique(taken, "abcd")
_, isFree := free[got]
assert.True(t, isFree, "PickUnique must return one of the free words; got %q", got)
@@ -65,8 +57,7 @@ func TestPickUnique_FallsBackWhenAllReserved(t *testing.T) {
taken[w] = struct{}{}
}
rng := rand.New(rand.NewSource(99))
got := PickUnique(rng, taken, "abcd")
got := PickUnique(taken, "abcd")
assert.True(t, strings.HasSuffix(got, "-abcd"), "Exhausted pool must produce <word>-<suffix>; got %q", got)
@@ -114,9 +105,8 @@ func TestPickTuple_ShapeAndPoolMembership(t *testing.T) {
inAdjectives[a] = struct{}{}
}
rng := rand.New(rand.NewSource(7))
for i := 0; i < 200; i++ {
got := PickTuple(rng)
got := PickTuple()
parts := strings.Split(got, "-")
require.Len(t, parts, 2, "PickTuple must produce exactly two hyphen-joined words; got %q", got)
@@ -158,22 +148,13 @@ func TestAdjectives_AreDNSSafeAndDeduplicated(t *testing.T) {
assert.Greater(t, len(adjectives), 150, "Adjective pool too small to give a useful namespace")
}
// TestPickTuple_DeterministicWithSeededRng documents that generation is a pure
// function of the rng, which is what makes allocation retries reproducible in tests.
func TestPickTuple_DeterministicWithSeededRng(t *testing.T) {
a := PickTuple(rand.New(rand.NewSource(42)))
b := PickTuple(rand.New(rand.NewSource(42)))
assert.Equal(t, a, b, "Same seed must yield the same tuple")
}
// TestPickTuple_SpansALargeNamespace guards the reason we moved to tuples: a
// single-word pool caps the GLOBAL namespace at 857. Drawing many tuples must
// yield overwhelmingly distinct values.
func TestPickTuple_SpansALargeNamespace(t *testing.T) {
rng := rand.New(rand.NewSource(11))
seen := make(map[string]struct{}, 2000)
for i := 0; i < 2000; i++ {
seen[PickTuple(rng)] = struct{}{}
seen[PickTuple()] = struct{}{}
}
assert.Greater(t, len(seen), 1900,
"2000 draws should be nearly all distinct across a ~200k namespace; got %d unique", len(seen))
@@ -4,7 +4,6 @@ import (
"context"
"errors"
"fmt"
"math/rand"
"slices"
"strings"
"sync"
@@ -13,6 +12,7 @@ import (
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/labelgen"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
@@ -50,6 +50,7 @@ type Manager interface {
CreateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
DeleteProvider(ctx context.Context, accountID, userID, providerID string) error
DiscoverProviderModels(ctx context.Context, accountID, userID string, req modeldiscovery.Request, recordID string) ([]modeldiscovery.Model, error)
GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error)
GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error)
@@ -83,6 +84,13 @@ type Manager interface {
RecordAccountBudgetUsage(ctx context.Context, accountID, userID string, groupIDs []string, tokensIn, tokensOut int64, costUSD float64) error
RecordUsage(ctx context.Context, in RecordUsageInput) error
SelectPolicyForRequest(ctx context.Context, in PolicySelectionInput) (*PolicySelectionResult, error)
// GetAgentConfigForUser backs the self-service agent-config endpoint.
// Caller-scoped, so it skips the role permission gate; see
// the implementation. The caller's own usage and requests come
// through GetUsageOverview / ListAccessLogs, which self-scope when
// the account-wide grant is missing.
GetAgentConfigForUser(ctx context.Context, accountID, userID string) (*types.AgentConfig, error)
}
// PolicySelectionInput is the per-request selection envelope. The
@@ -123,16 +131,35 @@ type managerImpl struct {
permissionsManager permissions.Manager
proxyController proxy.Controller
// modelDiscovery queries vendors for the models a credential can reach.
// An interface rather than the concrete client because it is now on a
// write path: the credential check runs inside CreateProvider and
// UpdateProvider, so every test that saves a provider would otherwise
// reach a vendor over the network to do it.
//
// One instance serves every request for the process's lifetime, so its
// fields must stay read-only after construction: lazy initialisation
// inside Fetch or httpClient would race across request goroutines.
modelDiscovery ModelLister
// reconcileCache holds the last set of synthesised proxy mappings
// per account, each paired with the proxy that served it, so a change
// of serving proxy can be diffed without re-deriving it.
reconcileMu sync.Mutex
reconcileCache map[string]map[string]syntheticMapping
}
// labelRngMu guards labelRng. PickUnique consumes math/rand.Source
// state; concurrent provider creates would otherwise race.
labelRngMu sync.Mutex
labelRng *rand.Rand
// ManagerOption replaces a manager dependency at construction. Production
// passes none; each option exists for something a test cannot let run for
// real.
type ManagerOption func(*managerImpl)
// WithModelLister replaces the vendor call behind the provider credential
// check. A test that saves a provider needs this — the check runs inside
// CreateProvider and UpdateProvider, so the write path reaches a vendor
// without it.
func WithModelLister(lister ModelLister) ManagerOption {
return func(m *managerImpl) { m.modelDiscovery = lister }
}
// NewManager constructs the persistent Agent Network manager. The
@@ -145,29 +172,192 @@ func NewManager(
permissionsManager permissions.Manager,
accountManager account.Manager,
proxyController proxy.Controller,
opts ...ManagerOption,
) Manager {
return &managerImpl{
m := &managerImpl{
store: store,
accountManager: accountManager,
permissionsManager: permissionsManager,
proxyController: proxyController,
modelDiscovery: &modeldiscovery.Client{},
reconcileCache: make(map[string]map[string]syntheticMapping),
labelRng: rand.New(rand.NewSource(time.Now().UnixNano())),
}
for _, opt := range opts {
opt(m)
}
return m
}
// GetAllProviders returns the account's providers for callers holding the
// providers read grant (connection config redacted unless they can also
// update). A caller without the grant self-scopes instead of being denied
// — mirroring the usage and log endpoints: they get the providers their
// own policies authorize, redacted to the display surface, which is what
// feeds the dashboard's provider filter for plain users.
func (m *managerImpl) GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil {
ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read)
if err != nil {
return nil, status.NewPermissionValidationError(err)
}
if !ok {
return m.callerScopedProviders(ctx, accountID, userID)
}
providers, err := m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID)
if err != nil {
return nil, err
}
return m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID)
return m.redactProvidersForViewer(ctx, accountID, userID, providers)
}
// GetProvider self-scopes like GetAllProviders: a caller without the read
// grant may fetch a provider their own policies authorize (redacted), and
// gets the same not-found answer for any other id — an out-of-scope
// provider must be indistinguishable from a nonexistent one.
func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil {
ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read)
if err != nil {
return nil, status.NewPermissionValidationError(err)
}
if !ok {
scoped, err := m.callerScopedProviders(ctx, accountID, userID)
if err != nil {
return nil, err
}
for _, p := range scoped {
if p.ID == providerID {
return p, nil
}
}
return nil, status.NewAgentNetworkProviderNotFoundError(providerID)
}
provider, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
if err != nil {
return nil, err
}
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
redacted, err := m.redactProvidersForViewer(ctx, accountID, userID, []*types.Provider{provider})
if err != nil {
return nil, err
}
return redacted[0], nil
}
// callerScopedProviders returns the providers the caller's own policies
// authorize — the same selection the self-service setup answer and the
// proxy's routing derive from — each reduced to the display surface. No
// role permission is needed: the answer is scoped strictly to the caller,
// and a caller outside every policy gets an empty list, indistinguishable
// from an account with nothing configured.
func (m *managerImpl) callerScopedProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) {
user, err := m.store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
if err != nil {
return nil, fmt.Errorf("get user: %w", err)
}
authorized, applicable, err := m.authorizedProvidersForGroups(ctx, accountID, user.AutoGroups)
if err != nil {
return nil, err
}
var guardrailsByID map[string]*types.Guardrail
if anyPolicyHasGuardrails(applicable) {
guardrailsByID, err = m.loadGuardrailsByID(ctx, accountID)
if err != nil {
return nil, err
}
}
out := make([]*types.Provider, 0, len(authorized))
for _, p := range authorized {
r := p.RedactedForViewer()
// The model list follows the same effective computation the setup
// answer and the proxy use: allowlist-restricted callers see only
// the models their guardrails permit, and an unrestricted policy
// on a provider without an operator declaration surfaces the
// catalog models, matching the setup response — so the dashboard's
// model filter never offers a model the caller's own requests
// could not use, and never comes up empty when the setup page
// lists models. Grant holders keep the full declared lists —
// their usage view spans everyone's requests.
_, effective := effectiveModelsForProvider(p, policiesForProvider(applicable, p.ID), guardrailsByID)
r.Models = providerModelsByID(p, effective)
out = append(out, r)
}
return out, nil
}
// redactProvidersForViewer strips the connection configuration from
// providers handed to a caller who holds only the read grant on
// agent_network.providers. Update is the managing signal: a role that can
// edit a provider sees its config in the edit form anyway, while a
// read-only role (usage_viewer) only needs the display surface the usage
// filters resolve against — upstream URLs and operator-supplied header
// values are not part of that. Validation errors fail closed.
func (m *managerImpl) redactProvidersForViewer(ctx context.Context, accountID, userID string, providers []*types.Provider) ([]*types.Provider, error) {
canManage, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Update)
if err != nil {
return nil, status.NewPermissionValidationError(err)
}
if canManage {
return providers, nil
}
out := make([]*types.Provider, 0, len(providers))
for _, p := range providers {
if p == nil {
out = append(out, nil)
continue
}
out = append(out, p.RedactedForViewer())
}
return out, nil
}
// DiscoverProviderModels asks the vendor which models a credential can reach.
//
// recordID, when set, names an existing provider whose stored credential is
// used instead of the one in req — so the dashboard can refresh the list
// without ever holding the key. An upstream in req overrides the stored one,
// which is what lets a form list against a URL the operator has typed but not
// saved yet, using the credential they cannot retype.
//
// Gated on Create rather than Read: this spends the operator's credential
// against a third party, which is not something a read-only role should be
// able to make the server do. That one check also covers reading the stored
// record — Create is strictly stronger than Read here, and the lookup is
// scoped to accountID, so another account's record is never reachable.
func (m *managerImpl) DiscoverProviderModels(ctx context.Context, accountID, userID string, req modeldiscovery.Request, recordID string) ([]modeldiscovery.Model, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil {
return nil, err
}
if recordID != "" {
record, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, recordID)
if err != nil {
return nil, err
}
// The catalog id comes from the stored record too: letting the caller
// name a different one would run a provider's credential against
// whichever vendor endpoint they picked.
req.CatalogID = record.ProviderID
req.APIKey = record.APIKey
// The upstream is the one field the caller may override, so that a URL
// typed into the form can be listed against before it is saved.
//
// It sends the stored credential to a host the caller named, which is
// a capability they already have: the same permission set updates the
// record's upstream, and that write runs this same check against
// whatever it is pointed at. What it would not otherwise be is silent,
// since the write leaves an activity event behind — so the override is
// recorded here.
if strings.TrimSpace(req.UpstreamURL) == "" {
req.UpstreamURL = record.UpstreamURL
} else if req.UpstreamURL != record.UpstreamURL {
log.WithContext(ctx).Infof("agent network provider %s listed against caller-supplied upstream %s by user %s",
recordID, req.UpstreamURL, userID)
}
}
models, err := m.modelDiscovery.Fetch(ctx, req)
if err != nil {
return nil, discoveryFailure(ctx, req.CatalogID, err)
}
return models, nil
}
// CreateProvider persists a new provider for the account. Providers have no
@@ -185,6 +375,18 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
if strings.TrimSpace(provider.APIKey) == "" {
return nil, status.Errorf(status.InvalidArgument, "api_key is required when creating an agent network provider")
}
// Stored as it will be sent. The vendor call below trims the key before
// building the auth header while the synthesiser substitutes the stored
// value verbatim, so a key pasted with surrounding whitespace would pass
// its check and then fail every request the provider serves.
provider.APIKey = strings.TrimSpace(provider.APIKey)
// Before anything is persisted: a record whose upstream or credential does
// not work is rejected here rather than discovered later as a failed
// request with nothing pointing back at it.
if err := m.checkProviderCredential(ctx, provider); err != nil {
return nil, err
}
if provider.ID == "" {
fresh := types.NewProvider(provider.AccountID)
@@ -220,11 +422,47 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide
// Preserve the API key if the caller didn't rotate it. A
// whitespace-only value is treated as "not rotated" rather than a
// real key, but it must not silently overwrite a valid stored key.
if provider.APIKey == "" {
provider.APIKey = existing.APIKey
} else if strings.TrimSpace(provider.APIKey) == "" {
switch trimmed := strings.TrimSpace(provider.APIKey); {
case provider.APIKey == "":
// Trimmed on the way through: a record stored before keys were
// normalised carries whitespace the proxy still sends, and an edit
// that preserves the key is the occasion to repair it. Doing so makes
// the comparison below see a change, which is correct — that key has
// never been tested in the form it is about to be sent in.
provider.APIKey = strings.TrimSpace(existing.APIKey)
case trimmed == "":
return nil, status.Errorf(status.InvalidArgument, "api_key must be non-blank when rotating an agent network provider")
default:
// See CreateProvider: the key is stored in the form the proxy will
// send, so the check below tests what the provider will actually use.
provider.APIKey = trimmed
}
// Only the fields the vendor would judge are worth a round-trip. This same
// call carries renames, model rows and price edits, and none of those
// should wait on a vendor — or be refused because one is having a bad day.
//
// The catalog entry counts as one of them: it decides which vendor is
// asked, under which auth header, so moving a record from one to another
// sends an unchanged credential somewhere it has never been accepted.
//
// The comparison runs after the merge above, so an update that changes only
// the URL reads as unchanged on the key and is checked against the stored
// one, which is the only credential the operator has to offer here.
//
// Turning TLS verification back on is the fourth: the record was stored
// unchecked precisely because that flag was set, so this is the first
// moment it can be checked at all, and nothing else about it need change
// for that to be true.
if provider.UpstreamURL != existing.UpstreamURL ||
provider.APIKey != existing.APIKey ||
provider.ProviderID != existing.ProviderID ||
(existing.SkipTLSVerification && !provider.SkipTLSVerification) {
if err := m.checkProviderCredential(ctx, provider); err != nil {
return nil, err
}
}
// Always preserve the session keypair across updates so existing
// session cookies stay valid. The keys are server-managed and
// never surfaced through the API.
@@ -798,6 +1036,18 @@ func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *type
if err != nil {
return status.Errorf(status.InvalidArgument, "invalid endpoint: %s", err)
}
if err := m.requireHostNotForeign(ctx, settings.AccountID, hostname); err != nil {
return err
}
// Another account's labeled pin beneath this hostname makes it their
// cluster: a proxy serving them there would never serve this endpoint.
// The domain unique index already arbitrates two endpoints on one name.
if err := m.requireNotClaimedByOtherAccount(ctx, settings.AccountID, hostname, m.store.HasGatewayClusterPinnedByOtherAccount); err != nil {
return err
}
if err := m.validateGatewayCluster(ctx, settings.AccountID, hostname); err != nil {
return err
}
settings.Domain = hostname
settings.ProxyAddress = hostname
@@ -816,6 +1066,99 @@ func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *type
return nil
}
// validateGatewayCluster rejects a bootstrap pinned to a cluster that cannot
// serve the account's gateway — a labeled endpoint beneath the cluster and a
// self-addressed one on the very address a proxy declares alike, since the
// service behind either is the same private one.
//
// The synthesised gateway service is unconditionally private
// (buildAccountService): agents reach it over the WireGuard tunnel and are
// authorised by ValidateTunnelPeer against the policies' source groups, and
// its single target is the cluster itself with DirectUpstream. Only a cluster
// with private capabilities can serve that. Management reports it per cluster
// as the `private` capability, the same flag the dashboard renders as
// supports_private when it gates NetBird-only services.
//
// Without this check the bootstrap happily pins to any cluster the caller
// names, including one without private capabilities — and the endpoint it
// allocates is immutable, so the account is left with a dead gateway that only
// a DeleteSettings/re-bootstrap can undo.
//
// Whether management knows the cluster is decided on the proxy rows
// themselves, never on how fresh their heartbeats are: a cluster's rows
// outlive its proxies' liveness (only the stale-proxy reaper removes them), so
// a cluster that exists stays judged as one. Judging on liveness instead would
// make the same centralised cluster pass or fail depending on whether its
// proxies happened to have heartbeated in the last couple of minutes.
//
// The single opening left is a cluster management holds no proxy row for at
// all: pinning ahead of a proxy's first connection is a legitimate order — the
// dedicated path claims an address the same way, before any proxy declares it.
func (m *managerImpl) validateGatewayCluster(ctx context.Context, accountID, clusterAddr string) error {
declared, err := m.accountClusterSpellings(ctx, accountID, clusterAddr)
if err != nil {
return err
}
if len(declared) == 0 {
// No proxy has ever declared this address: an address-first pin.
return nil
}
// A cluster management knows has to prove it can serve the gateway, and
// only a live proxy reporting the capability proves that. Both an explicit false and an
// unreported capability (nothing live in the cluster, or proxies predating
// capability reporting) fail here: unusable and unproven are the same
// answer for a decision that cannot be revisited later.
//
// The capability is read per declared spelling and taken as any-true, the
// same way it aggregates over a cluster's proxies: the store matches
// cluster_address exactly, so a host two proxies spelled differently must
// not come back unproven just because it was asked about under one of them.
for _, address := range declared {
if private := m.store.GetClusterSupportsPrivate(ctx, address); private != nil && *private {
return nil
}
}
return status.Errorf(status.InvalidArgument,
"proxy cluster %s has no private capabilities: the agent network gateway requires a reverse proxy cluster "+
"with private capabilities", clusterAddr)
}
// accountClusterSpellings returns every proxy cluster address in the account's
// view — its own (BYOP) clusters plus the shared ones — that names the same
// host as clusterAddr. Empty means management holds no proxy row for that host
// in this account's view.
//
// A proxy declares its cluster address as the operator spelled it, so identity
// is compared on the normalised form rather than byte-equal — an in-memory pass
// over the account's clusters, not a query. What comes back is the stored
// spelling, because the capability lookup matches cluster_address exactly and
// would silently find nothing under a spelling the store never held. The
// cluster listing is not gated on heartbeats, so this answer does not change
// while a cluster's proxies are merely offline.
func (m *managerImpl) accountClusterSpellings(ctx context.Context, accountID, clusterAddr string) ([]string, error) {
clusters, err := m.store.GetProxyClusters(ctx, accountID)
if err != nil {
return nil, fmt.Errorf("list proxy clusters: %w", err)
}
var spellings []string
for _, cluster := range clusters {
normalized, err := types.NormalizeHostname(cluster.Address)
if err != nil {
// An address declared in a shape we cannot normalise is not one an
// endpoint can be allocated beneath.
log.WithContext(ctx).Debugf("skipping unusable proxy cluster address %q: %s", cluster.Address, err)
continue
}
if normalized == clusterAddr {
spellings = append(spellings, cluster.Address)
}
}
return spellings, nil
}
// bootstrapLabeled allocates a labeled endpoint one label beneath the given
// cluster address: Domain = <label>.<proxyAddress>, served by whichever proxy
// declares the parent. Labels are adjective-noun tuples; a candidate is
@@ -827,11 +1170,23 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett
if err != nil {
return status.Errorf(status.InvalidArgument, "invalid proxy_address: %s", err)
}
if err := m.requireHostNotForeign(ctx, settings.AccountID, parent); err != nil {
return err
}
// Another account's endpoint at this exact hostname means the proxy that
// declares it is theirs, so nothing would serve a label beneath it. Other
// accounts' labeled pins under the same cluster are not asked about: a
// shared cluster carries many of them by design.
if err := m.requireNotClaimedByOtherAccount(ctx, settings.AccountID, parent, m.store.HasGatewayEndpointByOtherAccount); err != nil {
return err
}
if err := m.validateGatewayCluster(ctx, settings.AccountID, parent); err != nil {
return err
}
for attempt := 1; attempt <= maxDomainAllocationAttempts; attempt++ {
m.labelRngMu.Lock()
label := labelgen.PickTuple(m.labelRng)
m.labelRngMu.Unlock()
label := labelgen.PickTuple()
if label == "" {
// Only reachable if either word pool were emptied. An empty label
// would produce a broken endpoint like ".example.com", so fail
@@ -875,6 +1230,41 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett
return fmt.Errorf("allocate agent network endpoint for account %s: %d attempts exhausted", settings.AccountID, maxDomainAllocationAttempts)
}
// requireHostNotForeign refuses to pin the account's gateway onto a host that
// another account's proxy declares. The pin's proxy_address is what selects
// the proxy that serves the endpoint, and an account-scoped proxy only ever
// receives its own account's mappings, so such a pin could never be served —
// and the endpoint it assigns is immutable. Shared proxies are not foreign, and
// a host no proxy has declared stays pinnable: claiming the address before the
// proxy's first connection is the documented order.
func (m *managerImpl) requireHostNotForeign(ctx context.Context, accountID, host string) error {
foreign, err := m.store.HasForeignAccountProxyAtHost(ctx, host, accountID)
if err != nil {
return fmt.Errorf("check proxy host ownership: %w", err)
}
if foreign {
return errHostNotAvailable(host)
}
return nil
}
// requireNotClaimedByOtherAccount refuses the pin when another account's
// gateway settings already claim the host in the shape claimed answers for.
func (m *managerImpl) requireNotClaimedByOtherAccount(ctx context.Context, accountID, host string, claimed func(context.Context, string, string) (bool, error)) error {
taken, err := claimed(ctx, host, accountID)
if err != nil {
return fmt.Errorf("check agent network gateway claims at host: %w", err)
}
if taken {
return errHostNotAvailable(host)
}
return nil
}
func errHostNotAvailable(host string) error {
return status.Errorf(status.InvalidArgument, "proxy cluster %s is not available to this account", host)
}
// isUniqueConstraintError reports whether err is a database unique-constraint
// violation, matched on the driver message because CreateAgentNetworkSettings
// deliberately returns the driver error unwrapped.
@@ -901,8 +1291,11 @@ func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID str
// ListAccessLogs returns a paginated, server-side-filtered page of
// agent-network access logs plus the total count matching the filter.
// Callers without the account-wide logs grant get a self-scoped page —
// only their own requests — instead of a denial.
func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil {
filter, err := m.scopeFilterToCaller(ctx, accountID, userID, modules.AgentNetworkLogs, filter)
if err != nil {
return nil, 0, err
}
return m.store.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, filter)
@@ -910,18 +1303,23 @@ func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID stri
// ListAccessLogSessions returns a paginated, server-side-filtered page of
// agent-network access logs grouped by session, plus the total number of
// sessions matching the filter.
// sessions matching the filter. Self-scoped like ListAccessLogs for
// callers without the account-wide logs grant.
func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil {
filter, err := m.scopeFilterToCaller(ctx, accountID, userID, modules.AgentNetworkLogs, filter)
if err != nil {
return nil, 0, err
}
return m.store.GetAgentNetworkAccessLogSessions(ctx, store.LockingStrengthNone, accountID, filter)
}
// GetUsageOverview returns the filtered usage rows aggregated into time buckets
// at the requested granularity, oldest-first.
// at the requested granularity, oldest-first. Callers without the
// account-wide usage grant get their own rows aggregated instead of a
// denial, so the dashboard serves "my usage" from the same endpoint.
func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil {
filter, err := m.scopeFilterToCaller(ctx, accountID, userID, modules.AgentNetworkUsage, filter)
if err != nil {
return nil, err
}
rows, err := m.store.GetAgentNetworkUsageRows(ctx, store.LockingStrengthNone, accountID, filter)
@@ -931,6 +1329,25 @@ func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID st
return types.AggregateUsageByGranularity(rows, granularity), nil
}
// scopeFilterToCaller applies the account-wide read gate for module and,
// when the caller lacks the grant, pins the filter to the caller instead
// of denying: their own user id replaces any requested one and group
// filters are dropped. A caller may always see their own rows — strictly
// tighter than any role gate — which is what lets every authenticated
// user read their usage and requests through the regular endpoints.
// Validation errors (not denials) still fail closed.
func (m *managerImpl) scopeFilterToCaller(ctx context.Context, accountID, userID string, module modules.Module, filter types.AgentNetworkAccessLogFilter) (types.AgentNetworkAccessLogFilter, error) {
ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, module, operations.Read)
if err != nil {
return filter, status.NewPermissionValidationError(err)
}
if !ok {
filter.UserID = &userID
filter.GroupIDs = nil
}
return filter, nil
}
// StartAccessLogCleanup launches a background sweep that periodically deletes
// each account's agent-network access-log rows older than that account's
// AccessLogRetentionDays. Usage records are never swept. A non-positive
@@ -1017,6 +1434,10 @@ func (*mockManager) GetAllProviders(_ context.Context, _, _ string) ([]*types.Pr
return []*types.Provider{}, nil
}
func (*mockManager) DiscoverProviderModels(_ context.Context, _, _ string, _ modeldiscovery.Request, _ string) ([]modeldiscovery.Model, error) {
return nil, nil
}
func (*mockManager) GetProvider(_ context.Context, _, _, _ string) (*types.Provider, error) {
return &types.Provider{}, nil
}
@@ -0,0 +1,521 @@
// Package modeldiscovery asks a vendor which models an operator's own
// credential can reach, so the provider form can offer a live list instead of
// only the catalog's hand-curated one.
//
// The catalog cannot know two things that matter. It goes stale — its entries
// carry comments tracking which models a vendor retired on which date — and it
// cannot see an account: which OpenAI models an org is entitled to, which
// Bedrock inference profiles a given account and region hold, which Vertex
// models a project has enabled. Those are exactly the facts an operator needs
// when filling in a provider record, and only the vendor has them.
//
// The vendor is authoritative for the model ID. The catalog remains
// authoritative for pricing, and a discovered model the catalog cannot price
// is reported as such rather than silently registered at a rate of zero.
package modeldiscovery
import (
"context"
"encoding/base64"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/netip"
"net/url"
"strings"
"syscall"
"time"
"golang.org/x/oauth2/google"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
)
const (
// fetchTimeout bounds one vendor call end to end. A listing is a single
// small GET; anything slower is a vendor problem and the operator is
// waiting on a form.
fetchTimeout = 8 * time.Second
// maxListingBytes bounds the response we will buffer. The largest real
// listing observed is Bedrock's foundation-model catalogue at ~70KB, so
// this is a wide margin over anything legitimate.
maxListingBytes = 2 << 20
// gcpScope matches the scope llm_router mints Vertex tokens under, so a
// credential that works for discovery works for inference too.
gcpScope = "https://www.googleapis.com/auth/cloud-platform"
// vertexKeyfilePrefix marks an api_key that is a base64 service-account
// JSON key rather than a bearer token.
vertexKeyfilePrefix = "keyfile::"
)
// ErrNoDiscovery is returned for a catalog entry that declares no listing
// endpoint. The caller should fall back to the catalog list plus free-text
// entry rather than treating this as a failure.
var ErrNoDiscovery = errors.New("provider has no model-discovery endpoint")
// ErrInvalidRequest marks a discovery failure caused by the caller's own input
// rather than by the vendor or by this server. Every one of these is reachable
// from a well-formed request carrying a bad field value, so the handler owes
// the caller a 400 — a 500 would both misinform them and bury real server
// faults in the error rate.
var ErrInvalidRequest = errors.New("invalid discovery request")
// Model is one discovered model.
type Model struct {
// ID is the identifier to register on the provider record, in the form the
// vendor issues it. For Bedrock that is the region-prefixed inference
// profile id, which is the only form AWS accepts at invoke time.
ID string
// Label is the vendor's display name where it supplies one.
Label string
// PricingKnown reports whether the shipped pricing table can price this
// model. False means the operator must set rates, or the request would
// meter at zero.
PricingKnown bool
// The rates below are the defaults for this model, taken from the same
// table the proxy bills with, so the form prefills exactly what a request
// would cost. All zero when PricingKnown is false — an unpriced model is
// offered at zero and flagged, rather than withheld: the vendor says the
// credential can reach it, and refusing to show it would hide a model the
// operator genuinely has.
InputPer1k float64
OutputPer1k float64
CachedInputPer1k float64
CacheReadPer1k float64
CacheCreationPer1k float64
}
// Request identifies which vendor to ask and with what credential.
type Request struct {
// CatalogID selects the catalog entry, which supplies the endpoint, the
// auth header and the response shape. The caller never supplies those.
CatalogID string
// UpstreamURL is the provider record's configured upstream. It is used
// only when the catalog entry declares no discovery host of its own.
UpstreamURL string
// Region substitutes the catalog host's <region> placeholder.
Region string
// APIKey is the operator's credential, exactly as stored on the record.
APIKey string
}
// Client fetches model listings. The zero value is usable; Resolver and
// HTTPClient exist so tests can drive it against a local server.
type Client struct {
HTTPClient *http.Client
// Resolver looks up the host for the SSRF check. Nil uses the default.
Resolver *net.Resolver
// AllowPrivateHosts disables the private-address guard. Only tests set it:
// their server is on loopback, which is precisely what the guard blocks.
AllowPrivateHosts bool
}
// Fetch returns the models the credential can reach.
func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
entry, ok := catalog.Lookup(req.CatalogID)
if !ok {
return nil, fmt.Errorf("%w: unknown catalog provider %q", ErrInvalidRequest, req.CatalogID)
}
if entry.Discovery == nil {
return nil, ErrNoDiscovery
}
// One deadline over the whole operation. Both host lookups and the request
// itself run under it, so a vendor cannot be slow twice, and a caller that
// gives up is not left waiting on a resolver.
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
defer cancel()
// An entry with a listing host of its own answers from somewhere other
// than the upstream on the record — Bedrock lists from the control plane
// and infers on the runtime host. Reaching the listing therefore proves
// nothing about the host requests will actually go to, so that one is
// checked separately or not at all.
if entry.Discovery.Host != "" {
if err := c.checkUpstreamHost(ctx, entry, req.UpstreamURL); err != nil {
return nil, err
}
}
endpoint, err := c.discoveryURL(ctx, entry, req)
if err != nil {
return nil, err
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, fmt.Errorf("build discovery request: %w", err)
}
if err := applyAuth(httpReq, entry, req.APIKey); err != nil {
return nil, err
}
for name, value := range entry.Discovery.Headers {
httpReq.Header.Set(name, value)
}
httpReq.Header.Set("Accept", "application/json")
resp, err := c.httpClient().Do(httpReq)
if err != nil {
return nil, &UnreachableError{Provider: entry.Name, Err: err}
}
defer func() { _ = resp.Body.Close() }()
body, err := io.ReadAll(io.LimitReader(resp.Body, maxListingBytes))
if err != nil {
return nil, fmt.Errorf("read %s listing: %w", entry.Name, err)
}
if resp.StatusCode != http.StatusOK {
// Surface the vendor's own status. An operator whose key lacks a scope
// needs to see 403 rather than a generic failure.
return nil, &VendorStatusError{Provider: entry.Name, Status: resp.StatusCode}
}
ids, err := parseListing(entry.Discovery.Shape, body)
if err != nil {
return nil, err
}
return decorate(entry, ids), nil
}
// discoveryURL builds the listing URL and refuses one that does not point at a
// public host.
//
// The path, query and (for Bedrock) the host all come from the catalog rather
// than from the caller, so the only operator-controlled part is the host of an
// entry whose listing lives on its own upstream. That still has to be checked:
// management holds credentials for every provider, and an upstream pointed at
// an internal address would turn this endpoint into a probe of the management
// server's own network.
func (c *Client) discoveryURL(ctx context.Context, entry catalog.Provider, req Request) (string, error) {
host := entry.Discovery.Host
if host == "" {
parsed, err := url.Parse(strings.TrimSpace(req.UpstreamURL))
if err != nil || parsed.Host == "" {
// The URL is left out of the message on purpose: it reaches the
// operator through an endpoint that does not lowercase it, but the
// rest of this feature's copy never echoes what they typed, and one
// path that does is the one that ends up quoted in a bug report.
return "", fmt.Errorf("%w: the provider upstream is not a usable URL", ErrInvalidRequest)
}
host = parsed.Host
}
if strings.Contains(host, catalog.RegionPlaceholder) {
region := strings.TrimSpace(req.Region)
if region == "" {
// A provider record carries no region field: the region lives
// inside the upstream host the operator already configured, so
// read it back out rather than asking them for it twice.
region = RegionFromUpstream(entry, req.UpstreamURL)
}
if region == "" {
return "", fmt.Errorf("%w: %w: %s discovery needs a region, and none could be read from the provider upstream",
ErrInvalidRequest, ErrNoDiscoveryHost, entry.Name)
}
host = strings.ReplaceAll(host, catalog.RegionPlaceholder, region)
}
target := &url.URL{Scheme: "https", Host: host, Path: entry.Discovery.Path, RawQuery: entry.Discovery.Query}
if err := c.classifyHost(ctx, entry, target.Hostname()); err != nil {
return "", err
}
return target.String(), nil
}
// checkUpstreamHost verifies the host the operator configured, for entries
// whose listing lives elsewhere and so cannot vouch for it.
//
// A name that does not resolve is the record being wrong. One that resolves
// privately is not: an upstream behind a proxy is a supported configuration,
// and ErrPrivateHost carries that difference on to the caller, which treats it
// as unverifiable rather than as a failure.
func (c *Client) checkUpstreamHost(ctx context.Context, entry catalog.Provider, upstreamURL string) error {
parsed, err := url.Parse(strings.TrimSpace(upstreamURL))
if err != nil || parsed.Hostname() == "" {
return fmt.Errorf("%w: the provider upstream is not a usable URL", ErrInvalidRequest)
}
return c.classifyHost(ctx, entry, parsed.Hostname())
}
// classifyHost renders a failed host check as the two outcomes the caller
// distinguishes. A host that refuses to resolve is the commonest way for an
// upstream to be wrong and has to arrive as unreachable rather than as an
// unclassified fault. ErrPrivateHost means something else entirely — not a bad
// host, one we decline to dial.
func (c *Client) classifyHost(ctx context.Context, entry catalog.Provider, host string) error {
err := c.checkPublicHost(ctx, host)
if err == nil || errors.Is(err, ErrPrivateHost) {
return err
}
return &UnreachableError{Provider: entry.Name, Err: err}
}
// RegionFromUpstream recovers the region an operator embedded in the provider
// upstream, by matching it against the catalog's own host template. Bedrock's
// template is "bedrock-runtime.<region>.amazonaws.com" and Vertex's is
// "<region>-aiplatform.googleapis.com", so the region is whatever sits between
// the fixed halves. Returns empty when the upstream does not match the
// template, which is the case for a custom or proxied endpoint.
func RegionFromUpstream(entry catalog.Provider, upstreamURL string) string {
prefix, suffix, found := strings.Cut(entry.DefaultHost, catalog.RegionPlaceholder)
if !found {
return ""
}
parsed, err := url.Parse(strings.TrimSpace(upstreamURL))
if err != nil {
return ""
}
host := parsed.Hostname()
if host == "" {
// A bare host with no scheme parses as a path, not a host.
host = strings.TrimSpace(upstreamURL)
}
// The two halves must not overlap. "bedrock-runtime.amazonaws.com" carries
// both of Bedrock's — it is the regionless endpoint — and satisfies both
// checks above while leaving nothing between them, so slicing it would
// panic on an inverted range rather than report "no region here".
if !strings.HasPrefix(host, prefix) || !strings.HasSuffix(host, suffix) ||
len(host) < len(prefix)+len(suffix) {
return ""
}
region := host[len(prefix) : len(host)-len(suffix)]
if region == "" || strings.Contains(region, ".") {
return ""
}
return region
}
// checkPublicHost refuses hosts that resolve to an address the management
// server should never be asked to reach on an operator's behalf.
func (c *Client) checkPublicHost(ctx context.Context, host string) error {
if c.AllowPrivateHosts {
return nil
}
if host == "" {
return errors.New("discovery host is empty")
}
resolver := c.Resolver
if resolver == nil {
resolver = net.DefaultResolver
}
addrs, err := resolver.LookupNetIP(ctx, "ip", host)
if err != nil {
return fmt.Errorf("resolve discovery host %q: %w", host, err)
}
// Every address must be public: a name that resolves to one public and one
// loopback address is still a way to reach loopback.
for _, addr := range addrs {
if !isPublic(addr) {
return fmt.Errorf("%w: %w: discovery host %q resolves to a non-public address", ErrInvalidRequest, ErrPrivateHost, host)
}
}
return nil
}
// isPublic reports whether an address is one we are willing to dial.
func isPublic(addr netip.Addr) bool {
addr = addr.Unmap()
switch {
case !addr.IsValid(),
addr.IsLoopback(),
addr.IsPrivate(),
addr.IsLinkLocalUnicast(),
addr.IsLinkLocalMulticast(),
addr.IsInterfaceLocalMulticast(),
addr.IsMulticast(),
addr.IsUnspecified():
return false
}
// 100.64.0.0/10 (carrier NAT) is where NetBird's own overlay addresses
// live, so it is emphatically not somewhere to send a provider credential.
if addr.Is4() {
b := addr.As4()
if b[0] == 100 && b[1] >= 64 && b[1] <= 127 {
return false
}
}
return true
}
// applyAuth sets the credential header the catalog entry declares. A Vertex
// service-account key is exchanged for an OAuth token first, the same way the
// proxy does at request time.
func applyAuth(req *http.Request, entry catalog.Provider, apiKey string) error {
key := strings.TrimSpace(apiKey)
if key == "" {
return fmt.Errorf("%w: %s discovery needs an API key", ErrInvalidRequest, entry.Name)
}
if rest, ok := strings.CutPrefix(key, vertexKeyfilePrefix); ok {
token, err := mintGCPToken(req.Context(), rest)
if err != nil {
return err
}
key = token
}
name := entry.AuthHeaderName
if name == "" {
name = "Authorization"
}
template := entry.AuthHeaderTemplate
if template == "" {
template = "${API_KEY}"
}
req.Header.Set(name, strings.ReplaceAll(template, "${API_KEY}", key))
return nil
}
// mintGCPToken exchanges a base64 service-account key for an access token.
func mintGCPToken(ctx context.Context, saKeyB64 string) (string, error) {
jsonKey, err := base64.StdEncoding.DecodeString(strings.TrimSpace(saKeyB64))
if err != nil {
return "", fmt.Errorf("decode service-account key: %w", err)
}
conf, err := google.JWTConfigFromJSON(jsonKey, gcpScope)
if err != nil {
return "", fmt.Errorf("parse service-account key: %w", err)
}
tok, err := conf.TokenSource(ctx).Token()
if err != nil {
return "", fmt.Errorf("mint gcp token: %w", err)
}
return tok.AccessToken, nil
}
// decorate turns raw vendor ids into the models the caller renders, attaching
// the rates the request would actually be billed at.
//
// Rates come from the live default pricing table rather than the compiled-in
// catalog, because that is the table the synthesiser ships to the proxy: an
// operator running a defaults_llm_pricing.yaml would otherwise be shown one
// price in the form and charged another. It is also the same lookup the catalog
// endpoint prefills from, so a model reached by either route prices identically.
func decorate(entry catalog.Provider, ids []listedModel) []Model {
out := make([]Model, 0, len(ids))
seen := make(map[string]struct{}, len(ids))
for _, listed := range ids {
if listed.id == "" {
continue
}
if entry.Discovery.ExactModelsOnly && strings.Contains(listed.id, "*") {
continue
}
if _, dup := seen[listed.id]; dup {
continue
}
seen[listed.id] = struct{}{}
// The table keys pricing by the normalised id while the vendor issues
// the wire form, so normalise before looking it up — otherwise every
// Bedrock profile would report unpriced.
model := Model{ID: listed.id, Label: listed.label}
if rate, known := pricing.LookupDefault(entry.PricingSurfaces, normalizeForPricing(entry.ID, listed.id)); known {
model.PricingKnown = true
model.InputPer1k = rate.InputPer1k
model.OutputPer1k = rate.OutputPer1k
model.CachedInputPer1k = rate.CachedInputPer1k
model.CacheReadPer1k = rate.CacheReadPer1k
model.CacheCreationPer1k = rate.CacheCreationPer1k
}
out = append(out, model)
}
return out
}
// refuseRedirect is the redirect policy every discovery request runs under. A
// redirect is a way to move the request to a host checkPublicHost never saw,
// so none are followed.
func refuseRedirect(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
}
func (c *Client) httpClient() *http.Client {
if c.HTTPClient != nil {
if c.HTTPClient.CheckRedirect != nil {
return c.HTTPClient
}
// An injected client that states no policy still gets ours: the
// no-redirect guarantee should not depend on the caller remembering it.
//
// Copied rather than assigned into: one Client is shared by every
// request for the process's lifetime, so writing to its fields here
// would race across request goroutines. The copy shares the Transport,
// which is safe for concurrent use by design.
clone := *c.HTTPClient
clone.CheckRedirect = refuseRedirect
return &clone
}
transport := guardedTransport
if c.AllowPrivateHosts {
transport = http.DefaultTransport
}
return &http.Client{
Timeout: fetchTimeout,
Transport: transport,
CheckRedirect: refuseRedirect,
}
}
// guardedTransport dials only addresses isPublic accepts.
//
// checkPublicHost resolves the host itself, and the transport then resolves it
// again when it dials — two lookups of a name whose owner chooses the answers.
// A record that returns a public address to the first and 127.0.0.1 to the
// second passes the guard and reaches loopback anyway, which is the whole of
// DNS rebinding. Re-checking at the socket closes that window: whatever the
// second lookup returned is what Control is handed, and an address the guard
// refuses never gets connected.
//
// Shared package-wide rather than built per Fetch so connections and their
// pool survive between calls; the guard holds no state.
var guardedTransport = newGuardedTransport()
func newGuardedTransport() http.RoundTripper {
base, ok := http.DefaultTransport.(*http.Transport)
if !ok {
// Something replaced the default transport. Fall back to it rather
// than dropping its behaviour, and rely on checkPublicHost alone.
return http.DefaultTransport
}
// Cloned so proxy settings, TLS defaults and timeouts come from the
// standard transport rather than being restated here.
transport := base.Clone()
dialer := &net.Dialer{
Timeout: fetchTimeout,
KeepAlive: 30 * time.Second,
Control: func(_, address string, _ syscall.RawConn) error {
return guardDialAddress(address)
},
}
transport.DialContext = dialer.DialContext
return transport
}
// guardDialAddress refuses a resolved socket address the discovery client has
// no business connecting to. Control hands it over post-resolution and
// pre-connect, once per address the dialer tries, so a name with several A
// records is checked at each one.
func guardDialAddress(address string) error {
host, _, err := net.SplitHostPort(address)
if err != nil {
return fmt.Errorf("discovery dial address %q is unreadable", address)
}
addr, err := netip.ParseAddr(host)
if err != nil {
// Control is documented to receive a resolved address; anything else
// is a state we cannot vet, so it does not get dialled.
return fmt.Errorf("discovery dial address %q is not an IP", host)
}
if !isPublic(addr) {
// Deliberately not ErrPrivateHost, which means "this upstream is on a
// private network, so we cannot check it" and lets a save through
// unchecked. checkPublicHost has already cleared the target by the
// time anything is dialled, so an address refused here is not the
// operator's upstream: it is a rebinding attempt, or an HTTP proxy in
// the path. Neither may quietly skip the check — one is hostile, and
// the other would silently disable this on every provider.
return fmt.Errorf("discovery refused to dial non-public address %s", addr)
}
return nil
}
@@ -0,0 +1,678 @@
package modeldiscovery
import (
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
)
// stubTransport answers every request with one canned response and records the
// request it was given, so a test can assert on the URL and headers the client
// built without a network round trip.
type stubTransport struct {
status int
body string
got *http.Request
}
func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) {
s.got = req
status := s.status
if status == 0 {
status = http.StatusOK
}
return &http.Response{
StatusCode: status,
Body: io.NopCloser(strings.NewReader(s.body)),
Header: http.Header{"Content-Type": []string{"application/json"}},
Request: req,
}, nil
}
// newStubClient returns a client that never leaves the process. The host guard
// is disabled because it would otherwise resolve the vendor's real name, which
// would make these tests depend on DNS.
func newStubClient(status int, body string) (*Client, *stubTransport) {
tr := &stubTransport{status: status, body: body}
return &Client{
HTTPClient: &http.Client{Transport: tr},
AllowPrivateHosts: true,
}, tr
}
// The payloads below are trimmed from what the vendors actually returned in
// the discovery e2e, rather than invented, so a parser that only works against
// an idealised shape fails here.
const openAIListing = `{"object":"list","data":[
{"id":"gpt-4o-mini","object":"model","created":1721172741,"owned_by":"system"},
{"id":"gpt-4o","object":"model","created":1715367049,"owned_by":"system"}
]}`
const agentgatewayListing = `{"object":"list","data":[
{"id":"gpt-4o-mini","object":"model","created":1785166485,"owned_by":"openai"},
{"id":"claude-haiku-4-5","object":"model","created":1785166485,"owned_by":"anthropic"},
{"id":"openai/*","object":"model","created":1785166485,"owned_by":"openai"},
{"id":"*-latest","object":"model","created":1785166485,"owned_by":"openai"}
]}`
const anthropicListing = `{"data":[
{"type":"model","id":"claude-haiku-4-5-20251001","display_name":"Claude Haiku 4.5"},
{"type":"model","id":"claude-sonnet-4-6","display_name":"Claude Sonnet 4.6"}
],"has_more":false}`
const bedrockListing = `{"inferenceProfileSummaries":[
{"inferenceProfileId":"eu.anthropic.claude-haiku-4-5-20251001-v1:0",
"inferenceProfileName":"EU Anthropic Claude Haiku 4.5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
{"inferenceProfileId":"global.cohere.embed-v4:0",
"inferenceProfileName":"Global Cohere Embed v4","status":"ACTIVE","type":"SYSTEM_DEFINED"},
{"inferenceProfileId":"eu.meta.llama3-2-1b-instruct-v1:0",
"inferenceProfileName":"EU Meta Llama 3.2 1B","status":"INACTIVE","type":"SYSTEM_DEFINED"}
]}`
const vertexListing = `{"publisherModels":[
{"name":"publishers/anthropic/models/claude-3-opus","versionId":"20240229","launchStage":"GA"},
{"name":"publishers/anthropic/models/claude-sonnet-4-5","versionId":"20250929","launchStage":"GA"}
]}`
func TestFetchOpenAIListing(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, openAIListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
APIKey: "sk-test",
})
require.NoError(t, err)
assert.Equal(t, "https://api.openai.com/v1/models", tr.got.URL.String())
assert.Equal(t, "Bearer sk-test", tr.got.Header.Get("Authorization"),
"the credential must be injected through the catalog's auth template")
assert.Equal(t, []string{"gpt-4o-mini", "gpt-4o"}, ids(models))
for _, m := range models {
assert.True(t, m.PricingKnown, "both models are in the shipped catalog: %s", m.ID)
}
}
func TestFetchAgentgatewayListing(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, agentgatewayListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "agentgateway",
UpstreamURL: "https://gateway.example.com",
APIKey: "virtual-key",
})
require.NoError(t, err)
assert.Equal(t, "https://gateway.example.com/v1/models", tr.got.URL.String())
assert.Equal(t, "Bearer virtual-key", tr.got.Header.Get("Authorization"),
"agentgateway model discovery must use the configured virtual key")
assert.Equal(t, []string{"gpt-4o-mini", "claude-haiku-4-5"}, ids(models),
"model patterns must not be offered as exact NetBird authorization rows")
for _, m := range models {
assert.True(t, m.PricingKnown, "known upstream model must use NetBird catalog pricing: %s", m.ID)
}
}
func TestFetchAnthropicSendsTheVersionHeader(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, anthropicListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "anthropic_api",
UpstreamURL: "https://api.anthropic.com",
APIKey: "sk-ant-test",
})
require.NoError(t, err)
// Anthropic rejects a request without the version header, so a listing
// that reached us at all proves it was sent — but assert it, because the
// failure mode otherwise only shows up against the live API.
assert.Equal(t, "2023-06-01", tr.got.Header.Get("anthropic-version"))
assert.Equal(t, "sk-ant-test", tr.got.Header.Get("x-api-key"),
"Anthropic takes a bare key under its own header, not a Bearer token")
assert.Equal(t, "limit=1000", tr.got.URL.RawQuery)
assert.Equal(t, []string{"claude-haiku-4-5-20251001", "claude-sonnet-4-6"}, ids(models))
assert.Equal(t, "Claude Haiku 4.5", models[0].Label)
}
func TestFetchBedrockUsesTheControlPlaneAndKeepsWireIDs(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, bedrockListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
// The record's upstream is the RUNTIME host, which does not serve
// listings. The catalog's own discovery host must win over it.
UpstreamURL: "https://bedrock-runtime.eu-central-1.amazonaws.com",
Region: "eu-central-1",
APIKey: "aws-bearer",
})
require.NoError(t, err)
assert.Equal(t, "https://bedrock.eu-central-1.amazonaws.com/inference-profiles",
tr.got.URL.String(), "listings come from the control plane, not the runtime host")
// Region-prefixed ids verbatim: the prefix is what makes them invocable
// and it cannot be reconstructed — global.* alongside eu.* is exactly the
// case that defeats deriving it from the configured region.
assert.Equal(t, []string{
"eu.anthropic.claude-haiku-4-5-20251001-v1:0",
"global.cohere.embed-v4:0",
}, ids(models), "an INACTIVE profile must not be offered")
assert.True(t, models[0].PricingKnown,
"the catalog prices anthropic.claude-haiku-4-5, which this id normalises to")
assert.False(t, models[1].PricingKnown,
"cohere embed is not in the shipped Bedrock catalog, so the operator must price it")
// The rates travel with the model, so the form can prefill an editable row
// rather than making the operator look every price up by hand.
assert.Positive(t, models[0].InputPer1k, "a priced model must carry its input rate")
assert.Positive(t, models[0].OutputPer1k, "a priced model must carry its output rate")
// An unpriced model is offered at zero and flagged, not withheld: the
// vendor says the credential can reach it.
assert.Zero(t, models[1].InputPer1k)
assert.Zero(t, models[1].OutputPer1k)
}
// TestDiscoveredRatesMatchTheCatalogEndpoint pins the two prefill paths to one
// table. The provider form fills a model row either from the catalog response
// or from a discovery response, and an operator who switches between them must
// not see the price change — both must equal what the proxy will bill.
func TestDiscoveredRatesMatchTheCatalogEndpoint(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, openAIListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
APIKey: "sk-test",
})
require.NoError(t, err)
require.NotEmpty(t, models)
entry, ok := catalog.Lookup("openai_api")
require.True(t, ok)
for _, m := range models {
want, known := pricing.LookupDefault(entry.PricingSurfaces, m.ID)
require.True(t, known, "%s should be priced by the default table", m.ID)
assert.Equal(t, want.InputPer1k, m.InputPer1k, "input rate for %s", m.ID)
assert.Equal(t, want.OutputPer1k, m.OutputPer1k, "output rate for %s", m.ID)
assert.Equal(t, want.CachedInputPer1k, m.CachedInputPer1k, "cached-input rate for %s", m.ID)
}
}
func TestFetchVertexJoinsNameAndVersion(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, vertexListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "vertex_ai_api",
UpstreamURL: "https://us-east5-aiplatform.googleapis.com",
Region: "us-east5",
APIKey: "ya29.test-token",
})
require.NoError(t, err)
// Vertex addresses a model as "<id>@<version>" on rawPredict, and splits
// those across two fields in the listing.
assert.Equal(t, []string{"claude-3-opus@20240229", "claude-sonnet-4-5@20250929"}, ids(models))
assert.Equal(t, "claude-3-opus", models[0].Label)
}
func TestFetchSurfacesTheVendorStatus(t *testing.T) {
cl, _ := newStubClient(http.StatusForbidden, `{"error":{"message":"no access"}}`)
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
APIKey: "sk-test",
})
require.Error(t, err)
assert.Contains(t, err.Error(), "403",
"an operator whose key lacks access needs to see which status the vendor returned")
}
func TestFetchRejectsAProviderWithoutDiscovery(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, openAIListing)
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "litellm_proxy",
UpstreamURL: "https://gateway.example.com",
APIKey: "sk-test",
})
assert.ErrorIs(t, err, ErrNoDiscovery,
"a gateway with no listing endpoint must be distinguishable from a failure, so the caller can fall back")
}
func TestFetchRequiresACredential(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, openAIListing)
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
})
require.Error(t, err)
assert.Contains(t, err.Error(), "API key")
}
func TestDiscoveryURLNeedsARegionWhenTheHostTemplatesOne(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, bedrockListing)
// An upstream that matches no catalog template — a proxy in front of
// Bedrock, say — leaves nothing to read the region from. Refusing beats
// guessing: an unsubstituted placeholder would dial a host that does not
// exist, and a guessed region would dial the wrong account's endpoint.
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
UpstreamURL: "https://bedrock.internal-proxy.example.com",
APIKey: "aws-bearer",
})
require.Error(t, err)
assert.Contains(t, err.Error(), "region")
}
// TestHostGuardRejectsNonPublicAddresses is the SSRF guard. Management holds a
// credential for every provider, so an upstream pointed at an internal address
// would turn discovery into a way to probe — and hand a token to — the
// management server's own network.
func TestHostGuardRejectsNonPublicAddresses(t *testing.T) {
for _, tc := range []struct {
name string
addr string
want bool
}{
{"loopback v4", "127.0.0.1", false},
{"loopback v6", "::1", false},
{"private 10/8", "10.0.0.5", false},
{"private 172.16/12", "172.16.4.1", false},
{"private 192.168/16", "192.168.1.1", false},
{"link-local", "169.254.169.254", false}, // cloud metadata
{"unspecified", "0.0.0.0", false},
{"multicast", "224.0.0.1", false},
{"netbird overlay 100.64/10", "100.90.1.2", false},
{"v4-mapped loopback", "::ffff:127.0.0.1", false},
{"public v4", "1.1.1.1", true},
{"public v6", "2606:4700:4700::1111", true},
{"just outside CGNAT", "100.128.0.1", true},
} {
t.Run(tc.name, func(t *testing.T) {
addr, err := netip.ParseAddr(tc.addr)
require.NoError(t, err)
assert.Equal(t, tc.want, isPublic(addr))
})
}
}
func TestHostGuardResolvesAndRejectsLocalhost(t *testing.T) {
cl := &Client{}
err := cl.checkPublicHost(context.Background(), "localhost")
require.Error(t, err, "a name resolving to loopback must be refused, not just a literal address")
assert.Contains(t, err.Error(), "non-public")
}
// TestRedirectsAreNotFollowed covers a gap the other tests leave open: they all
// inject an HTTPClient, which bypasses httpClient() and therefore the redirect
// policy entirely. The policy is a security control — a 302 moves the request
// to a host checkPublicHost never resolved — so it needs a test that goes
// through the constructor the manager actually uses.
func TestRedirectsAreNotFollowed(t *testing.T) {
var hits int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hits++
http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound)
}))
t.Cleanup(srv.Close)
for name, cl := range map[string]*Client{
// The production shape: no injected client at all.
"default client": {AllowPrivateHosts: true},
// An injected client that states no policy must inherit ours rather
// than silently chasing the redirect.
"injected client with no policy": {
AllowPrivateHosts: true,
HTTPClient: &http.Client{},
},
} {
t.Run(name, func(t *testing.T) {
hits = 0
req, err := http.NewRequest(http.MethodGet, srv.URL, nil)
require.NoError(t, err)
resp, err := cl.httpClient().Do(req)
require.NoError(t, err)
t.Cleanup(func() { _ = resp.Body.Close() })
assert.Equal(t, http.StatusFound, resp.StatusCode,
"the redirect must be surfaced, not followed to an unchecked host")
assert.Equal(t, 1, hits, "exactly one request must leave the client")
})
}
}
// TestInjectedClientKeepsItsOwnRedirectPolicy pins that the default above is a
// default, not an override, and that supplying it does not mutate the caller's
// client — one Client is shared across every request, so a write here would
// race.
func TestInjectedClientKeepsItsOwnRedirectPolicy(t *testing.T) {
own := func(*http.Request, []*http.Request) error { return nil }
injected := &http.Client{CheckRedirect: own}
cl := &Client{HTTPClient: injected}
assert.Same(t, injected, cl.httpClient(),
"a client that states a policy must be handed back untouched")
bare := &http.Client{}
cl = &Client{HTTPClient: bare}
require.NotSame(t, bare, cl.httpClient(), "the policy must be applied to a copy")
assert.Nil(t, bare.CheckRedirect, "the caller's client must not be written to")
}
// TestDialGuardRejectsRebindingToANonPublicAddress covers the window between
// the two DNS lookups. checkPublicHost resolves the host, then the transport
// resolves it again to dial; a name whose owner answers the first with a public
// address and the second with 127.0.0.1 would otherwise pass the guard and
// still reach loopback. The dial-time check sees whatever the second lookup
// actually returned.
func TestDialGuardRejectsRebindingToANonPublicAddress(t *testing.T) {
for _, tc := range []struct {
name string
address string
wantErr string
}{
{"loopback", "127.0.0.1:443", "non-public"},
{"cloud metadata", "169.254.169.254:80", "non-public"},
{"rfc1918", "10.1.2.3:443", "non-public"},
{"netbird overlay", "100.90.1.2:443", "non-public"},
{"loopback v6", "[::1]:443", "non-public"},
{"unresolved name", "evil.example.com:443", "not an IP"},
{"no port", "1.1.1.1", "unreadable"},
} {
t.Run(tc.name, func(t *testing.T) {
err := guardDialAddress(tc.address)
require.Error(t, err)
assert.Contains(t, err.Error(), tc.wantErr)
})
}
assert.NoError(t, guardDialAddress("1.1.1.1:443"), "a public address must still be dialled")
assert.NoError(t, guardDialAddress("[2606:4700:4700::1111]:443"))
}
// TestDialGuardIsInstalledOnTheDefaultClient pins the wiring rather than the
// guard: a correct guard nothing calls protects nothing.
func TestDialGuardIsInstalledOnTheDefaultClient(t *testing.T) {
cl := &Client{}
transport, ok := cl.httpClient().Transport.(*http.Transport)
require.True(t, ok, "the default discovery client must carry the guarded transport")
require.NotNil(t, transport.DialContext, "the guarded transport must dial through the guard")
_, err := transport.DialContext(context.Background(), "tcp", "127.0.0.1:9")
require.Error(t, err, "the guard must refuse loopback even when the caller dials it directly")
assert.Contains(t, err.Error(), "non-public")
// Tests point the client at a loopback server on purpose, so the opt-out
// has to reach the dialer too.
relaxed := &Client{AllowPrivateHosts: true}
assert.Equal(t, http.DefaultTransport, relaxed.httpClient().Transport)
}
// TestCallerInputFailuresAreMarkedInvalid keeps the handler's 400 mapping
// honest: it branches on this sentinel, so an unmarked caller-input failure
// silently becomes a 500.
func TestCallerInputFailuresAreMarkedInvalid(t *testing.T) {
for _, tc := range []struct {
name string
req Request
}{
{"unknown provider", Request{CatalogID: "not_a_provider", APIKey: "k"}},
{"unusable upstream", Request{CatalogID: "openai_api", UpstreamURL: "://", APIKey: "k"}},
{"missing api key", Request{CatalogID: "openai_api", UpstreamURL: "https://api.openai.com"}},
{"no region to read", Request{
CatalogID: "bedrock_api",
UpstreamURL: "https://bedrock-runtime.amazonaws.com",
APIKey: "aws-bearer",
}},
} {
t.Run(tc.name, func(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, openAIListing)
_, err := cl.Fetch(context.Background(), tc.req)
require.Error(t, err)
assert.ErrorIs(t, err, ErrInvalidRequest)
})
}
}
// TestEveryDiscoveryEntryHasAParser keeps the catalog and the parser table from
// drifting: adding a Discovery block with a shape nothing parses would fail
// only at runtime, in front of an operator.
func TestEveryDiscoveryEntryHasAParser(t *testing.T) {
for _, entry := range catalog.All() {
if entry.Discovery == nil {
continue
}
t.Run(entry.ID, func(t *testing.T) {
assert.NotEmpty(t, entry.Discovery.Path, "a discovery entry needs a path")
_, err := parseListing(entry.Discovery.Shape, []byte(`{}`))
assert.NoError(t, err, "shape %q has no parser", entry.Discovery.Shape)
})
}
}
func ids(models []Model) []string {
out := make([]string, 0, len(models))
for _, m := range models {
out = append(out, m.ID)
}
return out
}
// TestRegionIsReadBackFromTheUpstream covers the reason the API takes no
// region field: a provider record has none, and the operator already encoded
// it in the upstream host when they configured inference.
func TestRegionIsReadBackFromTheUpstream(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, bedrockListing)
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
UpstreamURL: "https://bedrock-runtime.us-west-2.amazonaws.com",
APIKey: "aws-bearer",
})
require.NoError(t, err)
assert.Equal(t, "bedrock.us-west-2.amazonaws.com", tr.got.URL.Host)
}
func TestRegionFromUpstream(t *testing.T) {
bedrock, ok := catalog.Lookup("bedrock_api")
require.True(t, ok)
vertex, ok := catalog.Lookup("vertex_ai_api")
require.True(t, ok)
for _, tc := range []struct {
name string
entry catalog.Provider
upstream string
want string
}{
{"bedrock runtime host", bedrock, "https://bedrock-runtime.eu-central-1.amazonaws.com", "eu-central-1"},
{"bedrock without scheme", bedrock, "bedrock-runtime.ap-south-1.amazonaws.com", "ap-south-1"},
{"vertex regional host", vertex, "https://us-east5-aiplatform.googleapis.com", "us-east5"},
// A proxied or self-hosted upstream matches no template, and guessing
// a region from it would build a URL pointing somewhere arbitrary.
{"unrelated upstream", bedrock, "https://llm.internal.example.com", ""},
{"vertex global host has no region segment", vertex, "https://aiplatform.googleapis.com", ""},
// Bedrock's regionless endpoint carries both halves of the template at
// once, with nothing between them. It has to read as "no region here"
// rather than as an inverted slice range.
{"bedrock regionless endpoint", bedrock, "https://bedrock-runtime.amazonaws.com", ""},
{"bedrock regionless without scheme", bedrock, "bedrock-runtime.amazonaws.com", ""},
} {
t.Run(tc.name, func(t *testing.T) {
assert.Equal(t, tc.want, RegionFromUpstream(tc.entry, tc.upstream))
})
}
}
// bedrockGeoListing carries profiles from geographies the original prefix list
// did not name. Every one reduces to a catalog key, so every one must arrive
// priced — an unstripped geography is what made a real account's listing come
// back almost entirely at zero.
const bedrockGeoListing = `{"inferenceProfileSummaries":[
{"inferenceProfileId":"jp.anthropic.claude-sonnet-5-20260514-v1:0",
"inferenceProfileName":"JP Anthropic Claude Sonnet 5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
{"inferenceProfileId":"au.anthropic.claude-haiku-4-5-20251001-v1:0",
"inferenceProfileName":"AU Anthropic Claude Haiku 4.5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
{"inferenceProfileId":"us-gov.anthropic.claude-sonnet-5-20260514-v1:0",
"inferenceProfileName":"GovCloud Anthropic Claude Sonnet 5","status":"ACTIVE","type":"SYSTEM_DEFINED"}
]}`
func TestBedrockProfilesFromAnyGeographyArrivePriced(t *testing.T) {
cl, _ := newStubClient(http.StatusOK, bedrockGeoListing)
models, err := cl.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
UpstreamURL: "https://bedrock-runtime.eu-central-1.amazonaws.com",
APIKey: "aws-token",
})
require.NoError(t, err)
require.Len(t, models, 3)
for _, m := range models {
assert.True(t, m.PricingKnown, "%s must resolve to a catalog rate", m.ID)
assert.Greater(t, m.InputPer1k, 0.0, "input rate for %s", m.ID)
assert.Greater(t, m.OutputPer1k, 0.0, "output rate for %s", m.ID)
assert.Greater(t, m.CacheReadPer1k, 0.0, "cache-read rate for %s", m.ID)
}
// The wire id is preserved whatever the pricing key reduced to: it is the
// only form that works at invoke time.
assert.Equal(t, "jp.anthropic.claude-sonnet-5-20260514-v1:0", models[0].ID)
}
// TestFetch_AHostThatWillNotResolveIsUnreachable closes a gap the live suite
// found. The SSRF guard resolves the host before any request is built, so a
// name that does not resolve fails there rather than at the transport — and
// that error used to reach the caller unclassified. A wrong hostname is the
// commonest way for an upstream to be wrong, so it has to arrive as
// "unreachable" and not as an unrecognised fault.
func TestFetch_AHostThatWillNotResolveIsUnreachable(t *testing.T) {
// A resolver whose dial always fails, so the lookup errors without the
// test depending on real DNS.
refusing := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
return nil, errors.New("resolver unavailable")
},
}
client := &Client{Resolver: refusing}
_, err := client.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://not-a-real-vendor-host.example.invalid",
APIKey: "sk-test",
})
require.Error(t, err)
var unreachable *UnreachableError
require.ErrorAs(t, err, &unreachable, "a host that will not resolve must classify as unreachable")
require.NotErrorIs(t, err, ErrPrivateHost, "it is not a host we declined to dial")
}
// TestFetch_AProxyInThePathDoesNotSilentlyDisableTheCheck pins a fail-open the
// dial-time guard can produce. checkPublicHost clears the target before
// anything is dialled, so a private address refused at the socket is never the
// operator's upstream — it is a rebinding attempt, or an HTTP proxy the
// management server egresses through. Reporting either as ErrPrivateHost would
// read as "this provider cannot be checked" and let every save through
// unchecked, which is how a proxied deployment would install this feature and
// have it quietly do nothing.
func TestFetch_AProxyInThePathDoesNotSilentlyDisableTheCheck(t *testing.T) {
// A transport that refuses at the socket exactly as the guard does, with a
// loopback address standing in for the proxy the dial went to.
// AllowPrivateHosts short-circuits the resolve-stage check only; the
// injected transport below is still what the request goes through. Without
// it this test resolves api.openai.com for real, and on a runner with no
// egress that lookup fails as an UnreachableError too — so it would pass
// while never reaching the socket guard it is named for.
client := &Client{AllowPrivateHosts: true, HTTPClient: &http.Client{
Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return nil, guardDialAddress("127.0.0.1:38599")
}),
CheckRedirect: refuseRedirect,
}}
_, err := client.Fetch(context.Background(), Request{
CatalogID: "openai_api",
UpstreamURL: "https://api.openai.com",
APIKey: "sk-test",
})
require.Error(t, err)
require.NotErrorIs(t, err, ErrPrivateHost,
"a refusal at the socket must not read as an upstream we cannot check")
var unreachable *UnreachableError
require.ErrorAs(t, err, &unreachable, "it is the vendor we failed to reach")
}
// roundTripFunc adapts a function to http.RoundTripper.
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
// TestFetch_TheUpstreamIsCheckedWhenTheListingCannotVouchForIt covers the hole
// a separate listing host leaves. Bedrock lists from the control plane, so a
// record whose runtime upstream does not exist reaches a perfectly good
// listing and saves — the requests it then serves go nowhere.
//
// Both halves matter. A runtime host that cannot be resolved is the record
// being wrong, and blocks. A proxied one resolves and only leaves the region
// underivable, which stays the unverifiable outcome it already was.
func TestFetch_TheUpstreamIsCheckedWhenTheListingCannotVouchForIt(t *testing.T) {
refusing := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
return nil, errors.New("resolver unavailable")
},
}
client := &Client{Resolver: refusing}
_, err := client.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
// Matches no catalog template, so nothing here reaches the control
// plane the listing comes from: without its own check this upstream
// was never contacted at all.
UpstreamURL: "https://bedrock.typo.example.invalid",
APIKey: "aws-bearer",
})
require.Error(t, err)
var unreachable *UnreachableError
require.ErrorAs(t, err, &unreachable, "a runtime host that will not resolve must block the save")
}
// TestFetch_AListingHostOfItsOwnDoesNotReachThroughTheUpstream keeps the check
// above from reading the operator's upstream as the place to list from.
func TestFetch_AListingHostOfItsOwnDoesNotReachThroughTheUpstream(t *testing.T) {
cl, tr := newStubClient(http.StatusOK, bedrockListing)
_, err := cl.Fetch(context.Background(), Request{
CatalogID: "bedrock_api",
UpstreamURL: "https://bedrock-runtime.eu-central-1.amazonaws.com",
APIKey: "aws-bearer",
})
require.NoError(t, err)
assert.Equal(t, "bedrock.eu-central-1.amazonaws.com", tr.got.URL.Host,
"checking the runtime host must not turn it into the listing host")
}
@@ -0,0 +1,114 @@
package modeldiscovery
import (
"context"
"crypto/tls"
"errors"
"fmt"
"net"
"os"
"syscall"
)
// Fetch serves two callers with different needs: the model picker, which only
// needs to know it failed, and the provider credential check, which has to
// tell an operator whether the URL or the key is at fault. Each failure
// carries a type so the second does not have to branch on a message.
// VendorStatusError reports a listing answered with something other than 200.
// Only the vendor's own code separates a refused credential (401, 403) from a
// URL that does not serve this API (404, 405) from an unwell vendor (5xx).
type VendorStatusError struct {
Provider string
Status int
}
func (e *VendorStatusError) Error() string {
return fmt.Sprintf("%s returned %d for its model listing", e.Provider, e.Status)
}
// UnreachableError reports that the request never reached the vendor: the
// name did not resolve, the connection was refused, TLS failed, or it timed
// out. Nothing was authenticated, so only the URL is implicated.
type UnreachableError struct {
Provider string
Err error
}
func (e *UnreachableError) Error() string {
return fmt.Sprintf("reach %s: %v", e.Provider, e.Err)
}
func (e *UnreachableError) Unwrap() error { return e.Err }
// Reason names the transport failure in words an operator can act on: a wrong
// port and a wrong hostname fail differently and are worth telling apart.
// Empty means unrecognised, and the caller should say only that the host could
// not be reached rather than paste a Go error into the UI.
func (e *UnreachableError) Reason() string {
err := e.Err
var dns *net.DNSError
if errors.As(err, &dns) {
if dns.IsNotFound {
return "no such host"
}
// Named apart from the dial timeout below. A resolver that never
// answered and an upstream that never answered send an operator to
// different places, and the generic "connection timed out" would
// describe a connection that was never attempted.
if dns.IsTimeout {
return "dns lookup timed out"
}
return "dns lookup failed"
}
// Timeouts are checked before the syscall cases: a dial that times out is
// reported as a net.OpError wrapping a timeout, and the operator needs to
// hear "timed out" rather than the syscall underneath it.
if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, os.ErrDeadlineExceeded) {
return "connection timed out"
}
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
return "connection timed out"
}
if errors.Is(err, syscall.ECONNREFUSED) {
return "connection refused"
}
if errors.Is(err, syscall.EHOSTUNREACH) || errors.Is(err, syscall.ENETUNREACH) {
return "host unreachable"
}
var certErr *tls.CertificateVerificationError
if errors.As(err, &certErr) {
return "tls certificate not trusted"
}
var recordErr tls.RecordHeaderError
if errors.As(err, &recordErr) {
return "not a tls endpoint"
}
return ""
}
// ErrUnparseableListing marks a 200 whose body is not a listing in the shape
// the catalog declared. Distinct from a status refusal: the host answered and
// authenticated fine, it is just not the API — a login page, say.
var ErrUnparseableListing = errors.New("response is not a model listing")
// ErrNoDiscoveryHost marks a provider whose listing host cannot be derived
// from the record: Bedrock's control-plane host comes from the region in the
// upstream, so a proxied endpoint leaves nowhere to send it, and inventing one
// would spend the credential somewhere never configured.
//
// Wraps ErrInvalidRequest so the discovery endpoint still answers 400, while a
// credential check can read it as "cannot be checked" rather than "broken".
var ErrNoDiscoveryHost = errors.New("provider has no derivable discovery host")
// ErrPrivateHost marks an upstream resolving somewhere management will not
// dial. A self-hosted endpoint on a private network is a legitimate provider
// the proxy reaches through the tunnel, so this means the check cannot run,
// not that the record is wrong.
var ErrPrivateHost = errors.New("discovery host is not publicly routable")
@@ -0,0 +1,134 @@
package modeldiscovery
import (
"encoding/json"
"fmt"
"strings"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
sharedllm "github.com/netbirdio/netbird/shared/llm"
)
// listedModel is one entry lifted out of a vendor listing before the catalog
// is consulted about it.
type listedModel struct {
id string
label string
}
// parseListing extracts model ids from a vendor listing. Each vendor invented
// its own envelope, and the shape is declared by the catalog rather than
// sniffed, so a vendor that changes shape fails loudly instead of silently
// returning nothing.
func parseListing(shape catalog.ListingShape, body []byte) ([]listedModel, error) {
switch shape {
case catalog.ShapeOpenAIData:
return parseOpenAIData(body)
case catalog.ShapeBedrockInferenceProfiles:
return parseBedrockInferenceProfiles(body)
case catalog.ShapeVertexPublisherModels:
return parseVertexPublisherModels(body)
default:
return nil, fmt.Errorf("no parser for listing shape %q", shape)
}
}
// parseOpenAIData reads {"data":[{"id":…}]}, which OpenAI defined and
// Anthropic adopted. Anthropic additionally supplies display_name.
func parseOpenAIData(body []byte) ([]listedModel, error) {
var doc struct {
Data []struct {
ID string `json:"id"`
DisplayName string `json:"display_name"`
} `json:"data"`
}
if err := json.Unmarshal(body, &doc); err != nil {
return nil, fmt.Errorf("%w: decode model listing: %w", ErrUnparseableListing, err)
}
out := make([]listedModel, 0, len(doc.Data))
for _, entry := range doc.Data {
out = append(out, listedModel{id: entry.ID, label: entry.DisplayName})
}
return out, nil
}
// parseBedrockInferenceProfiles reads
// {"inferenceProfileSummaries":[{"inferenceProfileId":…}]}.
//
// The profile id is taken verbatim because its region prefix (eu., us.,
// global.) is what makes it invocable, and it is not derivable from the
// configured region — an account in one region legitimately holds global.*
// profiles alongside its regional ones.
//
// Only ACTIVE profiles are offered: AWS reports others, and registering one
// would produce a model that routes inside NetBird and fails at AWS.
func parseBedrockInferenceProfiles(body []byte) ([]listedModel, error) {
var doc struct {
Summaries []struct {
ID string `json:"inferenceProfileId"`
Name string `json:"inferenceProfileName"`
Status string `json:"status"`
} `json:"inferenceProfileSummaries"`
}
if err := json.Unmarshal(body, &doc); err != nil {
return nil, fmt.Errorf("%w: decode inference-profile listing: %w", ErrUnparseableListing, err)
}
out := make([]listedModel, 0, len(doc.Summaries))
for _, entry := range doc.Summaries {
if entry.Status != "" && !strings.EqualFold(entry.Status, "ACTIVE") {
continue
}
out = append(out, listedModel{id: entry.ID, label: entry.Name})
}
return out, nil
}
// parseVertexPublisherModels reads {"publisherModels":[{"name":…}]}, where
// name is a resource path ("publishers/anthropic/models/claude-3-opus") and
// the version lives in a separate field.
//
// Vertex addresses a model as "<id>@<version>" on the rawPredict path, so the
// two are joined here: reporting the bare name would hand the operator an id
// that looks usable and is not.
func parseVertexPublisherModels(body []byte) ([]listedModel, error) {
var doc struct {
Models []struct {
Name string `json:"name"`
VersionID string `json:"versionId"`
} `json:"publisherModels"`
}
if err := json.Unmarshal(body, &doc); err != nil {
return nil, fmt.Errorf("%w: decode publisher-model listing: %w", ErrUnparseableListing, err)
}
out := make([]listedModel, 0, len(doc.Models))
for _, entry := range doc.Models {
id := entry.Name
if slash := strings.LastIndex(id, "/"); slash >= 0 {
id = id[slash+1:]
}
if id == "" {
continue
}
label := id
if entry.VersionID != "" {
id += "@" + entry.VersionID
}
out = append(out, listedModel{id: id, label: label})
}
return out, nil
}
// normalizeForPricing maps a vendor's wire id onto the key the catalog prices
// it under. It mirrors the synthesiser's normalizePricingModelID: the two must
// agree, or a model reported here as priced would meter at the default rate
// instead of the operator's.
func normalizeForPricing(catalogProviderID, modelID string) string {
switch {
case catalog.IsBedrockPathStyle(catalogProviderID):
return sharedllm.NormalizeBedrockModel(modelID)
case catalog.IsVertexPathStyle(catalogProviderID):
return sharedllm.NormalizeVertexModel(modelID)
default:
return modelID
}
}
@@ -164,23 +164,12 @@ func (m *managerImpl) SelectPolicyForRequest(ctx context.Context, in PolicySelec
}
candidates := filterApplicablePolicies(policies, in)
// Model-allowlist gate scoped to the matched policies: keep candidates whose
// guardrails permit the model (none enabled = unrestricted), deny when
// policies apply but none permits it. Skip the load when none has a guardrail.
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
candidates, denied, err := m.applyModelGate(ctx, in, candidates)
if err != nil {
return nil, err
}
if denied != nil {
return denied, nil
}
// Prefetch every consumption counter the ceiling + candidate policies will
@@ -285,6 +274,59 @@ func anyPolicyHasGuardrails(policies []*types.Policy) bool {
return false
}
// applyModelGate is the model-allowlist gate scoped to the matched policies:
// it keeps the candidates whose guardrails permit the model (none enabled =
// unrestricted) and returns a deny result when policies apply but none
// permits it. The guardrail load is skipped when no candidate references a
// guardrail, and the provider's catalog id — which picks the model-id
// normalizer — is resolved only when a candidate actually restricts models:
// with no enabled allowlist every candidate is unrestricted, and a
// provider-store failure must not fail a request the gate would have waved
// through.
func (m *managerImpl) applyModelGate(ctx context.Context, in PolicySelectionInput, candidates []*types.Policy) ([]*types.Policy, *PolicySelectionResult, error) {
if len(candidates) == 0 || !anyPolicyHasGuardrails(candidates) {
return candidates, nil, nil
}
guardrailsByID, err := m.loadGuardrailsByID(ctx, in.AccountID)
if err != nil {
return nil, nil, err
}
if !anyEnabledModelAllowlist(candidates, guardrailsByID) {
return candidates, nil, nil
}
catalogID, err := m.providerCatalogID(ctx, in.AccountID, in.ProviderID)
if err != nil {
return nil, nil, err
}
permitted := filterModelPermittedPolicies(candidates, guardrailsByID, in.Model, catalogID)
if len(permitted) == 0 {
return nil, &PolicySelectionResult{
Allow: false,
DenyCode: denyCodeModelBlocked,
DenyReason: modelBlockedReason(in.Model),
}, nil
}
return permitted, nil, nil
}
// anyEnabledModelAllowlist reports whether any policy references a guardrail
// whose model allowlist is enabled — the only case the model gate restricts
// anything. Disabled allowlists, stale guardrail references, and guardrails
// carrying only other checks all leave every candidate unrestricted.
func anyEnabledModelAllowlist(policies []*types.Policy, byID map[string]*types.Guardrail) bool {
for _, p := range policies {
if p == nil {
continue
}
for _, gID := range p.GuardrailIDs {
if g, ok := byID[gID]; ok && g != nil && g.Checks.ModelAllowlist.Enabled {
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) {
@@ -301,12 +343,33 @@ func (m *managerImpl) loadGuardrailsByID(ctx context.Context, accountID string)
return byID, nil
}
// providerCatalogID resolves a provider record id to its catalog provider
// id, the key the model-id normalizers are picked by. A missing provider
// resolves to the empty catalog id — the compare then runs verbatim-only,
// which can never widen an allowlist — while a store failure propagates
// rather than degrading a security decision.
func (m *managerImpl) providerCatalogID(ctx context.Context, accountID, providerID string) (string, error) {
if providerID == "" {
return "", nil
}
provider, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
switch {
case err == nil:
return provider.ProviderID, nil
case isNotFound(err):
return "", nil
default:
return "", fmt.Errorf("get provider: %w", err)
}
}
// 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 {
// permit the model on the provider with the given catalog id. Order is
// preserved so downstream scoring is unaffected.
func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, model, catalogProviderID string) []*types.Policy {
out := make([]*types.Policy, 0, len(policies))
for _, p := range policies {
if policyPermitsModel(p, byID, model) {
if policyPermitsModel(p, byID, model, catalogProviderID) {
out = append(out, p)
}
}
@@ -316,8 +379,13 @@ func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*typ
// policyPermitsModel reports whether a policy permits the model. No
// allowlist-enabled guardrail = unrestricted (permits any, incl. empty);
// otherwise the model must be in the union of its allowlists, so an
// empty/undetermined model fails closed.
func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model string) bool {
// empty/undetermined model fails closed. An entry matches on its own
// normalised form or, for a path-style provider, its canonical form: the
// parser emits the canonical id for path-routed requests, while an
// allowlist may hold the raw declared id the dashboard's picker copies
// from the provider. The catalog id picks the normalizer, so a plain
// provider's entries always compare verbatim.
func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model, catalogProviderID string) bool {
if p == nil {
return false
}
@@ -333,7 +401,7 @@ func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model
continue
}
for _, allowed := range g.Checks.ModelAllowlist.Models {
if normaliseModelID(allowed) == wanted {
if normaliseModelID(allowed) == wanted || canonicalModelKey(catalogProviderID, allowed) == wanted {
return true
}
}
@@ -6,12 +6,13 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/status"
)
// guardedPolicy builds an enabled, uncapped policy that authorises sourceGroups
@@ -53,6 +54,17 @@ func expectGuardrails(mockStore *store.MockStore, account string, guardrails ...
Return(guardrails, nil)
}
// expectProviderCatalog resolves the destination provider to the given
// catalog provider id, which picks the model-id normalizer the allowlist
// gate compares through. AnyTimes: the lookup runs only when the guardrail
// gate is reached.
func expectProviderCatalog(mockStore *store.MockStore, account, providerID, catalog string) {
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), gomock.Any(), account, providerID).
Return(&types.Provider{ID: providerID, AccountID: account, ProviderID: catalog}, nil).
AnyTimes()
}
// 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.
@@ -63,6 +75,7 @@ func TestSelectPolicy_ModelBlockedByAllowlist(t *testing.T) {
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"))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
@@ -86,6 +99,7 @@ func TestSelectPolicy_ModelAllowedByAllowlist(t *testing.T) {
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"))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -109,6 +123,7 @@ func TestSelectPolicy_CaseInsensitiveModelMatch(t *testing.T) {
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 "))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -132,6 +147,7 @@ func TestSelectPolicy_UnguardedPolicyIsUnrestricted(t *testing.T) {
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"))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -159,6 +175,7 @@ func TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups(t *testing.T) {
allowlistGuardrail("g-a", "acc-1", "gpt-4o"),
allowlistGuardrail("g-b", "acc-1", "claude-opus-4"),
)
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
@@ -181,6 +198,7 @@ func TestSelectPolicy_UndeterminedModelFailsClosed(t *testing.T) {
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"))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
@@ -210,6 +228,8 @@ func TestSelectPolicy_DisabledAllowlistDoesNotRestrict(t *testing.T) {
}
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", disabled)
// Deliberately no provider expectation: with no enabled allowlist the
// gate must skip the catalog-id lookup entirely.
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -235,6 +255,7 @@ func TestSelectPolicy_UnionAcrossPolicyGuardrails(t *testing.T) {
allowlistGuardrail("g-1", "acc-1", "gpt-4o"),
allowlistGuardrail("g-2", "acc-1", "claude-opus-4"),
)
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -281,6 +302,8 @@ func TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted(t *testing.
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-missing")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1")
// Deliberately no provider expectation: an orphaned guardrail reference
// restricts nothing, so the gate must skip the catalog-id lookup.
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -314,6 +337,7 @@ func TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter(t *testing.T) {
allowlistGuardrail("g-restrict", "acc-1", "gpt-4o"),
allowlistGuardrail("g-permit", "acc-1", "claude-opus-4"),
)
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
@@ -327,3 +351,159 @@ func TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter(t *testing.T) {
assert.Equal(t, "pol-small", res.SelectedPolicyID,
"the model filter must exclude pol-big before cap scoring")
}
// TestSelectPolicy_RawDeclaredAllowlistPermitsCanonicalModel proves an
// allowlist holding the raw vendor-issued id — the form the dashboard's
// picker copies from a provider's declared models — permits the request:
// the parser emits the path-style canonical id, so the entry must match
// through the same canonicalization.
func TestSelectPolicy_RawDeclaredAllowlistPermitsCanonicalModel(t *testing.T) {
cases := []struct {
name string
catalog string
entry string
request string
}{
{"bedrock raw region/version form", "bedrock_api", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", "anthropic.claude-sonnet-4-5"},
{"vertex raw @version form", "vertex_ai_api", "claude-sonnet-4-5@20250929", "claude-sonnet-4-5"},
{"vertex raw dated @version form", "vertex_ai_api", "gpt-4o@2024-08-06", "gpt-4o"},
{"bedrock raw form with case and whitespace", "bedrock_api", " EU.Anthropic.Claude-Sonnet-4-5-20250929-V1:0 ", "anthropic.claude-sonnet-4-5"},
}
for _, tc := range cases {
t.Run(tc.name, func(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", tc.entry))
expectProviderCatalog(mockStore, "acc-1", "prov-1", tc.catalog)
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: tc.request,
})
require.NoError(t, err)
assert.True(t, res.Allow, "the raw declared allowlist entry must permit its canonical model")
assert.Equal(t, "pol-A", res.SelectedPolicyID)
})
}
// A model outside the allowlist stays denied under the same entry shape.
t.Run("unrelated canonical model stays denied", func(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", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0"))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "bedrock_api")
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "anthropic.claude-opus-4-8",
})
require.NoError(t, err)
assert.False(t, res.Allow, "a model the allowlist never names must stay denied")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
})
}
// TestSelectPolicy_PlainProviderEntriesStayVerbatim proves the canonical-form
// compare never relaxes an allowlist on a body-routed provider: its catalog
// id selects no normalizer, so a suffix that would be stripped under Bedrock
// ("-v2") or Vertex ("@...") stays part of the entry and must NOT also admit
// the stripped id — on this provider that is a different model.
func TestSelectPolicy_PlainProviderEntriesStayVerbatim(t *testing.T) {
cases := []struct {
name string
entry string
request string
}{
{"a -vN suffix is not a Bedrock version tag here", "claude-3-5-sonnet-v2", "claude-3-5-sonnet"},
{"an @word suffix is not a Vertex version tag here", "custom-model@team", "custom-model"},
{"an @digits suffix is not a Vertex version tag here", "custom-model@2024", "custom-model"},
}
for _, tc := range cases {
t.Run(tc.name, func(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", tc.entry))
expectProviderCatalog(mockStore, "acc-1", "prov-1", "openai_api")
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: tc.request,
})
require.NoError(t, err)
assert.False(t, res.Allow, "a plain provider's allowlist entry must not widen to its stripped form")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
})
}
}
// TestSelectPolicy_MissingProviderRecordComparesVerbatim proves a provider the
// store no longer holds degrades to the verbatim-only compare — the raw entry
// still matches itself, and nothing widens — rather than erroring or guessing
// a normalizer.
func TestSelectPolicy_MissingProviderRecordComparesVerbatim(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", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0"))
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), gomock.Any(), "acc-1", "prov-1").
Return(nil, status.Errorf(status.NotFound, "provider not found")).
AnyTimes()
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "anthropic.claude-sonnet-4-5",
})
require.NoError(t, err)
assert.False(t, res.Allow, "without the provider record the compare runs verbatim and must not widen")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
}
// TestSelectPolicy_ProviderLookupErrorPropagates proves a store failure while
// resolving the provider's catalog id surfaces as an error — the model gate is
// a security decision and must not silently degrade.
func TestSelectPolicy_ProviderLookupErrorPropagates(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"))
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), gomock.Any(), "acc-1", "prov-1").
Return(nil, errors.New("store unavailable"))
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "gpt-4o",
})
require.Error(t, err, "a provider-lookup failure must surface as an error")
assert.Nil(t, res)
}
@@ -6,7 +6,7 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -47,17 +47,11 @@ var supplementalDefaults = map[string]map[string]Entry{
"gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005},
},
"anthropic": {
// claude-opus-5 is not yet in the catalog lineup but gateway /
// grandfathered traffic uses it; priced so it isn't skipped.
"claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
// "kimi-k3[1m]" is the 1M-context alias some Claude Code guides
// configure against Moonshot's Anthropic-compatible endpoint;
// priced identically to kimi-k3 so those requests aren't skipped.
"kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003},
},
"bedrock": {
"anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
},
}
var (
@@ -82,6 +82,11 @@ anthropic:
output_per_1k: 0.015
cache_read_per_1k: 0.0003
cache_creation_per_1k: 0.00375
claude-sonnet-5:
input_per_1k: 0.003
output_per_1k: 0.015
cache_read_per_1k: 0.0003
cache_creation_per_1k: 0.00375
kimi-k3:
input_per_1k: 0.003
output_per_1k: 0.015
@@ -145,6 +150,11 @@ bedrock:
output_per_1k: 0.015
cache_read_per_1k: 0.0003
cache_creation_per_1k: 0.00375
anthropic.claude-sonnet-5:
input_per_1k: 0.003
output_per_1k: 0.015
cache_read_per_1k: 0.0003
cache_creation_per_1k: 0.00375
meta.llama3-3-70b-instruct:
input_per_1k: 0.00072
output_per_1k: 0.00072
@@ -116,11 +116,13 @@ func TestDefaultTable_PinnedRates(t *testing.T) {
assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input")
assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation")
// Supplementals present on their surfaces.
// Every id below must stay priced whichever source provides it: the
// catalog lineup for the current Claude 5 family, supplementalDefaults
// for the ids the dashboard deliberately doesn't offer.
for surface, ids := range map[string][]string{
"openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"},
"anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"},
"bedrock": {"anthropic.claude-opus-5"},
"anthropic": {"claude-opus-5", "claude-sonnet-5", "kimi-k3[1m]", "kimi-k3"},
"bedrock": {"anthropic.claude-opus-5", "anthropic.claude-sonnet-5"},
} {
for _, id := range ids {
_, ok := table[surface][id]
@@ -0,0 +1,289 @@
package agentnetwork
import (
"context"
"testing"
"time"
"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/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
nbtypes "github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
// These tests pin the provider read surface per grant: a caller holding
// providers read together with update (managers) gets the full record,
// while read-only viewers (usage_viewer) get the display surface only —
// connection configuration is redacted before it reaches the wire layer.
func TestGetAllProviders_RedactsConnectionConfigForReadOnlyViewer(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
saved := newSynthTestProvider()
saved.ExtraValues = map[string]string{"x-portkey-config": "cfg-123"}
saved.IdentityHeaderUserID = "X-User"
saved.IdentityHeaderGroups = "X-Groups"
saved.SkipTLSVerification = true
require.NoError(t, f.store.SaveAgentNetworkProvider(ctx, saved))
f.expectPermission(testAccountID, "viewer", modules.AgentNetworkProviders, operations.Read, true)
f.expectPermission(testAccountID, "viewer", modules.AgentNetworkProviders, operations.Update, false)
providers, err := f.manager.GetAllProviders(ctx, testAccountID, "viewer")
require.NoError(t, err)
require.Len(t, providers, 1)
p := providers[0]
assert.Equal(t, saved.ID, p.ID, "identity survives redaction")
assert.Equal(t, saved.Name, p.Name)
assert.Equal(t, saved.ProviderID, p.ProviderID)
assert.Equal(t, saved.Models, p.Models, "the model list backs the usage filters and stays")
assert.True(t, p.Enabled)
assert.Empty(t, p.UpstreamURL, "upstream URL is connection config")
assert.Empty(t, p.ExtraValues, "operator-typed header values are connection config")
assert.Empty(t, p.IdentityHeaderUserID)
assert.Empty(t, p.IdentityHeaderGroups)
assert.False(t, p.SkipTLSVerification)
assert.Empty(t, p.APIKey)
assert.Empty(t, p.SessionPrivateKey)
stored, err := f.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, testAccountID, saved.ID)
require.NoError(t, err)
assert.NotEmpty(t, stored.UpstreamURL, "redaction must not write back to the store")
}
func TestGetProvider_FullConfigForManagingCaller(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
saved := newSynthTestProvider()
saved.ExtraValues = map[string]string{"x-portkey-config": "cfg-123"}
require.NoError(t, f.store.SaveAgentNetworkProvider(ctx, saved))
f.expectPermission(testAccountID, "admin", modules.AgentNetworkProviders, operations.Read, true)
f.expectPermission(testAccountID, "admin", modules.AgentNetworkProviders, operations.Update, true)
p, err := f.manager.GetProvider(ctx, testAccountID, "admin", saved.ID)
require.NoError(t, err)
assert.Equal(t, saved.UpstreamURL, p.UpstreamURL, "a caller who can edit the provider sees its config")
assert.Equal(t, saved.ExtraValues, p.ExtraValues)
}
func TestGetProvider_RedactsForReadOnlyViewer(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
saved := newSynthTestProvider()
require.NoError(t, f.store.SaveAgentNetworkProvider(ctx, saved))
f.expectPermission(testAccountID, "viewer", modules.AgentNetworkProviders, operations.Read, true)
f.expectPermission(testAccountID, "viewer", modules.AgentNetworkProviders, operations.Update, false)
p, err := f.manager.GetProvider(ctx, testAccountID, "viewer", saved.ID)
require.NoError(t, err)
assert.Equal(t, saved.ID, p.ID)
assert.Empty(t, p.UpstreamURL)
}
// The self-scope tests drive the real permissions manager over the real
// store, so role resolution is the production one: a plain user holds no
// providers grant and must fall back to the caller-scoped list — the same
// selection the self-service setup answer derives from — while an admin
// keeps the account-wide view with full config.
// newSelfScopeStore seeds the account and its users only, so each test
// declares exactly the providers and policies it asserts on — the store
// rejects re-saving a policy id on MySQL, so tests never overwrite each
// other's rows.
func newSelfScopeStore(t *testing.T) (*managerImpl, store.Store) {
t.Helper()
mgr, s := newAgentConfigTestMgr(t)
mgr.permissionsManager = permissions.NewManager(s)
ctx := context.Background()
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: testAccountID}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "user-a", AccountID: testAccountID, Role: nbtypes.UserRoleUser, AutoGroups: []string{"grp-eng"},
}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "user-out", AccountID: testAccountID, Role: nbtypes.UserRoleUser,
}))
require.NoError(t, s.SaveUser(ctx, &nbtypes.User{
Id: "admin", AccountID: testAccountID, Role: nbtypes.UserRoleAdmin,
}))
return mgr, s
}
func newSelfScopeProvidersFixture(t *testing.T) (*managerImpl, store.Store) {
t.Helper()
mgr, s := newSelfScopeStore(t)
ctx := context.Background()
granted := newSynthTestProvider()
granted.ID = "prov-granted"
granted.Name = "Granted"
require.NoError(t, s.SaveAgentNetworkProvider(ctx, granted))
other := newSynthTestProvider()
other.ID = "prov-other"
other.Name = "Other"
other.CreatedAt = granted.CreatedAt.Add(time.Hour)
require.NoError(t, s.SaveAgentNetworkProvider(ctx, other))
disabled := newSynthTestProvider()
disabled.ID = "prov-disabled"
disabled.Name = "Disabled"
disabled.Enabled = false
require.NoError(t, s.SaveAgentNetworkProvider(ctx, disabled))
// user-a's group authorizes the granted and the disabled provider; the
// disabled one must still not surface (the proxy never routes it).
policy := newSynthTestPolicy(granted.ID, "grp-eng", "")
policy.DestinationProviderIDs = []string{granted.ID, disabled.ID}
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
return mgr, s
}
func TestGetAllProviders_SelfScopedForPlainUser(t *testing.T) {
ctx := context.Background()
mgr, _ := newSelfScopeProvidersFixture(t)
scoped, err := mgr.GetAllProviders(ctx, testAccountID, "user-a")
require.NoError(t, err, "a caller without the read grant self-scopes instead of being denied")
require.Len(t, scoped, 1)
assert.Equal(t, "prov-granted", scoped[0].ID)
assert.Empty(t, scoped[0].UpstreamURL, "the caller-scoped list is the redacted display surface")
assert.NotEmpty(t, scoped[0].Models, "model list backs the dashboard filters")
empty, err := mgr.GetAllProviders(ctx, testAccountID, "user-out")
require.NoError(t, err)
assert.Empty(t, empty, "a caller outside every policy gets an empty list, not an error")
all, err := mgr.GetAllProviders(ctx, testAccountID, "admin")
require.NoError(t, err)
assert.Len(t, all, 3, "grant holders keep the account-wide list, disabled providers included")
for _, p := range all {
if p.ID == "prov-granted" {
assert.NotEmpty(t, p.UpstreamURL, "a managing caller sees the connection config")
}
}
}
func TestGetProvider_SelfScopedForPlainUser(t *testing.T) {
ctx := context.Background()
mgr, _ := newSelfScopeProvidersFixture(t)
p, err := mgr.GetProvider(ctx, testAccountID, "user-a", "prov-granted")
require.NoError(t, err)
assert.Equal(t, "prov-granted", p.ID)
assert.Empty(t, p.UpstreamURL)
assertNotFound := func(id string) {
t.Helper()
_, err := mgr.GetProvider(ctx, testAccountID, "user-a", id)
require.Error(t, err)
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
assert.Equal(t, status.NotFound, sErr.Type(),
"out-of-scope and nonexistent providers must be indistinguishable")
}
assertNotFound("prov-other")
assertNotFound("prov-disabled")
assertNotFound("prov-does-not-exist")
}
func TestGetAllProviders_SelfScopedModelsFollowGuardrails(t *testing.T) {
ctx := context.Background()
mgr, s := newSelfScopeStore(t)
// A provider declaring two models, restricted by an allowlist admitting
// one declared model plus one the operator never declared (unreachable —
// the router only claims declared models, so it must not surface).
granted := newSynthTestProvider()
granted.ID = "prov-models"
granted.Name = "Granted"
granted.Models = []types.ProviderModel{
{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02},
{ID: "gpt-4o", InputPer1k: 0.0025, OutputPer1k: 0.01},
}
require.NoError(t, s.SaveAgentNetworkProvider(ctx, granted))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-models", "gpt-5.4", "gpt-undeclared")))
policy := newSynthTestPolicy(granted.ID, "grp-eng", "guard-models")
policy.ID = "pol-guard-models"
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
scoped, err := mgr.GetAllProviders(ctx, testAccountID, "user-a")
require.NoError(t, err)
require.Len(t, scoped, 1)
require.Len(t, scoped[0].Models, 1,
"the self-scoped model list is the effective set: allowlist ∩ declared")
assert.Equal(t, "gpt-5.4", scoped[0].Models[0].ID)
assert.Equal(t, 0.004, scoped[0].Models[0].InputPer1k, "declared entry survives, prices included")
all, err := mgr.GetAllProviders(ctx, testAccountID, "admin")
require.NoError(t, err)
for _, p := range all {
if p.ID == granted.ID {
assert.Len(t, p.Models, 2,
"grant holders keep the full declared list — their usage view spans everyone's requests")
}
}
}
func TestGetAllProviders_SelfScopedAllowlistWithoutDeclaredModels(t *testing.T) {
ctx := context.Background()
mgr, s := newSelfScopeStore(t)
// No operator declaration: the router claims every model, so the
// allowlist union is the effective set and comes back as bare entries.
granted := newSynthTestProvider()
granted.ID = "prov-bare"
granted.Name = "Granted"
granted.Models = nil
require.NoError(t, s.SaveAgentNetworkProvider(ctx, granted))
require.NoError(t, s.SaveAgentNetworkGuardrail(ctx, newSetupTestGuardrail("guard-bare", "gpt-5.4")))
policy := newSynthTestPolicy(granted.ID, "grp-eng", "guard-bare")
policy.ID = "pol-guard-bare"
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
scoped, err := mgr.GetAllProviders(ctx, testAccountID, "user-a")
require.NoError(t, err)
require.Len(t, scoped, 1)
require.Len(t, scoped[0].Models, 1)
assert.Equal(t, "gpt-5.4", scoped[0].Models[0].ID)
}
func TestGetAllProviders_SelfScopedUnrestrictedFallsBackToCatalogModels(t *testing.T) {
ctx := context.Background()
mgr, s := newSelfScopeStore(t)
// Unrestricted policy on a provider without an operator declaration:
// the setup answer advertises the catalog models, and the scoped
// provider list must match so the model filter is never emptier than
// the setup page.
granted := newSynthTestProvider()
granted.ID = "prov-catalog"
granted.Name = "Granted"
granted.Models = nil
require.NoError(t, s.SaveAgentNetworkProvider(ctx, granted))
policy := newSynthTestPolicy(granted.ID, "grp-eng", "")
policy.ID = "pol-catalog"
require.NoError(t, s.SaveAgentNetworkPolicy(ctx, policy))
scoped, err := mgr.GetAllProviders(ctx, testAccountID, "user-a")
require.NoError(t, err)
require.Len(t, scoped, 1)
require.NotEmpty(t, scoped[0].Models, "catalog models back the filter when the operator declared none")
ids := make([]string, 0, len(scoped[0].Models))
for _, m := range scoped[0].Models {
ids = append(ids, m.ID)
}
assert.Equal(t, declaredModelIDs(granted), ids, "the scoped list mirrors the setup answer's declared/catalog set")
}
@@ -4,7 +4,7 @@ import (
"context"
"testing"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -5,12 +5,14 @@ import (
"runtime"
"strings"
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
@@ -26,6 +28,10 @@ type bootstrapFixture struct {
manager Manager
store store.Store
perms *permissions.MockManager
// vendor stands in for the provider credential check's vendor call, which
// runs on every provider write. Without it these tests would reach a real
// vendor to save a record.
vendor *stubLister
}
func newBootstrapFixture(t *testing.T) *bootstrapFixture {
@@ -47,10 +53,12 @@ func newBootstrapFixture(t *testing.T) *bootstrapFixture {
accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
vendor := &stubLister{}
return &bootstrapFixture{
manager: NewManager(st, perms, accounts, nil),
manager: NewManager(st, perms, accounts, nil, WithModelLister(vendor)),
store: st,
perms: perms,
vendor: vendor,
}
}
@@ -64,6 +72,57 @@ func (f *bootstrapFixture) createSettings(ctx context.Context, accountID, userID
return f.manager.CreateSettings(ctx, userID, types.DefaultSettings(accountID), proxyAddress, endpoint)
}
func ptrTo[T any](v T) *T { return &v }
// seedProxy registers a proxy in clusterAddr, heartbeating now, so the labeled
// bootstrap path has a real cluster to validate against. accountID empty makes
// it a shared (NetBird-operated) cluster; private mirrors the capability an
// proxy with private capabilities reports, nil an unreported one.
func (f *bootstrapFixture) seedProxy(t *testing.T, proxyID, accountID, clusterAddr string, private *bool) {
t.Helper()
f.seedProxyAt(t, proxyID, accountID, clusterAddr, private, time.Now().UTC())
}
// seedProxyAt is seedProxy with an explicit last-seen, for cases that need a
// proxy whose heartbeat has aged past the active window while its row (and so
// its cluster) is still on record.
func (f *bootstrapFixture) seedProxyAt(t *testing.T, proxyID, accountID, clusterAddr string, private *bool, lastSeen time.Time) {
t.Helper()
p := &proxy.Proxy{
ID: proxyID,
ClusterAddress: clusterAddr,
Status: proxy.StatusConnected,
LastSeen: lastSeen,
Capabilities: proxy.Capabilities{Private: private},
}
if accountID != "" {
p.AccountID = &accountID
}
require.NoError(t, f.store.SaveProxy(context.Background(), p), "seeding a proxy must succeed")
}
// seedPrivateCluster is the common case: a shared cluster with a connected
// proxy that has private capabilities, which is what a bootstrap requires.
func (f *bootstrapFixture) seedPrivateCluster(t *testing.T, clusterAddr string) {
t.Helper()
f.seedProxy(t, "proxy-"+clusterAddr, "", clusterAddr, ptrTo(true))
}
// requireForeignClusterRefusal asserts the refusal a pin onto another
// account's host gets, and that it left no row behind.
func (f *bootstrapFixture) requireForeignClusterRefusal(t *testing.T, err error, accountID string) {
t.Helper()
require.Error(t, err, "another account's host must be refused")
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
assert.Contains(t, err.Error(), "not available to this account",
"the error must say the host is not the account's to use")
_, err = f.store.GetAgentNetworkSettings(context.Background(), store.LockingStrengthNone, accountID)
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
}
// TestCreateSettingsRequiresPermission pins the gate: bootstrap assigns the
// account's immutable endpoint, a settings write requiring the settings
// Create permission — and a denial leaves no row behind.
@@ -88,6 +147,7 @@ func TestCreateSettingsRequiresPermission(t *testing.T) {
func TestCreateSettingsLabeled(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.seedPrivateCluster(t, "cluster1.example.com")
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "Cluster1.Example.com", "")
@@ -161,6 +221,7 @@ func TestCreateSettingsIdentityFieldValidation(t *testing.T) {
func TestCreateSettingsConflictsOnSecondBootstrap(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.seedPrivateCluster(t, "cluster1.example.com")
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
first, err := f.createSettings(ctx, "account1", "user1", "cluster1.example.com", "")
@@ -211,6 +272,7 @@ func TestCreateProviderHasNoSettingsSideEffects(t *testing.T) {
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
provider := types.NewProvider("account1")
provider.ProviderID = "openai_api"
provider.Name = "openai"
provider.UpstreamURL = "https://api.openai.com"
provider.APIKey = "sk-test"
@@ -223,3 +285,279 @@ func TestCreateProviderHasNoSettingsSideEffects(t *testing.T) {
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
assert.Error(t, err, "provider create must not conjure a settings row")
}
// TestCreateSettingsRejectsOfflineCluster is the guard against deciding on
// heartbeat freshness. A centralised cluster is refused while its proxies are
// live; the same cluster must stay refused once they stop heartbeating, which
// takes only a couple of minutes (proxyActiveThreshold). Judging on liveness
// would turn "wait for the proxy to go quiet" into a way to pin the account's
// immutable endpoint to a cluster that can never serve it.
func TestCreateSettingsRejectsOfflineCluster(t *testing.T) {
ctx := context.Background()
notPrivate := false
cases := map[string]*bool{
"centralised proxy gone quiet": &notPrivate,
// A cluster that could serve the gateway still has to have something
// live in it to prove so at bootstrap: refusing is the safe direction
// (reconnect the proxy and retry) where accepting is permanent.
"private proxy gone quiet": ptrTo(true),
}
for name, private := range cases {
t.Run(name, func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxyAt(t, "proxy1", "", "offline.example.com", private,
time.Now().UTC().Add(-time.Hour))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "offline.example.com", "")
require.Error(t, err, "a known cluster with nothing live in it must be rejected")
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
assert.Contains(t, err.Error(), "private capabilities",
"the error must say private capabilities are what is missing")
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
})
}
}
// TestCreateSettingsRequiresPrivateCluster pins the capability gate: the
// synthesised gateway service is always private, so a live cluster whose
// proxies lack private capabilities cannot serve it and must not
// become the account's immutable endpoint.
func TestCreateSettingsRequiresPrivateCluster(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
notPrivate := false
f.seedProxy(t, "proxy1", "", "central.example.com", &notPrivate)
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "central.example.com", "")
require.Error(t, err, "a cluster without private capabilities must be rejected")
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
assert.Contains(t, err.Error(), "private capabilities", "the error must name what the cluster is missing")
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
}
// TestCreateSettingsAcceptsOwnPrivateCluster pins the BYOP happy path: the
// account's own cluster with a connected private-capable proxy is a valid pin.
func TestCreateSettingsAcceptsOwnPrivateCluster(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.seedProxy(t, "proxy1", "account1", "byop.account1.example.com", ptrTo(true))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "byop.account1.example.com", "")
require.NoError(t, err, "the account's own private cluster must be accepted")
assert.Equal(t, "byop.account1.example.com", created.ProxyAddress)
}
// TestCreateSettingsMatchesClusterCasing pins that a cluster spelled with
// capitals in the store is still recognised as the same cluster the normalised
// proxy_address names, in both directions: a private cluster is accepted and a
// centralised one is refused, whatever the casing. The comparison is in memory
// over the account's cluster list; the capability lookup is still asked under
// the spelling the store actually holds, which is what an exact match needs.
func TestCreateSettingsMatchesClusterCasing(t *testing.T) {
ctx := context.Background()
t.Run("own private cluster is found", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "proxy1", "", "EU.Proxy.Example.com", ptrTo(true))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "eu.proxy.example.com", "")
require.NoError(t, err, "a private cluster declared with capitals must still be accepted")
assert.Equal(t, "eu.proxy.example.com", created.ProxyAddress)
})
t.Run("non-private cluster is still refused", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "proxy1", "", "Central.Example.com", ptrTo(false))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "central.example.com", "")
require.Error(t, err, "casing must not become a way past the capability check")
assert.Contains(t, err.Error(), "private capabilities")
})
}
// TestCreateSettingsRejectsForeignCluster pins tenant consistency on the pin:
// an account may not pin its gateway onto a host another account's proxy
// declares. That proxy only ever receives its own account's mappings, so the
// pin could never be served, and the endpoint it assigns is immutable.
// Ownership is decided on the proxy rows, not on heartbeat freshness — a
// cluster whose proxies are merely offline is still somebody's — and on the
// normalised host, since proxies declare their address as the operator
// spelled it.
func TestCreateSettingsRejectsForeignCluster(t *testing.T) {
ctx := context.Background()
cases := map[string]struct {
spelling string
lastSeen time.Time
}{
"live": {"byop.account2.example.com", time.Now().UTC()},
"offline": {"byop.account2.example.com", time.Now().UTC().Add(-time.Hour)},
"spelled in caps": {"BYOP.Account2.Example.com", time.Now().UTC()},
}
for name, tc := range cases {
t.Run("labeled "+name, func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxyAt(t, "proxy1", "account2", tc.spelling, ptrTo(true), tc.lastSeen)
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "byop.account2.example.com", "")
f.requireForeignClusterRefusal(t, err, "account1")
})
t.Run("self-addressed "+name, func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxyAt(t, "proxy1", "account2", tc.spelling, ptrTo(true), tc.lastSeen)
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "", "byop.account2.example.com")
f.requireForeignClusterRefusal(t, err, "account1")
})
}
}
// TestCreateSettingsSharedClusterStaysPinnable pins the constraint the
// ownership check must respect: a shared (NetBird-operated) cluster is not
// anybody's, so any number of accounts pin their gateways to it — including
// an account that also runs a proxy of its own elsewhere.
func TestCreateSettingsSharedClusterStaysPinnable(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.seedProxy(t, "shared", "", "eu.proxy.netbird.io", ptrTo(true))
f.seedProxy(t, "own", "account1", "byop.account1.example.com", ptrTo(true))
for _, account := range []string{"account1", "account2"} {
f.expectPermission(account, "user", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, account, "user", "eu.proxy.netbird.io", "")
require.NoError(t, err, "a shared cluster must stay pinnable by %s", account)
assert.Equal(t, "eu.proxy.netbird.io", created.ProxyAddress)
}
}
// TestCreateSettingsOwnClusterIsPinnable is the BYOP order in both directions:
// the account's own proxy is not a competing claim, whether the pin is labeled
// beneath its cluster or self-addressed onto the very host it declares.
func TestCreateSettingsOwnClusterIsPinnable(t *testing.T) {
ctx := context.Background()
t.Run("labeled", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "own", "account1", "byop.account1.example.com", ptrTo(true))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "byop.account1.example.com", "")
require.NoError(t, err, "the account's own cluster must be pinnable")
assert.True(t, strings.HasSuffix(created.Domain, ".byop.account1.example.com"))
})
t.Run("self-addressed", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "own", "account1", "gw.account1.example.com", ptrTo(true))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "", "gw.account1.example.com")
require.NoError(t, err, "the host the account's own proxy declares must be pinnable")
assert.Equal(t, "gw.account1.example.com", created.ProxyAddress)
})
}
// TestCreateSettingsUnknownHostIsPinnable pins the address-first order: a host
// no proxy has ever declared is nobody's, so the pin goes through and the
// proxy is deployed after.
func TestCreateSettingsUnknownHostIsPinnable(t *testing.T) {
ctx := context.Background()
f := newBootstrapFixture(t)
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "future.example.com", "")
require.NoError(t, err, "a host no proxy has declared must stay pinnable")
assert.Equal(t, "future.example.com", created.ProxyAddress)
}
// TestCreateSettingsRejectsHostAnotherAccountPinned covers claims made by pins
// rather than proxies, which the proxy-row check cannot see. A labeled pin
// beneath a host makes that host the other account's cluster, so a
// self-addressed endpoint on it would never be served; a self-addressed
// endpoint on a host makes the proxy declaring it theirs, so a label beneath
// it would never be served either. Neither is a shared-cluster shape: many
// labeled pins under one cluster are asked about in neither direction.
func TestCreateSettingsRejectsHostAnotherAccountPinned(t *testing.T) {
ctx := context.Background()
t.Run("self-addressed onto another account's cluster", func(t *testing.T) {
f := newBootstrapFixture(t)
f.expectPermission("account2", "user2", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account2", "user2", "gw.example.com", "")
require.NoError(t, err, "account2's labeled pin beneath the host must go through first")
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err = f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
f.requireForeignClusterRefusal(t, err, "account1")
})
t.Run("labeled beneath another account's endpoint", func(t *testing.T) {
f := newBootstrapFixture(t)
f.expectPermission("account2", "user2", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account2", "user2", "", "gw.example.com")
require.NoError(t, err, "account2's self-addressed endpoint must go through first")
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err = f.createSettings(ctx, "account1", "user1", "gw.example.com", "")
f.requireForeignClusterRefusal(t, err, "account1")
})
t.Run("labeled beside another account's labeled pin stays allowed", func(t *testing.T) {
f := newBootstrapFixture(t)
for _, account := range []string{"account1", "account2"} {
f.expectPermission(account, "user", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, account, "user", "eu.proxy.netbird.io", "")
require.NoError(t, err, "labeled pins under one cluster are the shared-cluster shape and must not refuse each other")
}
})
}
// TestCreateSettingsSelfAddressedRequiresPrivateCluster pins that the
// capability gate applies to a self-addressed endpoint too: the service behind
// it is the same private one, so a proxy that already declares the hostname
// must have private capabilities, whether the account's own or a shared cluster's. A
// hostname no proxy declares yet stays claimable (TestCreateSettingsSelfAddressed).
func TestCreateSettingsSelfAddressedRequiresPrivateCluster(t *testing.T) {
ctx := context.Background()
t.Run("centralised proxy at the hostname is refused", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "central", "", "gw.example.com", ptrTo(false))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
_, err := f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
require.Error(t, err, "a self-addressed endpoint on a centralised proxy can never be served")
var sErr *status.Error
require.ErrorAs(t, err, &sErr)
assert.Equal(t, status.InvalidArgument, sErr.Type())
assert.Contains(t, err.Error(), "private capabilities")
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
})
t.Run("private proxy at the hostname is accepted", func(t *testing.T) {
f := newBootstrapFixture(t)
f.seedProxy(t, "private", "", "gw.example.com", ptrTo(true))
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
created, err := f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
require.NoError(t, err)
assert.Equal(t, "gw.example.com", created.ProxyAddress)
})
}
@@ -10,6 +10,7 @@ import (
"strings"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
@@ -210,8 +211,21 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
}
groupIndex := indexProviderGroups(enabledPolicies)
catalogByProvider := catalogIDsByProvider(enabledProviders)
routerCfgJSON, err := buildRouterConfigJSON(enabledProviders, groupIndex)
// The proxy guardrail is a per-provider fail-closed backstop; the
// authoritative per-policy/group decision is management's
// SelectPolicyForRequest. A provider lands in that map only when every
// authorising policy restricts models.
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID, catalogByProvider)
// Discovery gets the finer view: per policy rather than flattened per
// provider, so a listing can be bounded to what the calling groups may
// actually use instead of the union across everyone who reaches the
// provider.
modelPolicies := buildModelPolicies(enabledPolicies, guardrailsByID, catalogByProvider)
routerCfgJSON, err := buildRouterConfigJSON(enabledProviders, groupIndex, modelPolicies)
if err != nil {
return nil, err
}
@@ -228,11 +242,6 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
applyAccountCollectionControls(&mergedGuardrails, settings)
// The proxy guardrail is a per-provider fail-closed backstop; the
// authoritative per-policy/group decision is management's
// SelectPolicyForRequest. A provider lands in this map only when every
// authorising policy restricts models.
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID)
guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture)
if err != nil {
return nil, err
@@ -344,6 +353,7 @@ type routerConfig struct {
type routerProviderRoute struct {
ID string `json:"id"`
Vendor string `json:"vendor,omitempty"`
Vendors []string `json:"vendors,omitempty"`
Models []string `json:"models"`
UpstreamScheme string `json:"upstream_scheme"`
UpstreamHost string `json:"upstream_host"`
@@ -351,6 +361,11 @@ type routerProviderRoute struct {
AuthHeaderName string `json:"auth_header_name"`
AuthHeaderValue string `json:"auth_header_value"`
AllowedGroupIDs []string `json:"allowed_group_ids,omitempty"`
// ModelPolicies is one entry per enabled policy authorising this provider,
// carrying that policy's source groups and the models it permits. The
// router bounds a model listing with it, so a provider two groups reach
// under different allowlists offers each only its own.
ModelPolicies []routerModelPolicy `json:"model_policies,omitempty"`
// Vertex marks a Google Vertex AI provider, whose requests carry the
// model in the URL path. The router selects it by path, bypassing the
// model/vendor table.
@@ -368,6 +383,9 @@ type routerProviderRoute struct {
// proxy dials this provider's upstream. For self-hosted / internal gateways
// behind a private or self-signed certificate.
SkipTLSVerify bool `json:"skip_tls_verify,omitempty"`
// DiscoveryHost, when set, is the host serving this provider's model
// listing, for a vendor that does not serve it from the inference host.
DiscoveryHost string `json:"discovery_host,omitempty"`
}
// indexProviderGroups walks the enabled policies and returns, per
@@ -422,7 +440,7 @@ func indexProviderGroups(policies []*types.Policy) map[string][]string {
// path-prefix tiebreak. Providers no enabled policy authorises
// (orphans) are intentionally OMITTED so the router never observes a
// route with an empty ACL.
func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) {
func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]string, modelPolicies map[string][]routerModelPolicy) ([]byte, error) {
cfg := routerConfig{Providers: make([]routerProviderRoute, 0, len(providers))}
for _, p := range providers {
groups, hasPolicy := groupIndex[p.ID]
@@ -435,6 +453,9 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
if err != nil {
return nil, fmt.Errorf("router config for provider %s: %w", p.ID, err)
}
// Lookup rather than assume: an unknown provider id yields the zero
// entry, which declares no discovery and so contributes nothing.
catalogEntry, _ := catalog.Lookup(p.ProviderID)
headerName, headerValue, gcpSAKeyB64, err := providerAuthHeader(p)
if err != nil {
return nil, err
@@ -442,6 +463,7 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
cfg.Providers = append(cfg.Providers, routerProviderRoute{
ID: p.ID,
Vendor: providerVendor(p),
Vendors: providerVendors(p),
Models: providerModelIDs(p),
UpstreamScheme: scheme,
UpstreamHost: host,
@@ -449,10 +471,12 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
AuthHeaderName: headerName,
AuthHeaderValue: headerValue,
AllowedGroupIDs: groups,
ModelPolicies: modelPolicies[p.ID],
Vertex: catalog.IsVertexPathStyle(p.ProviderID),
Bedrock: catalog.IsBedrockPathStyle(p.ProviderID),
GCPServiceAccountKeyB64: gcpSAKeyB64,
SkipTLSVerify: p.SkipTLSVerification,
DiscoveryHost: discoveryHost(catalogEntry, p.UpstreamURL),
})
}
out, err := json.Marshal(cfg)
@@ -462,6 +486,33 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
return out, nil
}
// discoveryHost returns the host serving this provider's model listing when it
// differs from the inference host, and empty when the two are the same — which
// is true of every vendor but Bedrock, whose ListInferenceProfiles is a control
// plane operation on bedrock.<region> while InvokeModel must go to
// bedrock-runtime.<region>. One provider record therefore needs two hosts.
//
// The catalog declares the listing host; the region is recovered from the
// upstream the operator configured, since a provider record carries no region
// field. An upstream matching no catalog template yields empty rather than a
// guess: a proxied or self-hosted Bedrock endpoint may serve both from one
// place, and inventing a host would send the credential somewhere the operator
// never configured.
func discoveryHost(entry catalog.Provider, upstreamURL string) string {
if entry.Discovery == nil || entry.Discovery.Host == "" {
return ""
}
host := entry.Discovery.Host
if !strings.Contains(host, catalog.RegionPlaceholder) {
return host
}
region := modeldiscovery.RegionFromUpstream(entry, upstreamURL)
if region == "" {
return ""
}
return strings.ReplaceAll(host, catalog.RegionPlaceholder, region)
}
// providerVendor returns the parser surface ("openai", "anthropic", …)
// the provider speaks, sourced from its catalog entry's ParserID. The
// router uses it to keep a request the parser tagged with a vendor on a
@@ -477,6 +528,17 @@ func providerVendor(p *types.Provider) string {
return entry.ParserID
}
// providerVendors returns the parser surfaces a multi-surface gateway route
// accepts. Single-surface providers keep using the singular vendor field so
// existing proxy versions and configurations retain their wire shape.
func providerVendors(p *types.Provider) []string {
entry, ok := catalog.Lookup(p.ProviderID)
if !ok || len(entry.RouterVendors) == 0 {
return nil
}
return append([]string(nil), entry.RouterVendors...)
}
// providerModelIDs returns the model identifiers exposed by the
// provider, deduplicated and in the operator's declared order. Empty
// slice when no models are configured — the router treats that as
@@ -846,7 +908,9 @@ func marshalGuardrailConfig(providerAllowlists map[string][]string, capture Merg
// buildProviderAllowlists returns the proxy's per-provider backstop: a provider
// is included only when every authorising policy restricts models (their union);
// if any leaves it unrestricted it is omitted, so management decides per group.
func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]string {
// Entries carry their provider-specific canonical form alongside the verbatim
// one, resolved through catalogByProvider.
func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail, catalogByProvider map[string]string) map[string][]string {
type providerAcc struct {
models map[string]struct{}
anyUnrestricted bool
@@ -870,7 +934,7 @@ func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Gu
acc.anyUnrestricted = true
continue
}
for _, m := range models {
for _, m := range expandModelsForProvider(models, catalogByProvider[providerID]) {
acc.models[m] = struct{}{}
}
}
@@ -891,8 +955,10 @@ func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Gu
}
// policyModelAllowlist reports whether a policy restricts models (has an
// allowlist-enabled guardrail) and the union of allowed models. Models are
// verbatim; the proxy factory lowercases/trims them at decode time.
// allowlist-enabled guardrail) and the union of allowed models, verbatim.
// Consumers expand the entries per destination provider with
// expandModelsForProvider — the canonical form is provider-specific — and
// the proxy factory lowercases/trims them at decode time.
func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bool, []string) {
restricted := false
var models []string
@@ -911,6 +977,45 @@ func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bo
return restricted, models
}
// expandModelsForProvider returns the allowlist entries for one destination
// provider: each entry verbatim plus, when it differs, its canonical form
// under that provider's catalog id — the id the proxy's parser emits at
// request time — deduplicated. The proxy-side compares (guardrail backstop,
// per-group router rules) then admit an allowlist however the operator
// wrote it, raw declared id or canonical, while a plain provider's entries
// stay verbatim and can never widen.
func expandModelsForProvider(models []string, catalogProviderID string) []string {
out := make([]string, 0, len(models))
seen := make(map[string]struct{}, len(models))
add := func(m string) {
if m == "" {
return
}
if _, dup := seen[m]; dup {
return
}
seen[m] = struct{}{}
out = append(out, m)
}
for _, m := range models {
add(m)
add(canonicalModelKey(catalogProviderID, m))
}
return out
}
// catalogIDsByProvider indexes providers' catalog ids by provider record id,
// the lookup the per-provider allowlist expansion keys the normalizer on.
func catalogIDsByProvider(providers []*types.Provider) map[string]string {
out := make(map[string]string, len(providers))
for _, p := range providers {
if p != nil {
out[p.ID] = p.ProviderID
}
}
return out
}
// buildAccountService composes the per-account gateway Service. The
// target carries the noop placeholder URL — the router middleware
// rewrites every request to the matched provider's upstream before the
@@ -1098,3 +1203,48 @@ func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) {
}
}
}
// routerModelPolicy mirrors the router's ModelPolicyRule: one authorising
// policy's source groups plus the models it permits. Models is nil for a
// policy that sets no model allowlist, which lifts the restriction for the
// groups it binds — so nil and empty must survive the round trip distinctly.
type routerModelPolicy struct {
GroupIDs []string `json:"group_ids"`
Models []string `json:"models"`
}
// buildModelPolicies indexes, per provider, one rule for each enabled policy
// authorising it: the policy's source groups and the models its guardrail
// permits.
//
// This is deliberately finer than buildProviderAllowlists, which flattens the
// same inputs into one list per provider for the proxy's fail-closed guardrail.
// A flattened list cannot answer "what may THIS caller see", so a provider two
// teams reach under different allowlists would offer each team the other's
// models — a picker full of entries the next request refuses. Keeping the
// source groups alongside the models lets the router answer it at request time,
// where it knows the caller's groups.
func buildModelPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, catalogByProvider map[string]string) map[string][]routerModelPolicy {
out := make(map[string][]routerModelPolicy)
for _, p := range policies {
if p == nil || len(p.SourceGroups) == 0 {
continue
}
restricted, models := policyModelAllowlist(p, byID)
for _, providerID := range p.DestinationProviderIDs {
if providerID == "" {
continue
}
rule := routerModelPolicy{GroupIDs: append([]string(nil), p.SourceGroups...)}
if restricted {
// Never nil when restricted: an allowlist permitting nothing
// must stay distinguishable from no allowlist at all. The
// expansion is per provider — the canonical form of an entry
// depends on the destination's catalog id.
rule.Models = append([]string{}, expandModelsForProvider(models, catalogByProvider[providerID])...)
}
out[providerID] = append(out[providerID], rule)
}
}
return out
}
@@ -103,3 +103,37 @@ func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) {
assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry")
assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced")
}
// TestBuildCostMeterConfig_BedrockGeographyOutsideTheOriginalFour is the
// accounting half of the geography bug. The docs tell operators to register a
// Bedrock id exactly as AWS issues it, region prefix included, and the cost
// meter keys its table by the normalized form. While the geography was matched
// against a list of four, a profile issued anywhere else kept its prefix,
// missed the catalog entry it was meant to inherit from, and billed with a
// zero entry underneath the operator's own rates — so every cache bucket
// metered free and a model priced only by catalog defaults metered at nothing
// at all.
func TestBuildCostMeterConfig_BedrockGeographyOutsideTheOriginalFour(t *testing.T) {
for _, geo := range []string{"jp", "au", "ca", "sa", "us-gov"} {
t.Run(geo, func(t *testing.T) {
bedrock := &types.Provider{
ID: "prov-bedrock",
ProviderID: "bedrock_api",
Enabled: true,
Models: []types.ProviderModel{
{ID: geo + ".anthropic.claude-sonnet-5-20260514-v1:0", InputPer1k: 0.003, OutputPer1k: 0.015},
},
}
raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}})
require.NoError(t, err)
cfg := decodeCostMeterConfig(t, raw)
e, ok := cfg.Pricing.Providers["prov-bedrock"]["anthropic.claude-sonnet-5"]
require.True(t, ok, "a %s profile must key by the same normalized id the parser emits", geo)
assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9,
"cache read must be inherited from the bedrock default entry, not left at zero")
assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9,
"cache creation must be inherited from the bedrock default entry, not left at zero")
})
}
}
@@ -4,6 +4,7 @@ import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
)
@@ -32,7 +33,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", []string{"g-opus"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
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")
})
@@ -42,7 +43,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", nil, "prov-x"), // no guardrail
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
assert.NotContains(t, got, "prov-x",
"a provider reachable by an un-guardrailed policy must be omitted so the proxy treats it as unrestricted")
})
@@ -51,7 +52,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-disabled"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
assert.NotContains(t, got, "prov-x",
"a policy whose only guardrail has a disabled allowlist is unrestricted")
})
@@ -61,7 +62,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", []string{"g-opus"}, "prov-y"),
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
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")
})
@@ -70,7 +71,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o"}, "prov-x", "prov-y"),
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
assert.Equal(t, []string{"gpt-4o"}, got["prov-x"])
assert.Equal(t, []string{"gpt-4o"}, got["prov-y"])
})
@@ -79,7 +80,7 @@ func TestBuildProviderAllowlists(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o", "g-opus"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
got := buildProviderAllowlists(policies, byID, nil)
assert.ElementsMatch(t, []string{"claude-opus-4", "gpt-4o"}, got["prov-x"],
"a policy's own multiple allowlist guardrails union together")
})
@@ -88,8 +89,144 @@ func TestBuildProviderAllowlists(t *testing.T) {
empty := map[string]*types.Guardrail{"g-empty": allowlistGuardrail("g-empty", "acc-1")}
got := buildProviderAllowlists([]*types.Policy{
policyForProviders("p1", []string{"g-empty"}, "prov-x"),
}, empty)
}, empty, nil)
assert.Equal(t, map[string][]string{"prov-x": {}}, got,
"an enabled-but-empty allowlist is restricted with an empty set, not unrestricted")
})
}
// policyForGroups builds an enabled policy binding the given source groups to
// the given providers under an optional guardrail.
func policyForGroups(id string, groups []string, guardrailIDs []string, providerIDs ...string) *types.Policy {
return &types.Policy{
ID: id,
Enabled: true,
SourceGroups: groups,
DestinationProviderIDs: providerIDs,
GuardrailIDs: guardrailIDs,
}
}
// TestBuildModelPolicies covers the finer index discovery needs. Where
// buildProviderAllowlists flattens every authorising policy into one list per
// provider — enough for a fail-closed backstop, but blind to who is asking —
// this keeps each policy's source groups beside its models so the router can
// bound a listing to the calling groups.
func TestBuildModelPolicies(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("each policy keeps its own groups and models", func(t *testing.T) {
policies := []*types.Policy{
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
policyForGroups("p2", []string{"grp-sales"}, []string{"g-opus"}, "prov-x"),
}
got := buildModelPolicies(policies, byID, nil)
assert.Equal(t, []routerModelPolicy{
{GroupIDs: []string{"grp-eng"}, Models: []string{"gpt-4o"}},
{GroupIDs: []string{"grp-sales"}, Models: []string{"claude-opus-4"}},
}, got["prov-x"],
"the two policies must stay separable so neither group is offered the other's models")
})
t.Run("an unrestricted policy carries nil models", func(t *testing.T) {
policies := []*types.Policy{
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
policyForGroups("p2", []string{"grp-admin"}, nil, "prov-x"),
}
got := buildModelPolicies(policies, byID, nil)
assert.Nil(t, got["prov-x"][1].Models,
"no allowlist must reach the router as nil, which lifts the restriction for its groups")
})
t.Run("a disabled allowlist is not a restriction", func(t *testing.T) {
policies := []*types.Policy{policyForGroups("p1", []string{"grp-eng"}, []string{"g-disabled"}, "prov-x")}
got := buildModelPolicies(policies, byID, nil)
assert.Nil(t, got["prov-x"][0].Models,
"a guardrail with the allowlist check off restricts nothing")
})
t.Run("an enabled allowlist with no models permits nothing", func(t *testing.T) {
byIDEmpty := map[string]*types.Guardrail{
"g-empty": {ID: "g-empty", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true}}},
}
policies := []*types.Policy{policyForGroups("p1", []string{"grp-eng"}, []string{"g-empty"}, "prov-x")}
got := buildModelPolicies(policies, byIDEmpty, nil)
require.NotNil(t, got["prov-x"][0].Models,
"an empty allowlist must not arrive as nil — that would read as unrestricted")
assert.Empty(t, got["prov-x"][0].Models)
})
t.Run("a policy binding no groups is skipped", func(t *testing.T) {
policies := []*types.Policy{policyForGroups("p1", nil, []string{"g-4o"}, "prov-x")}
assert.Empty(t, buildModelPolicies(policies, byID, nil),
"a policy with no source groups authorises nobody, so it bounds nobody's listing")
})
}
// TestSynthesizedAllowlists_ExpandDeclaredIDsPerProvider proves the
// synthesized allowlists carry the canonical form alongside a raw declared
// entry — under the destination provider's own catalog id, never another's —
// so the proxy-side compares (guardrail backstop, per-group router rules)
// admit the allowlist however the operator wrote it, while a plain provider's
// "-vN"- or "@"-suffixed entries stay verbatim and cannot widen.
func TestSynthesizedAllowlists_ExpandDeclaredIDsPerProvider(t *testing.T) {
byID := map[string]*types.Guardrail{
"g-raw": allowlistGuardrail("g-raw", "acc-1",
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
"claude-sonnet-4-5@20250929",
"gpt-4o"),
}
catalogByProvider := map[string]string{
"prov-bedrock": "bedrock_api",
"prov-vertex": "vertex_ai_api",
"prov-plain": "openai_api",
}
policies := []*types.Policy{
policyForGroups("p1", []string{"grp-eng"}, []string{"g-raw"},
"prov-bedrock", "prov-vertex", "prov-plain"),
}
t.Run("guardrail backstop expands under each provider's own normalizer", func(t *testing.T) {
got := buildProviderAllowlists(policies, byID, catalogByProvider)
assert.ElementsMatch(t, []string{
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
"anthropic.claude-sonnet-4-5",
"claude-sonnet-4-5@20250929",
"gpt-4o",
}, got["prov-bedrock"],
"the Bedrock destination strips geography/version, but must not apply Vertex's @-strip")
assert.ElementsMatch(t, []string{
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
"claude-sonnet-4-5@20250929",
"claude-sonnet-4-5",
"gpt-4o",
}, got["prov-vertex"],
"the Vertex destination strips @version, but must not apply Bedrock's suffix strip")
assert.ElementsMatch(t, []string{
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
"claude-sonnet-4-5@20250929",
"gpt-4o",
}, got["prov-plain"],
"a body-routed provider keeps every entry verbatim — no alternate can widen it")
})
t.Run("router model rules expand the same way", func(t *testing.T) {
got := buildModelPolicies(policies, byID, catalogByProvider)
require.Len(t, got["prov-bedrock"], 1)
assert.Contains(t, got["prov-bedrock"][0].Models, "anthropic.claude-sonnet-4-5")
assert.NotContains(t, got["prov-bedrock"][0].Models, "claude-sonnet-4-5")
require.Len(t, got["prov-vertex"], 1)
assert.Contains(t, got["prov-vertex"][0].Models, "claude-sonnet-4-5")
assert.NotContains(t, got["prov-vertex"][0].Models, "anthropic.claude-sonnet-4-5")
require.Len(t, got["prov-plain"], 1)
assert.ElementsMatch(t, []string{
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0",
"claude-sonnet-4-5@20250929",
"gpt-4o",
}, got["prov-plain"][0].Models)
})
}
@@ -6,10 +6,11 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/store"
@@ -496,6 +497,55 @@ func TestSynthesizeServices_IdentityInject_LiteLLM(t *testing.T) {
assert.Equal(t, "x-litellm-tags", entry.HeaderPair.TagsHeader)
}
func TestBuildIdentityInjectConfigJSON_Agentgateway(t *testing.T) {
provider := &types.Provider{
ID: "prov-agentgateway",
ProviderID: "agentgateway",
}
raw, err := buildIdentityInjectConfigJSON(
[]*types.Provider{provider},
map[string][]string{provider.ID: []string{"grp-eng"}},
)
require.NoError(t, err)
var cfg identityInjectConfig
require.NoError(t, json.Unmarshal(raw, &cfg))
require.Len(t, cfg.Providers, 1)
rule := cfg.Providers[0]
assert.Equal(t, provider.ID, rule.ProviderID)
require.NotNil(t, rule.HeaderPair)
assert.Nil(t, rule.JSONMetadata)
assert.Equal(t, "x-netbird-user-id", rule.HeaderPair.EndUserIDHeader)
assert.Equal(t, "x-netbird-groups", rule.HeaderPair.TagsHeader)
assert.False(t, rule.HeaderPair.EndUserIDInBody)
assert.False(t, rule.HeaderPair.TagsInBody)
}
func TestBuildRouterConfigJSON_AgentgatewayVendors(t *testing.T) {
provider := &types.Provider{
ID: "prov-agentgateway",
ProviderID: "agentgateway",
UpstreamURL: "https://gateway.example.com",
APIKey: "virtual-key",
}
raw, err := buildRouterConfigJSON(
[]*types.Provider{provider},
map[string][]string{provider.ID: {"grp-eng"}},
nil,
)
require.NoError(t, err)
var cfg routerConfig
require.NoError(t, json.Unmarshal(raw, &cfg))
require.Len(t, cfg.Providers, 1)
assert.Empty(t, cfg.Providers[0].Vendor,
"the singular vendor remains empty for a multi-surface gateway")
assert.Equal(t, []string{"openai", "anthropic"}, cfg.Providers[0].Vendors)
}
// TestSynthesizeServices_IdentityInject_Bifrost_OperatorOverrides
// covers the customizable HeaderPair contract. The Bifrost catalog
// entry sets HeaderPair.Customizable=true with x-bf-dim-* defaults
@@ -1245,3 +1295,57 @@ func TestSynthesizeServices_EmptyAPIKey_FailsClosed(t *testing.T) {
require.Error(t, err, "synthesis must refuse a provider with no api key")
assert.Contains(t, err.Error(), "no api key", "error must surface the missing credential")
}
// TestDiscoveryHost pins which providers get a separate listing host. Getting
// this wrong in either direction is costly: a missing host leaves Bedrock
// discovery 404ing at AWS, and a host on the wrong provider would send that
// provider's listing — and its credential — somewhere the operator never
// configured.
func TestDiscoveryHost(t *testing.T) {
entry := func(id string) catalog.Provider {
p, ok := catalog.Lookup(id)
require.True(t, ok, "catalog entry %s must exist", id)
return p
}
for _, tc := range []struct {
name string
entry catalog.Provider
upstream string
want string
}{
{
// ListInferenceProfiles is a control-plane operation; the runtime
// host answers <UnknownOperationException/> for it.
name: "bedrock splits the listing off the runtime host",
entry: entry("bedrock_api"), upstream: "https://bedrock-runtime.eu-central-1.amazonaws.com",
want: "bedrock.eu-central-1.amazonaws.com",
},
{
name: "bedrock in another region",
entry: entry("bedrock_api"), upstream: "https://bedrock-runtime.us-west-2.amazonaws.com",
want: "bedrock.us-west-2.amazonaws.com",
},
{
// A proxied Bedrock endpoint may well serve both from one place,
// and there is no region to read back out of it.
name: "proxied bedrock upstream yields no discovery host",
entry: entry("bedrock_api"), upstream: "https://bedrock.internal.example.com",
want: "",
},
{
name: "openai serves its listing from the same host",
entry: entry("openai_api"), upstream: "https://api.openai.com",
want: "",
},
{
name: "vertex serves its listing from the same host",
entry: entry("vertex_ai_api"), upstream: "https://us-east5-aiplatform.googleapis.com",
want: "",
},
} {
t.Run(tc.name, func(t *testing.T) {
assert.Equal(t, tc.want, discoveryHost(tc.entry, tc.upstream))
})
}
}
@@ -0,0 +1,42 @@
package types
// AgentConfig is the caller-scoped answer to "what may this caller
// use on the Agent Network?" — the account's proxy endpoint plus the
// providers and models the caller's groups authorize. It intentionally
// carries display metadata only: no keys, no upstream URLs, no policy or
// guardrail structure, and no hint of providers the caller cannot reach.
type AgentConfig struct {
// Configured is false only when the account has no Agent Network set
// up. A caller no policy covers yet still reads as configured, with an
// empty Providers list: every member gets the same connection config,
// and the empty list is what tells them to ask for access.
Configured bool
// Endpoint is the account's proxy base URL
// ("https://<subdomain>.<cluster>"), reachable over the NetBird tunnel
// only. Empty when Configured is false. Handing it to a member the
// policies do not cover authorizes nothing on its own — the proxy
// still refuses every request no policy permits.
Endpoint string
// Providers lists the providers at least one applicable policy
// authorizes for the caller, in the account's created_at order.
Providers []AgentConfigProvider
}
// AgentConfigProvider is one authorized provider in an AgentConfig.
type AgentConfigProvider struct {
// Name is the operator-assigned label, e.g. "Bedrock prod".
Name string
// CatalogID names the catalog entry, e.g. "anthropic_api".
CatalogID string
// APIFlavor is the request-body shape the provider speaks — the
// catalog entry's parser id ("anthropic", "openai"); empty when the
// proxy dispatches the provider by URL path instead.
APIFlavor string
// AllModelsAllowed is true when no model allowlist restricts this
// provider for the caller. Models then lists the declared/catalog
// models as a courtesy (possibly none for gateway-style providers).
AllModelsAllowed bool
// Models is the effective model allowlist for the caller, or the
// declared/catalog models when AllModelsAllowed is true.
Models []string
}
@@ -175,6 +175,26 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
// ToAPIResponse renders the provider as the API representation. The API
// key is intentionally never surfaced.
// RedactedForViewer returns a copy with the connection configuration
// blanked: upstream URL, operator-typed extra header values, identity
// header names, the TLS-verification override, and (defence in depth —
// they never reach the wire anyway) the sealed credentials. Read-only
// viewers such as usage_viewer only need the display surface — id,
// catalog id, name, enabled state, and the model list the usage filters
// resolve against — so their responses carry nothing about how the
// operator connects to the vendor.
func (p *Provider) RedactedForViewer() *Provider {
c := *p
c.UpstreamURL = ""
c.APIKey = ""
c.ExtraValues = nil
c.IdentityHeaderUserID = ""
c.IdentityHeaderGroups = ""
c.SkipTLSVerification = false
c.SessionPrivateKey = ""
return &c
}
func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
models := make([]api.AgentNetworkProviderModel, 0, len(p.Models))
for _, m := range p.Models {
@@ -5,7 +5,7 @@ import (
"encoding/json"
"testing"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -7,7 +7,7 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
@@ -1,6 +1,6 @@
package peers
//go:generate go run github.com/golang/mock/mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
//go:generate go tool mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
import (
"context"
@@ -1,5 +1,10 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./manager.go
//
// Generated by this command:
//
// mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
//
// Package peers is a generated GoMock package.
package peers
@@ -9,18 +14,19 @@ import (
net "net"
reflect "reflect"
gomock "github.com/golang/mock/gomock"
network_map "github.com/netbirdio/netbird/management/internals/controllers/network_map"
account "github.com/netbirdio/netbird/management/server/account"
integrated_validator "github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
peer "github.com/netbirdio/netbird/management/server/peer"
types "github.com/netbirdio/netbird/management/server/types"
gomock "go.uber.org/mock/gomock"
)
// MockManager is a mock of Manager interface.
type MockManager struct {
ctrl *gomock.Controller
recorder *MockManagerMockRecorder
isgomock struct{}
}
// MockManagerMockRecorder is the mock recorder for MockManager.
@@ -49,7 +55,7 @@ func (m *MockManager) CreateProxyPeer(ctx context.Context, accountID, peerKey, c
}
// CreateProxyPeer indicates an expected call of CreateProxyPeer.
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateProxyPeer", reflect.TypeOf((*MockManager)(nil).CreateProxyPeer), ctx, accountID, peerKey, cluster)
}
@@ -63,7 +69,7 @@ func (m *MockManager) DeletePeers(ctx context.Context, accountID string, peerIDs
}
// DeletePeers indicates an expected call of DeletePeers.
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeletePeers", reflect.TypeOf((*MockManager)(nil).DeletePeers), ctx, accountID, peerIDs, userID, checkConnected)
}
@@ -78,7 +84,7 @@ func (m *MockManager) GetAllPeers(ctx context.Context, accountID, userID string)
}
// GetAllPeers indicates an expected call of GetAllPeers.
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllPeers", reflect.TypeOf((*MockManager)(nil).GetAllPeers), ctx, accountID, userID)
}
@@ -93,7 +99,7 @@ func (m *MockManager) GetPeer(ctx context.Context, accountID, userID, peerID str
}
// GetPeer indicates an expected call of GetPeer.
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeer", reflect.TypeOf((*MockManager)(nil).GetPeer), ctx, accountID, userID, peerID)
}
@@ -108,7 +114,7 @@ func (m *MockManager) GetPeerAccountID(ctx context.Context, peerID string) (stri
}
// GetPeerAccountID indicates an expected call of GetPeerAccountID.
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerAccountID", reflect.TypeOf((*MockManager)(nil).GetPeerAccountID), ctx, peerID)
}
@@ -123,7 +129,7 @@ func (m *MockManager) GetPeerByTunnelIP(ctx context.Context, accountID string, i
}
// GetPeerByTunnelIP indicates an expected call of GetPeerByTunnelIP.
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerByTunnelIP", reflect.TypeOf((*MockManager)(nil).GetPeerByTunnelIP), ctx, accountID, ip)
}
@@ -138,7 +144,7 @@ func (m *MockManager) GetPeerID(ctx context.Context, peerKey string) (string, er
}
// GetPeerID indicates an expected call of GetPeerID.
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerID", reflect.TypeOf((*MockManager)(nil).GetPeerID), ctx, peerKey)
}
@@ -154,7 +160,7 @@ func (m *MockManager) GetPeerWithGroups(ctx context.Context, accountID, peerID s
}
// GetPeerWithGroups indicates an expected call of GetPeerWithGroups.
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerWithGroups", reflect.TypeOf((*MockManager)(nil).GetPeerWithGroups), ctx, accountID, peerID)
}
@@ -169,7 +175,7 @@ func (m *MockManager) GetPeersByGroupIDs(ctx context.Context, accountID string,
}
// GetPeersByGroupIDs indicates an expected call of GetPeersByGroupIDs.
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeersByGroupIDs", reflect.TypeOf((*MockManager)(nil).GetPeersByGroupIDs), ctx, accountID, groupsIDs)
}
@@ -181,7 +187,7 @@ func (m *MockManager) SetAccountManager(accountManager account.Manager) {
}
// SetAccountManager indicates an expected call of SetAccountManager.
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetAccountManager", reflect.TypeOf((*MockManager)(nil).SetAccountManager), accountManager)
}
@@ -193,7 +199,7 @@ func (m *MockManager) SetIntegratedPeerValidator(integratedPeerValidator integra
}
// SetIntegratedPeerValidator indicates an expected call of SetIntegratedPeerValidator.
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetIntegratedPeerValidator", reflect.TypeOf((*MockManager)(nil).SetIntegratedPeerValidator), integratedPeerValidator)
}
@@ -205,7 +211,7 @@ func (m *MockManager) SetNetworkMapController(networkMapController network_map.C
}
// SetNetworkMapController indicates an expected call of SetNetworkMapController.
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetNetworkMapController", reflect.TypeOf((*MockManager)(nil).SetNetworkMapController), networkMapController)
}
@@ -5,7 +5,7 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -1,5 +1,13 @@
package domain
import "time"
// ValidationTTL is the time available to validate a custom domain registration.
const ValidationTTL = 48 * time.Hour
// ID identifies a custom domain registration.
type ID string
type Type string
const (
@@ -8,12 +16,13 @@ const (
)
type Domain struct {
ID string `gorm:"unique;primaryKey;autoIncrement"`
Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts.
AccountID string `gorm:"index"`
TargetCluster string // The proxy cluster this domain should be validated against
Type Type `gorm:"-"`
Validated bool
ID string `gorm:"unique;primaryKey;autoIncrement"`
Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts.
AccountID string `gorm:"index"`
TargetCluster string // The proxy cluster this domain should be validated against
Type Type `gorm:"-"`
Validated bool
ValidationExpiresAt *time.Time `gorm:"index"`
// SupportsCustomPorts is populated at query time for free domains from the
// proxy cluster capabilities. Not persisted.
SupportsCustomPorts *bool `gorm:"-"`
@@ -36,7 +45,12 @@ func (d *Domain) EventMeta() map[string]any {
}
}
// Copy returns a copy with an independent validation deadline.
func (d *Domain) Copy() *Domain {
dCopy := *d
if d.ValidationExpiresAt != nil {
expiresAt := *d.ValidationExpiresAt
dCopy.ValidationExpiresAt = &expiresAt
}
return &dCopy
}
@@ -66,8 +66,8 @@ func TestExtractClusterFromFreeDomain(t *testing.T) {
func TestExtractClusterFromCustomDomains(t *testing.T) {
customDomains := []*domain.Domain{
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io"},
{Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io"},
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: true},
{Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io", Validated: true},
}
tests := []struct {
@@ -120,19 +120,49 @@ func TestExtractClusterFromCustomDomains(t *testing.T) {
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains)
assert.Equal(t, tc.wantOK, ok)
if ok {
assert.Equal(t, tc.wantVal, cluster)
cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains)
if !tc.wantOK {
assert.Equal(t, customDomainNoMatch, match, "unrelated domain should not match any custom domain")
return
}
assert.Equal(t, customDomainValidated, match, "validated custom domain should resolve a cluster")
assert.Equal(t, tc.wantVal, cluster)
})
}
}
// An unvalidated row must never yield a cluster: the account has not shown it
// controls the name, so no service may be bound to it.
func TestExtractClusterFromCustomDomains_UnvalidatedDomainRefused(t *testing.T) {
customDomains := []*domain.Domain{
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: false},
}
for _, serviceDomain := range []string{"example.com", "app.example.com"} {
t.Run(serviceDomain, func(t *testing.T) {
cluster, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
assert.Equal(t, customDomainUnvalidated, match, "unvalidated row must be reported as such")
assert.Empty(t, cluster, "unvalidated row must not resolve a cluster")
})
}
}
// A more specific unvalidated row must not shadow a validated parent domain.
func TestExtractClusterFromCustomDomains_ValidatedParentWinsOverUnvalidatedChild(t *testing.T) {
customDomains := []*domain.Domain{
{Domain: "example.com", TargetCluster: "cluster-generic", Validated: true},
{Domain: "app.example.com", TargetCluster: "cluster-app", Validated: false},
}
cluster, match := extractClusterFromCustomDomains("app.example.com", customDomains)
assert.Equal(t, customDomainValidated, match)
assert.Equal(t, "cluster-generic", cluster, "validated parent domain should provide the cluster")
}
func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) {
customDomains := []*domain.Domain{
{Domain: "example.com", TargetCluster: "cluster-generic"},
{Domain: "app.example.com", TargetCluster: "cluster-app"},
{Domain: "example.com", TargetCluster: "cluster-generic", Validated: true},
{Domain: "app.example.com", TargetCluster: "cluster-app", Validated: true},
}
tests := []struct {
@@ -164,8 +194,8 @@ func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) {
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains)
assert.True(t, ok)
cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains)
assert.Equal(t, customDomainValidated, match)
assert.Equal(t, tc.wantVal, cluster)
})
}
@@ -0,0 +1,73 @@
package manager
import (
"context"
"time"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
"github.com/netbirdio/netbird/management/server/activity"
)
const (
validationCleanupInterval = 60 * time.Minute
validationCleanupBatch = 100
)
// RunValidationCleanup removes expired registrations on startup and hourly until cancellation.
func (m Manager) RunValidationCleanup(ctx context.Context) {
ticker := time.NewTicker(validationCleanupInterval)
defer ticker.Stop()
for {
m.cleanupExpiredDomains(ctx, time.Now().UTC())
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
func (m Manager) cleanupExpiredDomains(ctx context.Context, now time.Time) {
var afterID domain.ID
for ctx.Err() == nil {
domains, err := m.store.GetExpiredCustomDomains(ctx, now, afterID, validationCleanupBatch)
if err != nil {
if ctx.Err() == nil {
log.WithContext(ctx).WithError(err).Error("list expired custom domain registrations")
}
return
}
for _, d := range domains {
if ctx.Err() != nil {
return
}
m.deleteExpiredDomain(ctx, d, now)
afterID = domain.ID(d.ID)
}
if len(domains) < validationCleanupBatch {
return
}
}
}
func (m Manager) deleteExpiredDomain(ctx context.Context, d *domain.Domain, now time.Time) {
deleted, err := m.store.DeleteExpiredCustomDomain(ctx, d, now)
if err != nil {
if ctx.Err() == nil {
log.WithContext(ctx).WithFields(log.Fields{"accountID": d.AccountID, "domainID": d.ID}).
WithError(err).Warn("could not expire custom domain registration")
}
return
}
if !deleted {
return
}
meta := d.EventMeta()
if d.ValidationExpiresAt != nil {
meta["validation_expires_at"] = d.ValidationExpiresAt.UTC().Format(time.RFC3339)
}
m.accountManager.StoreEvent(ctx, activity.SystemInitiator, d.ID, d.AccountID,
activity.CustomDomainValidationExpired, meta)
}
@@ -0,0 +1,274 @@
package manager
import (
"context"
"fmt"
"sync"
"testing"
"testing/synctest"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/mock_server"
nbstore "github.com/netbirdio/netbird/management/server/store"
)
func TestValidateDomain_ExpiredRegistration(t *testing.T) {
env := setupDomainTest(t)
ctx := context.Background()
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "expired.example.com", testCluster)
require.NoError(t, err)
expiresAt := time.Now().Add(-time.Second)
db := env.store.(*nbstore.SqlStore).GetDB()
require.NoError(t, db.Model(&domain.Domain{}).Where("id = ?", d.ID).
Update("validation_expires_at", expiresAt).Error)
env.resolver.set("validation.expired.example.com", testCluster)
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
stored := storedDomain(t, env.store, accountA, d.Domain)
require.NotNil(t, stored)
assert.False(t, stored.Validated, "an expired registration must not become usable before cleanup runs")
}
func TestCreateDomain_ValidationDeadline(t *testing.T) {
env := setupClockDomainTest(t)
synctest.Test(t, func(t *testing.T) {
ctx := context.Background()
createdAt := time.Now().UTC()
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "pending.example.com", testCluster)
require.NoError(t, err)
require.NotNil(t, d.ValidationExpiresAt)
assert.Equal(t, createdAt.Add(48*time.Hour), *d.ValidationExpiresAt, "new registrations get 48 hours")
time.Sleep(time.Hour)
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
stored := storedDomain(t, env.store, accountA, d.Domain)
require.NotNil(t, stored)
require.NotNil(t, stored.ValidationExpiresAt)
assert.WithinDuration(t, *d.ValidationExpiresAt, *stored.ValidationExpiresAt, 0, "failed validation must not extend the deadline")
})
}
func TestCleanupExpiredDomains_Boundaries(t *testing.T) {
env := setupDomainTest(t)
events := captureDomainEvents(env)
ctx := context.Background()
now := time.Now().UTC().Truncate(time.Second)
tests := []struct {
name string
expiresAt time.Time
validated bool
deleted bool
}{
{"expired", now.Add(-time.Second), false, true},
{"deadline", now, false, true},
{"pending", now.Add(time.Second), false, false},
{"validated", now.Add(-time.Hour), true, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
d := createExpiringDomain(t, env, tt.name+".example.com", tt.expiresAt)
if tt.validated {
require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validated", true).Error)
}
env.manager.cleanupExpiredDomains(ctx, now)
stored := storedDomain(t, env.store, accountA, d.Domain)
if !tt.deleted {
assert.NotNil(t, stored, "pending and validated registrations must survive cleanup")
return
}
assert.Nil(t, stored, "expired unused registrations must be removed")
replacement, err := env.manager.CreateDomain(ctx, accountB, accountBUser, d.Domain, testCluster)
require.NoError(t, err)
assert.NotEqual(t, d.ID, replacement.ID, "the released name must receive a fresh registration")
assert.False(t, replacement.Validated, "the new account must validate its own registration")
})
}
got := events.get()
require.Len(t, got, 2, "only successful expiration deletions emit events")
for _, event := range got {
assert.Equal(t, activity.CustomDomainValidationExpired, event.Activity, "use the requested expiration event")
assert.Equal(t, activity.SystemInitiator, event.InitiatorID, "cleanup is attributed to the system")
assert.Equal(t, accountA, event.AccountID, "expiration belongs to the original account")
assert.NotEmpty(t, event.TargetID, "retain the deleted domain ID")
assert.NotEmpty(t, event.Meta["domain"], "retain the deleted domain name")
assert.NotEmpty(t, event.Meta["validation_expires_at"], "include the validation deadline")
}
}
func TestCleanupExpiredDomains_ContinuesPastProtectedBatch(t *testing.T) {
env := setupDomainTest(t)
ctx := context.Background()
now := time.Now().UTC()
for i := range validationCleanupBatch {
d := createExpiringDomain(t, env, fmt.Sprintf("protected-%d.example.com", i), now.Add(-time.Hour))
require.NoError(t, env.store.CreateService(ctx, &rpservice.Service{
ID: fmt.Sprintf("service-%d", i), AccountID: accountA, Domain: "app." + d.Domain,
}))
}
unprotected := createExpiringDomain(t, env, "unused.example.com", now.Add(-time.Hour))
env.manager.cleanupExpiredDomains(ctx, now)
assert.Nil(t, storedDomain(t, env.store, accountA, unprotected.Domain), "protected registrations must not starve later batches")
remaining, err := env.store.ListCustomDomains(ctx, accountA)
require.NoError(t, err)
assert.Len(t, remaining, validationCleanupBatch, "all registrations with dependent services must survive")
}
func TestCleanupExpiredDomains_ConcurrentWorkers(t *testing.T) {
env := setupDomainTest(t)
events := captureDomainEvents(env)
now := time.Now().UTC()
d := createExpiringDomain(t, env, "concurrent.example.com", now.Add(-time.Hour))
var workers sync.WaitGroup
for range 2 {
workers.Go(func() { env.manager.cleanupExpiredDomains(context.Background(), now) })
}
workers.Wait()
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "one worker must remove the expired registration")
assert.Len(t, events.get(), 1, "only the worker that deletes the row may emit the event")
}
func TestRunValidationCleanup_HourlyAndRestart(t *testing.T) {
env := setupClockDomainTest(t)
synctest.Test(t, func(t *testing.T) {
events := captureDomainEvents(env)
now := time.Now().UTC()
startup := createExpiringDomain(t, env, "startup.example.com", now.Add(-time.Hour))
hourly := createExpiringDomain(t, env, "hourly.example.com", now.Add(time.Minute))
ctx, cancel := context.WithCancel(context.Background())
done := make(chan struct{})
go func() {
defer close(done)
env.manager.RunValidationCleanup(ctx)
}()
synctest.Wait()
assert.Nil(t, storedDomain(t, env.store, accountA, startup.Domain), "startup must collect overdue registrations")
time.Sleep(59 * time.Minute)
synctest.Wait()
assert.NotNil(t, storedDomain(t, env.store, accountA, hourly.Domain), "cleanup must wait for the 60-minute interval")
time.Sleep(time.Minute)
synctest.Wait()
assert.Nil(t, storedDomain(t, env.store, accountA, hourly.Domain), "the hourly scan must collect expired registrations")
cancel()
<-done
offline := createExpiringDomain(t, env, "offline.example.com", time.Now().UTC().Add(time.Minute))
time.Sleep(2 * time.Hour)
assert.NotNil(t, storedDomain(t, env.store, accountA, offline.Domain), "a stopped worker must not continue deleting")
ctx, cancel = context.WithCancel(context.Background())
done = make(chan struct{})
go func() {
defer close(done)
env.manager.RunValidationCleanup(ctx)
}()
synctest.Wait()
assert.Nil(t, storedDomain(t, env.store, accountA, offline.Domain), "restart must use the persisted deadline")
cancel()
<-done
assert.Len(t, events.get(), 3, "each deletion should emit an expiration event")
})
}
type blockingDomainResolver struct {
started chan struct{}
release chan struct{}
}
func (r blockingDomainResolver) LookupCNAME(context.Context, string) (string, error) {
close(r.started)
<-r.release
return testCluster + ".", nil
}
func TestValidateDomain_DeadlinePassesDuringLookup(t *testing.T) {
for _, cleanup := range []bool{false, true} {
t.Run(fmt.Sprintf("cleanup=%t", cleanup), func(t *testing.T) {
env := setupClockDomainTest(t)
synctest.Test(t, func(t *testing.T) {
events := captureDomainEvents(env)
ctx := context.Background()
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "late.example.com", testCluster)
require.NoError(t, err)
resolver := blockingDomainResolver{started: make(chan struct{}), release: make(chan struct{})}
env.manager.validator.Resolver = resolver
done := make(chan struct{})
go func() {
defer close(done)
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
}()
<-resolver.started
time.Sleep(48 * time.Hour)
if cleanup {
env.manager.cleanupExpiredDomains(ctx, time.Now().UTC())
_, err = env.store.CreateCustomDomain(ctx, accountB, d.Domain, testCluster, false)
require.NoError(t, err)
}
close(resolver.release)
<-done
owner := accountA
if cleanup {
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "late validation must not restore the old claim")
owner = accountB
}
stored := storedDomain(t, env.store, owner, d.Domain)
require.NotNil(t, stored)
assert.False(t, stored.Validated, "late validation must not validate either claim")
for _, event := range events.get() {
assert.NotEqual(t, activity.DomainValidated, event.Activity, "a rejected write must not emit a validation event")
}
})
})
}
}
func setupClockDomainTest(t *testing.T) *domainTestEnv {
t.Helper()
// Network driver watchers cannot share cancellation channels across synctest bubbles.
// Store boundary and concurrency tests still exercise the selected database engine.
t.Setenv("NETBIRD_STORE_ENGINE", "sqlite")
return setupDomainTest(t)
}
func createExpiringDomain(t *testing.T, env *domainTestEnv, name string, expiresAt time.Time) *domain.Domain {
t.Helper()
d, err := env.store.CreateCustomDomain(context.Background(), accountA, name, testCluster, false)
require.NoError(t, err)
require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validation_expires_at", expiresAt).Error)
d.ValidationExpiresAt = &expiresAt
return d
}
type domainEvents struct {
mu sync.Mutex
events []*activity.Event
}
func captureDomainEvents(env *domainTestEnv) *domainEvents {
events := &domainEvents{}
env.manager.accountManager = &mock_server.MockAccountManager{
StoreEventFunc: func(_ context.Context, initiator, target, account string, code activity.ActivityDescriber, meta map[string]any) {
if code == activity.DomainAdded {
return
}
events.mu.Lock()
defer events.mu.Unlock()
events.events = append(events.events, &activity.Event{
InitiatorID: initiator, TargetID: target, AccountID: account,
Activity: code.(activity.Activity), Meta: meta,
})
},
}
return events
}
func (e *domainEvents) get() []*activity.Event {
e.mu.Lock()
defer e.mu.Unlock()
return append([]*activity.Event(nil), e.events...)
}
@@ -6,6 +6,7 @@ import (
"fmt"
"net"
"strings"
"time"
log "github.com/sirupsen/logrus"
@@ -18,6 +19,7 @@ import (
"github.com/netbirdio/netbird/management/server/permissions/operations"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
nbdomain "github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/status"
)
@@ -26,11 +28,14 @@ type store interface {
GetAgentNetworkSettings(ctx context.Context, lockStrength nbstore.LockingStrength, accountID string) (*agentnetworkTypes.Settings, error)
GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error)
GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error)
ListFreeDomains(ctx context.Context, accountID string) ([]string, error)
ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error)
CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error)
UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error)
DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error
GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error)
DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error)
}
type proxyManager interface {
@@ -105,12 +110,13 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
// Add custom domains.
for _, d := range domains {
cd := &domain.Domain{
ID: d.ID,
Domain: d.Domain,
AccountID: accountID,
TargetCluster: d.TargetCluster,
Type: domain.TypeCustom,
Validated: d.Validated,
ID: d.ID,
Domain: d.Domain,
AccountID: accountID,
TargetCluster: d.TargetCluster,
Type: domain.TypeCustom,
Validated: d.Validated,
ValidationExpiresAt: d.ValidationExpiresAt,
}
if d.TargetCluster != "" {
cd.SupportsCustomPorts = m.proxyManager.ClusterSupportsCustomPorts(ctx, d.TargetCluster)
@@ -125,6 +131,7 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
return ret, nil
}
// CreateDomain registers a normalized custom domain and attempts DNS validation.
func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName, targetCluster string) (*domain.Domain, error) {
ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Create)
if err != nil {
@@ -134,6 +141,15 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
return nil, status.NewPermissionDeniedError()
}
parsed, err := nbdomain.FromString(strings.TrimSuffix(domainName, "."))
if err != nil {
return nil, status.Errorf(status.InvalidArgument, "invalid domain: %v", err)
}
domainName = parsed.PunycodeString()
if !nbdomain.IsValidDomainNoWildcard(domainName) {
return nil, status.Errorf(status.InvalidArgument, "invalid domain format")
}
// Verify the target cluster is in the available clusters for this account
allowList, err := m.getClusterAllowList(ctx, accountID)
if err != nil {
@@ -150,6 +166,10 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
return nil, fmt.Errorf("target cluster %s is not available", targetCluster)
}
if err := m.checkDomainAvailable(ctx, domainName); err != nil {
return nil, err
}
// Attempt an initial validation against the specified cluster only
var validated bool
if m.validator.IsValid(ctx, domainName, []string{targetCluster}) {
@@ -166,6 +186,23 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
return d, nil
}
// checkDomainAvailable reports whether the domain is free to claim. The unique
// index on the column is the real guard; this turns the violation into a
// conflict the caller can act on instead of a database error, and says nothing
// about which account holds the domain.
func (m Manager) checkDomainAvailable(ctx context.Context, domainName string) error {
_, err := m.store.GetCustomDomainByName(ctx, domainName)
if err == nil {
return status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName)
}
if sErr, ok := status.FromError(err); ok && sErr.Type() == status.NotFound {
return nil
}
return fmt.Errorf("look up domain: %w", err)
}
func (m Manager) DeleteDomain(ctx context.Context, accountID, userID, domainID string) error {
ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Delete)
if err != nil {
@@ -203,7 +240,9 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
log.WithFields(log.Fields{
"accountID": accountID,
"domainID": domainID,
}).WithError(err).Error("validate domain")
"userID": userID,
}).Error("validate domain: permission denied")
return
}
log.WithFields(log.Fields{
@@ -219,6 +258,14 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
}).WithError(err).Error("get custom domain from store")
return
}
if d.Validated {
return
}
if d.ValidationExpiresAt == nil || !time.Now().Before(*d.ValidationExpiresAt) {
log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}).
Debug("custom domain validation window has expired")
return
}
// Validate only against the domain's target cluster
targetCluster := d.TargetCluster
@@ -239,20 +286,21 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
}).Info("validating domain against target cluster")
if m.validator.IsValid(context.Background(), d.Domain, []string{targetCluster}) {
log.WithFields(log.Fields{
"accountID": accountID,
"domainID": domainID,
"domain": d.Domain,
}).Info("domain validated successfully")
d.Validated = true
if _, err := m.store.UpdateCustomDomain(context.Background(), accountID, d); err != nil {
log.WithFields(log.Fields{
entry := log.WithFields(log.Fields{
"accountID": accountID,
"domainID": domainID,
"domain": d.Domain,
}).WithError(err).Error("update custom domain in store")
}).WithError(err)
if sErr, ok := status.FromError(err); ok && sErr.Type() == status.PreconditionFailed {
entry.Debug("custom domain registration is no longer pending validation")
return
}
entry.Error("update custom domain in store")
return
}
log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}).
Info("custom domain validated successfully")
m.accountManager.StoreEvent(context.Background(), userID, domainID, accountID, activity.DomainValidated, d.EventMeta())
} else {
@@ -298,9 +346,12 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain
return "", fmt.Errorf("list custom domains: %w", err)
}
targetCluster, valid := extractClusterFromCustomDomains(domain, customDomains)
if valid {
targetCluster, match := extractClusterFromCustomDomains(domain, customDomains)
switch match {
case customDomainValidated:
return targetCluster, nil
case customDomainUnvalidated:
return "", status.Errorf(status.PreconditionFailed, "domain %s is not validated", domain)
}
return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain)
@@ -363,19 +414,46 @@ func (m Manager) reservedGatewayAddress(ctx context.Context, accountID string) (
return settings.ProxyAddress, nil
}
func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, bool) {
// customDomainMatch describes how a service domain relates to the account's
// custom domain rows.
type customDomainMatch int
const (
customDomainNoMatch customDomainMatch = iota
customDomainUnvalidated
customDomainValidated
)
// extractClusterFromCustomDomains finds the longest custom domain covering the
// service domain and reports its target cluster. Only a validated row yields a
// cluster: until the CNAME check has passed the account has not shown it
// controls the name, so no traffic may be routed for it.
func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, customDomainMatch) {
bestCluster := ""
bestLen := -1
matched := false
for _, cd := range customDomains {
if serviceDomain != cd.Domain && !strings.HasSuffix(serviceDomain, "."+cd.Domain) {
continue
}
matched = true
if !cd.Validated {
continue
}
if l := len(cd.Domain); l > bestLen {
bestLen = l
bestCluster = cd.TargetCluster
}
}
return bestCluster, bestLen >= 0
switch {
case bestLen >= 0:
return bestCluster, customDomainValidated
case matched:
return "", customDomainUnvalidated
default:
return "", customDomainNoMatch
}
}
// ExtractClusterFromFreeDomain extracts the cluster address from a free domain.
@@ -0,0 +1,321 @@
package manager
import (
"context"
"fmt"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/metric/noop"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/mock_server"
"github.com/netbirdio/netbird/management/server/permissions"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
const (
testCluster = "eu.proxy.test"
accountA = "account-a"
accountAUser = "account-a-admin"
accountB = "account-b"
accountBUser = "account-b-admin"
accountAMember = "account-a-member"
)
// stubResolver answers CNAME lookups from a table the test controls, so a
// domain can point at the cluster or nowhere without touching a real resolver.
type stubResolver struct {
mu sync.Mutex
cnames map[string]string
}
func (r *stubResolver) LookupCNAME(_ context.Context, host string) (string, error) {
r.mu.Lock()
defer r.mu.Unlock()
cname, ok := r.cnames[host]
if !ok {
return "", fmt.Errorf("lookup %s: no such host", host)
}
return cname + ".", nil
}
func (r *stubResolver) set(host, cname string) {
r.mu.Lock()
defer r.mu.Unlock()
r.cnames[host] = cname
}
type domainTestEnv struct {
manager Manager
store nbstore.Store
resolver *stubResolver
}
// setupDomainTest builds the domain manager on a real SQLite store with two
// accounts and one active public proxy cluster.
func setupDomainTest(t *testing.T) *domainTestEnv {
t.Helper()
ctx := context.Background()
testStore, cleanup, err := nbstore.NewTestStoreFromSQL(ctx, "", t.TempDir())
require.NoError(t, err)
t.Cleanup(cleanup)
for accountID, userID := range map[string]string{accountA: accountAUser, accountB: accountBUser} {
users := map[string]*types.User{
userID: {
Id: userID,
AccountID: accountID,
Role: types.UserRoleAdmin,
},
}
if accountID == accountA {
// A real member of the account whose role denies Services:Create, so
// permission denial is exercised as ok=false rather than as a lookup
// error for a user who is not in the account at all.
users[accountAMember] = &types.User{
Id: accountAMember,
AccountID: accountID,
Role: types.UserRoleUser,
}
}
require.NoError(t, testStore.SaveAccount(ctx, &types.Account{
Id: accountID,
CreatedBy: userID,
Settings: &types.Settings{},
Users: users,
}))
}
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil)
require.NoError(t, err)
resolver := &stubResolver{cnames: make(map[string]string)}
mgr := Manager{
store: testStore,
proxyManager: proxyMgr,
validator: domain.Validator{Resolver: resolver},
permissionsManager: permissions.NewManager(testStore),
accountManager: &mock_server.MockAccountManager{
StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {},
},
}
return &domainTestEnv{manager: mgr, store: testStore, resolver: resolver}
}
// storedDomain reads a domain row back through the store so assertions are made
// on what was persisted rather than on the value the manager returned.
func storedDomain(t *testing.T, s nbstore.Store, accountID, domainName string) *domain.Domain {
t.Helper()
domains, err := s.ListCustomDomains(context.Background(), accountID)
require.NoError(t, err)
for _, d := range domains {
if d.Domain == domainName {
return d
}
}
return nil
}
// A domain whose CNAME check fails is stored unvalidated and must not resolve a
// cluster, which is what service creation gates on.
func TestCreateDomain_FailedLookupIsNotServable(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "apps.example.com", testCluster)
require.NoError(t, err)
assert.False(t, created.Validated, "a domain whose CNAME lookup fails must not be created validated")
stored := storedDomain(t, env.store, accountA, "apps.example.com")
require.NotNil(t, stored, "domain row should exist")
assert.False(t, stored.Validated, "persisted row must be unvalidated")
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "apps.example.com")
require.Error(t, err, "an unvalidated domain must not resolve a cluster")
assert.Empty(t, cluster)
assert.Contains(t, err.Error(), "not validated", "error should tell the caller what to fix")
sErr, ok := status.FromError(err)
require.True(t, ok, "error should be a typed status error")
assert.Equal(t, status.PreconditionFailed, sErr.Type())
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "sub.apps.example.com")
assert.Error(t, err, "subdomains of an unvalidated custom domain are not servable either")
}
// A second account claiming a registered domain gets a clean conflict, not a
// database error surfaced as a 500.
func TestCreateDomain_DuplicateIsAConflict(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "shared.example.com", testCluster)
require.NoError(t, err)
_, err = env.manager.CreateDomain(ctx, accountB, accountBUser, "shared.example.com", testCluster)
require.Error(t, err)
sErr, ok := status.FromError(err)
require.True(t, ok, "conflict must be a typed status error, not a raw database error")
assert.Equal(t, status.AlreadyExists, sErr.Type(), "conflict should map to 409, not 500")
assert.NotContains(t, sErr.Message, accountA, "the response must not reveal the holding account")
assert.Nil(t, storedDomain(t, env.store, accountB, "shared.example.com"), "no row should be written on conflict")
}
// The same account re-adding one of its own domains is a conflict too.
func TestCreateDomain_SameAccountDuplicateIsAConflict(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster)
require.NoError(t, err)
_, err = env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster)
require.Error(t, err)
sErr, ok := status.FromError(err)
require.True(t, ok)
assert.Equal(t, status.AlreadyExists, sErr.Type())
}
// The negative control: a validated domain still derives its cluster, for the
// bare name and for subdomains, exactly as before.
func TestCreateDomain_ValidatedDomainDerivesCluster(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
env.resolver.set("validation.valid.example.com", testCluster)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "valid.example.com", testCluster)
require.NoError(t, err)
require.True(t, created.Validated, "a matching CNAME should validate on create")
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "valid.example.com")
require.NoError(t, err)
assert.Equal(t, testCluster, cluster)
cluster, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "app.valid.example.com")
require.NoError(t, err)
assert.Equal(t, testCluster, cluster, "subdomains of a validated custom domain resolve too")
}
// Validating a domain flips the gate: the same lookup that failed before now
// resolves a cluster.
func TestValidateDomain_UnlocksClusterDerivation(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "later.example.com", testCluster)
require.NoError(t, err)
require.False(t, created.Validated)
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com")
require.Error(t, err)
env.resolver.set("validation.later.example.com", testCluster)
env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID)
require.True(t, storedDomain(t, env.store, accountA, "later.example.com").Validated)
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com")
require.NoError(t, err)
assert.Equal(t, testCluster, cluster)
}
// Free cluster domains are unaffected by the custom domain gate.
func TestDeriveClusterFromDomain_FreeDomainUnaffected(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "myapp.abc123."+testCluster)
require.NoError(t, err)
assert.Equal(t, testCluster, cluster)
}
// The manager pre-check exists to turn a conflict into a 409, but the unique
// index on the column is what actually guarantees the domain is claimed once.
//
// Two requests can clear the pre-check concurrently and race to the insert.
// Inserting twice through the store reaches the same code path the loser of
// that race takes, without the nondeterminism of driving it from goroutines,
// and the loser must still see a conflict rather than an internal error.
func TestStore_DuplicateDomainRejectedByIndexAsConflict(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
_, err := env.store.CreateCustomDomain(ctx, accountA, "indexed.example.com", testCluster, false)
require.NoError(t, err)
_, err = env.store.CreateCustomDomain(ctx, accountB, "indexed.example.com", testCluster, false)
require.Error(t, err, "the unique index must reject the same domain in a second account")
sErr, ok := status.FromError(err)
require.True(t, ok, "the losing insert must return a typed status error")
assert.Equal(t, status.AlreadyExists, sErr.Type(), "a lost race is a 409, not a 500")
}
// Validation is what decides whether a domain routes traffic, so a caller
// without permission to it must not be able to flip the flag. The check logged
// the denial and then carried on, which was inert while nothing read Validated
// and is not once cluster derivation gates on it.
func TestValidateDomain_PermissionDeniedDoesNotValidate(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "guarded.example.com", testCluster)
require.NoError(t, err)
require.False(t, created.Validated)
// The CNAME is in place, so the only thing standing between this caller and
// a validated domain is the permission check.
env.resolver.set("validation.guarded.example.com", testCluster)
env.manager.ValidateDomain(ctx, accountA, accountAMember, created.ID)
stored := storedDomain(t, env.store, accountA, "guarded.example.com")
require.NotNil(t, stored)
assert.False(t, stored.Validated, "a caller without permission must not validate the domain")
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "guarded.example.com")
assert.Error(t, err, "the domain must still be unservable")
}
// A validation finishing after deletion must reject the stale write, without
// restoring the registration or reporting successful validation.
func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "racy.example.com", testCluster)
require.NoError(t, err)
stale := storedDomain(t, env.store, accountA, "racy.example.com")
require.NotNil(t, stale)
require.NoError(t, env.manager.DeleteDomain(ctx, accountA, accountAUser, created.ID))
require.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), "the domain should be gone")
// What an in-flight validation would write once its CNAME check succeeded.
stale.Validated = true
_, err = env.store.UpdateCustomDomain(ctx, accountA, stale)
require.Error(t, err, "a deleted registration must reject a late validation")
assert.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"),
"a late validation write must not recreate a deleted domain")
}
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -184,6 +185,10 @@ func (s *stubStore) GetCustomDomain(context.Context, string, string) (*domain.Do
panic("not used in allow-list tests")
}
func (s *stubStore) GetCustomDomainByName(context.Context, string) (*domain.Domain, error) {
panic("not used in allow-list tests")
}
func (s *stubStore) ListFreeDomains(context.Context, string) ([]string, error) {
panic("not used in allow-list tests")
}
@@ -204,6 +209,14 @@ func (s *stubStore) DeleteCustomDomain(context.Context, string, string) error {
panic("not used in allow-list tests")
}
func (s *stubStore) GetExpiredCustomDomains(context.Context, time.Time, domain.ID, int) ([]*domain.Domain, error) {
panic("not used in allow-list tests")
}
func (s *stubStore) DeleteExpiredCustomDomain(context.Context, *domain.Domain, time.Time) (bool, error) {
panic("not used in allow-list tests")
}
// TestGetClusterAllowList_DedicatedGatewayAddressExcluded pins invariant (B)'s
// chokepoint: a self-addressed settings pin reserves the account's gateway
// address, so it is dropped from the allow list — which, because the
@@ -0,0 +1,83 @@
package manager
import (
"context"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/shared/management/status"
)
func TestCreateDomain_NormalizesName(t *testing.T) {
for _, tt := range []struct {
name string
input string
canonical string
}{
{"mixed case", "Apps.Example.COM", "apps.example.com"},
{"unicode", "münchen.example.com", "xn--mnchen-3ya.example.com"},
{"trailing dot", "apps.example.com.", "apps.example.com"},
{"underscore", "My_App.example.com", "my_app.example.com"},
} {
t.Run(tt.name, func(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
env.resolver.set("validation."+tt.canonical, testCluster)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, tt.input, testCluster)
require.NoError(t, err)
assert.Equal(t, tt.canonical, created.Domain, "the response must use the normalized name")
assert.True(t, created.Validated, "the CNAME lookup must use the normalized name")
stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID)
require.NoError(t, err)
assert.Equal(t, tt.canonical, stored.Domain, "the database must retain the normalized name")
_, err = env.manager.CreateDomain(ctx, accountB, accountBUser, tt.canonical, testCluster)
require.Error(t, err)
sErr, ok := status.FromError(err)
require.True(t, ok, "an equivalent name must return a typed conflict")
assert.Equal(t, status.AlreadyExists, sErr.Type(), "normalization must precede the availability check")
})
}
}
func TestCreateDomain_NormalizedNameCanValidateLater(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "Apps.Example.COM.", testCluster)
require.NoError(t, err)
require.False(t, created.Validated, "a missing CNAME must leave the normalized registration pending")
env.resolver.set("validation.apps.example.com", testCluster)
env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID)
stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID)
require.NoError(t, err)
assert.Equal(t, "apps.example.com", stored.Domain, "retrying validation must retain the normalized name")
assert.True(t, stored.Validated, "later validation must look up the normalized name")
}
func TestCreateDomain_RejectsInvalidName(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
for _, name := range []string{
"", ".", "app..example.com", "app.example.com..", "-app.example.com",
"app%.example.com", "app!.example.com", "*.example.com", "app example.com",
"https://example.com", strings.Repeat("a", 64) + ".example.com",
} {
t.Run(name, func(t *testing.T) {
// A matching DNS response must not make a malformed name acceptable.
env.resolver.set("validation."+name, testCluster)
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, name, testCluster)
require.Error(t, err)
sErr, ok := status.FromError(err)
require.True(t, ok, "invalid names must return a typed client error")
assert.Equal(t, status.InvalidArgument, sErr.Type(), "malformed names must be rejected before storage")
})
}
stored, err := env.store.ListCustomDomains(ctx, accountA)
require.NoError(t, err)
assert.Empty(t, stored, "invalid registration attempts must not reserve any names")
}
@@ -1,6 +1,6 @@
package proxy
//go:generate go run github.com/golang/mock/mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
//go:generate go tool mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
import (
"context"
@@ -1,5 +1,10 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./manager.go
//
// Generated by this command:
//
// mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
//
// Package proxy is a generated GoMock package.
package proxy
@@ -9,14 +14,15 @@ import (
reflect "reflect"
time "time"
gomock "github.com/golang/mock/gomock"
proto "github.com/netbirdio/netbird/shared/management/proto"
gomock "go.uber.org/mock/gomock"
)
// MockManager is a mock of Manager interface.
type MockManager struct {
ctrl *gomock.Controller
recorder *MockManagerMockRecorder
isgomock struct{}
}
// MockManagerMockRecorder is the mock recorder for MockManager.
@@ -45,25 +51,11 @@ func (m *MockManager) CleanupStale(ctx context.Context, inactivityDuration time.
}
// CleanupStale indicates an expected call of CleanupStale.
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
}
// ClusterSupportsCustomPorts mocks base method.
func (m *MockManager) ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ClusterSupportsCustomPorts", ctx, clusterAddr)
ret0, _ := ret[0].(*bool)
return ret0
}
// ClusterSupportsCustomPorts indicates an expected call of ClusterSupportsCustomPorts.
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCustomPorts", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCustomPorts), ctx, clusterAddr)
}
// ClusterRequireSubdomain mocks base method.
func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
@@ -73,7 +65,7 @@ func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr s
}
// ClusterRequireSubdomain indicates an expected call of ClusterRequireSubdomain.
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterRequireSubdomain", reflect.TypeOf((*MockManager)(nil).ClusterRequireSubdomain), ctx, clusterAddr)
}
@@ -87,11 +79,25 @@ func (m *MockManager) ClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
}
// ClusterSupportsCrowdSec indicates an expected call of ClusterSupportsCrowdSec.
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCrowdSec", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCrowdSec), ctx, clusterAddr)
}
// ClusterSupportsCustomPorts mocks base method.
func (m *MockManager) ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ClusterSupportsCustomPorts", ctx, clusterAddr)
ret0, _ := ret[0].(*bool)
return ret0
}
// ClusterSupportsCustomPorts indicates an expected call of ClusterSupportsCustomPorts.
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCustomPorts", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCustomPorts), ctx, clusterAddr)
}
// ClusterSupportsPrivate mocks base method.
func (m *MockManager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
@@ -101,7 +107,7 @@ func (m *MockManager) ClusterSupportsPrivate(ctx context.Context, clusterAddr st
}
// ClusterSupportsPrivate indicates an expected call of ClusterSupportsPrivate.
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
}
@@ -116,11 +122,40 @@ func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAd
}
// Connect indicates an expected call of Connect.
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
}
// CountAccountProxies mocks base method.
func (m *MockManager) CountAccountProxies(ctx context.Context, accountID string) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "CountAccountProxies", ctx, accountID)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// CountAccountProxies indicates an expected call of CountAccountProxies.
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountAccountProxies", reflect.TypeOf((*MockManager)(nil).CountAccountProxies), ctx, accountID)
}
// DeleteAccountCluster mocks base method.
func (m *MockManager) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteAccountCluster", ctx, clusterAddress, accountID)
ret0, _ := ret[0].(error)
return ret0
}
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID)
}
// Disconnect mocks base method.
func (m *MockManager) Disconnect(ctx context.Context, proxyID, sessionID string) error {
m.ctrl.T.Helper()
@@ -130,11 +165,26 @@ func (m *MockManager) Disconnect(ctx context.Context, proxyID, sessionID string)
}
// Disconnect indicates an expected call of Disconnect.
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Disconnect", reflect.TypeOf((*MockManager)(nil).Disconnect), ctx, proxyID, sessionID)
}
// GetAccountProxy mocks base method.
func (m *MockManager) GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetAccountProxy", ctx, accountID)
ret0, _ := ret[0].(*Proxy)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetAccountProxy indicates an expected call of GetAccountProxy.
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountProxy", reflect.TypeOf((*MockManager)(nil).GetAccountProxy), ctx, accountID)
}
// GetActiveClusterAddresses mocks base method.
func (m *MockManager) GetActiveClusterAddresses(ctx context.Context) ([]string, error) {
m.ctrl.T.Helper()
@@ -145,11 +195,12 @@ func (m *MockManager) GetActiveClusterAddresses(ctx context.Context) ([]string,
}
// GetActiveClusterAddresses indicates an expected call of GetActiveClusterAddresses.
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddresses", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddresses), ctx)
}
// GetActiveClusterAddressesForAccount mocks base method.
func (m *MockManager) GetActiveClusterAddressesForAccount(ctx context.Context, accountID string) ([]string, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetActiveClusterAddressesForAccount", ctx, accountID)
@@ -158,7 +209,8 @@ func (m *MockManager) GetActiveClusterAddressesForAccount(ctx context.Context, a
return ret0, ret1
}
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID interface{}) *gomock.Call {
// GetActiveClusterAddressesForAccount indicates an expected call of GetActiveClusterAddressesForAccount.
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddressesForAccount", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddressesForAccount), ctx, accountID)
}
@@ -172,41 +224,11 @@ func (m *MockManager) Heartbeat(ctx context.Context, p *Proxy) error {
}
// Heartbeat indicates an expected call of Heartbeat.
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Heartbeat", reflect.TypeOf((*MockManager)(nil).Heartbeat), ctx, p)
}
// GetAccountProxy mocks base method.
func (m *MockManager) GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetAccountProxy", ctx, accountID)
ret0, _ := ret[0].(*Proxy)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetAccountProxy indicates an expected call of GetAccountProxy.
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountProxy", reflect.TypeOf((*MockManager)(nil).GetAccountProxy), ctx, accountID)
}
// CountAccountProxies mocks base method.
func (m *MockManager) CountAccountProxies(ctx context.Context, accountID string) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "CountAccountProxies", ctx, accountID)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// CountAccountProxies indicates an expected call of CountAccountProxies.
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountAccountProxies", reflect.TypeOf((*MockManager)(nil).CountAccountProxies), ctx, accountID)
}
// IsClusterAddressAvailable mocks base method.
func (m *MockManager) IsClusterAddressAvailable(ctx context.Context, clusterAddress, accountID string) (bool, error) {
m.ctrl.T.Helper()
@@ -217,29 +239,16 @@ func (m *MockManager) IsClusterAddressAvailable(ctx context.Context, clusterAddr
}
// IsClusterAddressAvailable indicates an expected call of IsClusterAddressAvailable.
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "IsClusterAddressAvailable", reflect.TypeOf((*MockManager)(nil).IsClusterAddressAvailable), ctx, clusterAddress, accountID)
}
// DeleteAccountCluster mocks base method.
func (m *MockManager) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteAccountCluster", ctx, clusterAddress, accountID)
ret0, _ := ret[0].(error)
return ret0
}
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID)
}
// MockController is a mock of Controller interface.
type MockController struct {
ctrl *gomock.Controller
recorder *MockControllerMockRecorder
isgomock struct{}
}
// MockControllerMockRecorder is the mock recorder for MockController.
@@ -282,7 +291,7 @@ func (m *MockController) GetProxiesForCluster(clusterAddr string) []string {
}
// GetProxiesForCluster indicates an expected call of GetProxiesForCluster.
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr interface{}) *gomock.Call {
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetProxiesForCluster", reflect.TypeOf((*MockController)(nil).GetProxiesForCluster), clusterAddr)
}
@@ -296,7 +305,7 @@ func (m *MockController) RegisterProxyToCluster(ctx context.Context, clusterAddr
}
// RegisterProxyToCluster indicates an expected call of RegisterProxyToCluster.
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterProxyToCluster", reflect.TypeOf((*MockController)(nil).RegisterProxyToCluster), ctx, clusterAddr, proxyID)
}
@@ -308,7 +317,7 @@ func (m *MockController) SendServiceUpdateToCluster(ctx context.Context, account
}
// SendServiceUpdateToCluster indicates an expected call of SendServiceUpdateToCluster.
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr interface{}) *gomock.Call {
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendServiceUpdateToCluster", reflect.TypeOf((*MockController)(nil).SendServiceUpdateToCluster), ctx, accountID, update, clusterAddr)
}
@@ -322,7 +331,7 @@ func (m *MockController) UnregisterProxyFromCluster(ctx context.Context, cluster
}
// UnregisterProxyFromCluster indicates an expected call of UnregisterProxyFromCluster.
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnregisterProxyFromCluster", reflect.TypeOf((*MockController)(nil).UnregisterProxyFromCluster), ctx, clusterAddr, proxyID)
}
@@ -9,7 +9,7 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -1,6 +1,6 @@
package service
//go:generate go run github.com/golang/mock/mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
//go:generate go tool mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
import (
"context"
@@ -1,5 +1,10 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./interface.go
//
// Generated by this command:
//
// mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
//
// Package service is a generated GoMock package.
package service
@@ -8,14 +13,15 @@ import (
context "context"
reflect "reflect"
gomock "github.com/golang/mock/gomock"
proxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
gomock "go.uber.org/mock/gomock"
)
// MockManager is a mock of Manager interface.
type MockManager struct {
ctrl *gomock.Controller
recorder *MockManagerMockRecorder
isgomock struct{}
}
// MockManagerMockRecorder is the mock recorder for MockManager.
@@ -45,7 +51,7 @@ func (m *MockManager) CreateService(ctx context.Context, accountID, userID strin
}
// CreateService indicates an expected call of CreateService.
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateService", reflect.TypeOf((*MockManager)(nil).CreateService), ctx, accountID, userID, service)
}
@@ -60,7 +66,7 @@ func (m *MockManager) CreateServiceFromPeer(ctx context.Context, accountID, peer
}
// CreateServiceFromPeer indicates an expected call of CreateServiceFromPeer.
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateServiceFromPeer", reflect.TypeOf((*MockManager)(nil).CreateServiceFromPeer), ctx, accountID, peerID, req)
}
@@ -74,7 +80,7 @@ func (m *MockManager) DeleteAccountCluster(ctx context.Context, accountID, userI
}
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, accountID, userID, clusterAddress)
}
@@ -88,7 +94,7 @@ func (m *MockManager) DeleteAllServices(ctx context.Context, accountID, userID s
}
// DeleteAllServices indicates an expected call of DeleteAllServices.
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAllServices", reflect.TypeOf((*MockManager)(nil).DeleteAllServices), ctx, accountID, userID)
}
@@ -102,7 +108,7 @@ func (m *MockManager) DeleteService(ctx context.Context, accountID, userID, serv
}
// DeleteService indicates an expected call of DeleteService.
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteService", reflect.TypeOf((*MockManager)(nil).DeleteService), ctx, accountID, userID, serviceID)
}
@@ -117,7 +123,7 @@ func (m *MockManager) GetAccountServices(ctx context.Context, accountID string)
}
// GetAccountServices indicates an expected call of GetAccountServices.
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountServices", reflect.TypeOf((*MockManager)(nil).GetAccountServices), ctx, accountID)
}
@@ -132,7 +138,7 @@ func (m *MockManager) GetAllServices(ctx context.Context, accountID, userID stri
}
// GetAllServices indicates an expected call of GetAllServices.
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllServices", reflect.TypeOf((*MockManager)(nil).GetAllServices), ctx, accountID, userID)
}
@@ -147,7 +153,7 @@ func (m *MockManager) GetClusters(ctx context.Context, accountID, userID string)
}
// GetClusters indicates an expected call of GetClusters.
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusters", reflect.TypeOf((*MockManager)(nil).GetClusters), ctx, accountID, userID)
}
@@ -162,7 +168,7 @@ func (m *MockManager) GetGlobalServices(ctx context.Context) ([]*Service, error)
}
// GetGlobalServices indicates an expected call of GetGlobalServices.
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetGlobalServices", reflect.TypeOf((*MockManager)(nil).GetGlobalServices), ctx)
}
@@ -177,7 +183,7 @@ func (m *MockManager) GetService(ctx context.Context, accountID, userID, service
}
// GetService indicates an expected call of GetService.
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetService", reflect.TypeOf((*MockManager)(nil).GetService), ctx, accountID, userID, serviceID)
}
@@ -192,7 +198,7 @@ func (m *MockManager) GetServiceByDomain(ctx context.Context, domain string) (*S
}
// GetServiceByDomain indicates an expected call of GetServiceByDomain.
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByDomain", reflect.TypeOf((*MockManager)(nil).GetServiceByDomain), ctx, domain)
}
@@ -207,7 +213,7 @@ func (m *MockManager) GetServiceByID(ctx context.Context, accountID, serviceID s
}
// GetServiceByID indicates an expected call of GetServiceByID.
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByID", reflect.TypeOf((*MockManager)(nil).GetServiceByID), ctx, accountID, serviceID)
}
@@ -222,7 +228,7 @@ func (m *MockManager) GetServiceIDByTargetID(ctx context.Context, accountID, res
}
// GetServiceIDByTargetID indicates an expected call of GetServiceIDByTargetID.
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceIDByTargetID", reflect.TypeOf((*MockManager)(nil).GetServiceIDByTargetID), ctx, accountID, resourceID)
}
@@ -236,7 +242,7 @@ func (m *MockManager) ReloadAllServicesForAccount(ctx context.Context, accountID
}
// ReloadAllServicesForAccount indicates an expected call of ReloadAllServicesForAccount.
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadAllServicesForAccount", reflect.TypeOf((*MockManager)(nil).ReloadAllServicesForAccount), ctx, accountID)
}
@@ -250,7 +256,7 @@ func (m *MockManager) ReloadService(ctx context.Context, accountID, serviceID st
}
// ReloadService indicates an expected call of ReloadService.
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadService", reflect.TypeOf((*MockManager)(nil).ReloadService), ctx, accountID, serviceID)
}
@@ -264,7 +270,7 @@ func (m *MockManager) RenewServiceFromPeer(ctx context.Context, accountID, peerI
}
// RenewServiceFromPeer indicates an expected call of RenewServiceFromPeer.
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RenewServiceFromPeer", reflect.TypeOf((*MockManager)(nil).RenewServiceFromPeer), ctx, accountID, peerID, serviceID)
}
@@ -278,7 +284,7 @@ func (m *MockManager) SetCertificateIssuedAt(ctx context.Context, accountID, ser
}
// SetCertificateIssuedAt indicates an expected call of SetCertificateIssuedAt.
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetCertificateIssuedAt", reflect.TypeOf((*MockManager)(nil).SetCertificateIssuedAt), ctx, accountID, serviceID)
}
@@ -292,7 +298,7 @@ func (m *MockManager) SetStatus(ctx context.Context, accountID, serviceID string
}
// SetStatus indicates an expected call of SetStatus.
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetStatus", reflect.TypeOf((*MockManager)(nil).SetStatus), ctx, accountID, serviceID, status)
}
@@ -304,7 +310,7 @@ func (m *MockManager) StartExposeReaper(ctx context.Context) {
}
// StartExposeReaper indicates an expected call of StartExposeReaper.
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StartExposeReaper", reflect.TypeOf((*MockManager)(nil).StartExposeReaper), ctx)
}
@@ -318,7 +324,7 @@ func (m *MockManager) StopServiceFromPeer(ctx context.Context, accountID, peerID
}
// StopServiceFromPeer indicates an expected call of StopServiceFromPeer.
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StopServiceFromPeer", reflect.TypeOf((*MockManager)(nil).StopServiceFromPeer), ctx, accountID, peerID, serviceID)
}
@@ -333,7 +339,7 @@ func (m *MockManager) UpdateService(ctx context.Context, accountID, userID strin
}
// UpdateService indicates an expected call of UpdateService.
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service interface{}) *gomock.Call {
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateService", reflect.TypeOf((*MockManager)(nil).UpdateService), ctx, accountID, userID, service)
}
@@ -0,0 +1,127 @@
package manager
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/metric/noop"
domainmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain/manager"
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/mock_server"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/status"
)
const validationTestCluster = "eu.proxy.test"
// withRealDomainManager swaps the stub cluster deriver for the real domain
// manager backed by the same store, so service creation is gated by the actual
// domain rows rather than by a test double that always agrees.
func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
t.Helper()
ctx := context.Background()
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil)
require.NoError(t, err)
accountMgr := &mock_server.MockAccountManager{
StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {},
}
mgr.clusterDeriver = domainmanager.NewManager(testStore, proxyMgr, permissions.NewManager(testStore), accountMgr)
}
func newTestService(domain string) *rpservice.Service {
return &rpservice.Service{
Name: "test-service",
Domain: domain,
Enabled: true,
Mode: rpservice.ModeHTTP,
Targets: []*rpservice.Target{{
Host: "10.0.0.1",
Port: 8080,
Protocol: "http",
TargetId: testPeerID,
TargetType: "peer",
Enabled: true,
}},
}
}
// A service must not bind to a domain the account has not validated, and
// nothing may be persisted for the attempt.
func TestCreateService_RefusesUnvalidatedDomain(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false)
require.NoError(t, err)
_, err = mgr.CreateService(ctx, testAccountID, testUserID, newTestService("unproven.example.com"))
require.Error(t, err, "an unvalidated domain must not bind a service")
assert.Contains(t, err.Error(), "not validated", "the API error should name the actual problem")
sErr, ok := status.FromError(err)
require.True(t, ok, "error should be a typed status error")
assert.Equal(t, status.PreconditionFailed, sErr.Type())
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "no service row should be written for a refused domain")
}
// The negative control: a validated domain still binds a service and derives
// its cluster exactly as before.
func TestCreateService_ValidatedDomainBindsService(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
require.NoError(t, err)
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com"))
require.NoError(t, err)
assert.Equal(t, validationTestCluster, created.ProxyCluster, "service should bind to the domain's target cluster")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
require.Len(t, services, 1, "the service should be persisted")
assert.Equal(t, "app.proven.example.com", services[0].Domain)
}
// An update must not be a way around the creation gate: moving a live service
// onto an unvalidated domain has to fail rather than silently keep the old
// cluster and start serving the new hostname.
func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
require.NoError(t, err)
_, err = testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false)
require.NoError(t, err)
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com"))
require.NoError(t, err)
moved := *created
moved.Domain = "app.unproven.example.com"
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, &moved)
require.Error(t, err, "moving to an unvalidated domain must fail")
assert.Contains(t, err.Error(), "not validated")
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
require.NoError(t, err)
assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain")
}
@@ -6,7 +6,7 @@ import (
"testing"
"time"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -606,16 +606,19 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string,
return existing.ProxyCluster, nil
}
if m.clusterDeriver != nil {
derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain)
if err != nil {
log.WithError(err).Warnf("could not derive cluster from domain %s", svc.Domain)
} else {
return derived, nil
}
if m.clusterDeriver == nil {
return existing.ProxyCluster, nil
}
return existing.ProxyCluster, nil
// Falling back to the old cluster here would let an update move a service
// onto a domain the account has not validated, bypassing the check that
// creation makes.
derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain)
if err != nil {
return "", status.Errorf(status.PreconditionFailed, "could not derive cluster from domain %s: %v", svc.Domain, err)
}
return derived, nil
}
func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error {
@@ -7,11 +7,10 @@ import (
"testing"
"time"
cachestore "github.com/eko/gocache/lib/v4/store"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/metric/noop"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
@@ -31,7 +30,7 @@ import (
"github.com/netbirdio/netbird/shared/management/status"
)
func testCacheStore(t *testing.T) cachestore.StoreInterface {
func testCacheStore(t *testing.T) nbcache.Store {
t.Helper()
s, err := nbcache.NewStore(context.Background(), 30*time.Minute, 10*time.Minute, 100)
require.NoError(t, err)
@@ -295,6 +294,7 @@ func TestPersistNewService(t *testing.T) {
assert.Equal(t, status.AlreadyExists, sErr.Type())
})
}
func TestPreserveExistingAuthSecrets(t *testing.T) {
mgr := &Manager{}
@@ -55,6 +55,8 @@ const (
SourceEphemeral = "ephemeral"
)
var ErrUnsupportedIPAddressUpstreamHost = errors.New("unsupported ip address for a direct upstream host")
type TargetOptions struct {
SkipTLSVerify bool `json:"skip_tls_verify"`
RequestTimeout time.Duration `json:"request_timeout,omitempty"`
@@ -393,6 +395,7 @@ func (s *Service) ToProtoMapping(operation Operation, authToken string, oidcConf
if s.Auth.BearerAuth != nil && s.Auth.BearerAuth.Enabled {
auth.Oidc = true
auth.AllowedGroupIds = append([]string(nil), s.Auth.BearerAuth.DistributionGroups...)
}
for _, h := range s.Auth.HeaderAuths {
@@ -977,8 +980,8 @@ func (s *Service) validateHTTPTargets() error {
return err
}
case TargetTypeSubnet:
if target.Host == "" {
return fmt.Errorf("target %d has empty host but target_type is %q", i, target.TargetType)
if err := validateSubnetTarget(i, target); err != nil {
return err
}
case TargetTypeCluster:
if err := validateClusterTarget(i, target); err != nil {
@@ -1001,6 +1004,34 @@ func (s *Service) validateHTTPTargets() error {
return nil
}
func validateSubnetTarget(idx int, target *Target) error {
host := strings.TrimSpace(target.Host)
if host == "" {
return fmt.Errorf("target %d has empty host but target_type is %q", idx, target.TargetType)
}
if strings.ContainsAny(host, " \t/") {
return fmt.Errorf("target %d: host %q contains invalid characters", idx, host)
}
if _, _, err := net.SplitHostPort(host); err == nil {
return fmt.Errorf("target %d: host %q must not include a port (set target.port instead)", idx, host)
}
noBrackets := strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
maybeip, err := netip.ParseAddr(noBrackets)
if err != nil { // not an ip
return nil //nolint:nilerr
}
if maybeip.Zone() != "" {
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
}
if !target.Options.DirectUpstream {
return nil
}
if maybeip.IsLoopback() || maybeip.IsMulticast() || maybeip.IsLinkLocalUnicast() {
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
}
return nil
}
// validateClusterTarget cluster targets should not have empty hosts and should have direct upstream enabled.
func validateClusterTarget(idx int, target *Target) error {
host := strings.TrimSpace(target.Host)
@@ -1035,6 +1066,15 @@ func validateDirectUpstreamHost(idx int, target *Target) error {
if _, _, err := net.SplitHostPort(host); err == nil {
return fmt.Errorf("target %d: host %q must not include a port (set target.port instead)", idx, host)
}
noBrackets := strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
maybeip, err := netip.ParseAddr(noBrackets)
if err != nil { // not an ip
return nil //nolint:nilerr
}
if maybeip.Zone() != "" || maybeip.IsLoopback() || maybeip.IsMulticast() || maybeip.IsLinkLocalUnicast() {
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
}
return nil
}
@@ -216,6 +216,64 @@ func TestValidateTargetOptions_CustomHeaders(t *testing.T) {
})
}
func TestValidate_DirectUpstreamHost(t *testing.T) {
target := Target{TargetId: "id-1", TargetType: TargetTypePeer, Host: "10.0.0.1", Port: 80, Protocol: "http", Enabled: true, Options: TargetOptions{DirectUpstream: true}}
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "127.0.0.2")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "127.0.0.2:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "::1")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "::1%lo0")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1]")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1]:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1%lo0]")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1%lo0]:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "169.254.100.100")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "fe80::1")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[fe80::1]")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "224.100.100.100")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "ff00::ffff")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[ff00::ffff]")), ErrUnsupportedIPAddressUpstreamHost)
// empty host
assert.Nil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: " "}))
// host with a space
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with space"}))
// host with a tab
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with\ttab"}))
// host with a slash
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with/slash"}))
}
func TestValidate_ValidateSubnetTarget(t *testing.T) {
target := Target{TargetId: "id-1", TargetType: TargetTypeSubnet, Host: "10.0.0.1", Port: 80, Protocol: "http", Enabled: true, Options: TargetOptions{DirectUpstream: true}}
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "127.0.0.2")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "127.0.0.2:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "::1")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "[::1]:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "::1%lo0")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[::1%lo0]")), ErrUnsupportedIPAddressUpstreamHost)
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "[::1%lo0]:80")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "169.254.100.100")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "fe80::1")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[fe80::1]")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "224.100.100.100")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "ff00::ffff")), ErrUnsupportedIPAddressUpstreamHost)
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[ff00::ffff]")), ErrUnsupportedIPAddressUpstreamHost)
// empty host
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: " "}))
// host with a space
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with space"}))
// host with a tab
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with\ttab"}))
// host with a slash
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with/slash"}))
}
func targetWithHost(t *Target, host string) *Target {
t.Host = host
return t
}
func TestToProtoMapping_TargetOptions(t *testing.T) {
rp := &Service{
ID: "svc-1",
@@ -250,6 +308,44 @@ func TestToProtoMapping_TargetOptions(t *testing.T) {
assert.Equal(t, int64(30), opts.RequestTimeout.Seconds)
}
// TestToProtoMapping_AllowedGroupIds covers the list the proxy gates session
// cookies on: without it the proxy can only check a cookie's signature, which
// makes a token minted for a user outside the groups a bearer credential.
func TestToProtoMapping_AllowedGroupIds(t *testing.T) {
t.Run("distribution groups reach the proxy", func(t *testing.T) {
rp := &Service{
ID: "svc-1",
AccountID: "acc-1",
Domain: "example.com",
Auth: AuthConfig{
BearerAuth: &BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{"grp-1", "grp-2"},
},
},
}
pm := rp.ToProtoMapping(Create, "token", proxy.OIDCValidationConfig{})
assert.True(t, pm.GetAuth().GetOidc())
assert.Equal(t, []string{"grp-1", "grp-2"}, pm.GetAuth().GetAllowedGroupIds())
})
t.Run("a service open to the account carries no groups", func(t *testing.T) {
rp := &Service{
ID: "svc-1",
AccountID: "acc-1",
Domain: "example.com",
Auth: AuthConfig{
BearerAuth: &BearerAuthConfig{Enabled: true},
},
}
pm := rp.ToProtoMapping(Create, "token", proxy.OIDCValidationConfig{})
assert.True(t, pm.GetAuth().GetOidc())
assert.Empty(t, pm.GetAuth().GetAllowedGroupIds(), "an empty list must not restrict access")
})
}
func TestToProtoMapping_NoOptionsWhenDefault(t *testing.T) {
rp := &Service{
ID: "svc-1",
@@ -5,7 +5,7 @@ import (
"fmt"
"testing"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -4,7 +4,7 @@ import (
"context"
"testing"
"github.com/golang/mock/gomock"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"