mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 13:39:07 +02:00
Merge branch 'main' into fix/pkce-flow-session-extend
# Conflicts: # client/ios/NetBirdSDK/login.go # client/server/server.go # shared/management/proto/management.pb.go
This commit is contained in:
@@ -511,6 +511,7 @@ func (c *Controller) fetchNetworkMapData(ctx context.Context, accountID string)
|
||||
}
|
||||
|
||||
nmData.Services = c.proxyServicesFromRepo(ctx, accountID)
|
||||
nmData.BuildPrivateServiceCandidates()
|
||||
nmData.InjectProxyPolicies()
|
||||
nmData.PrecomputePostureValidation()
|
||||
|
||||
@@ -566,39 +567,38 @@ func NetworkMapFromData(ctx context.Context, nmData *networkmap.NetworkMapData,
|
||||
return nm
|
||||
}
|
||||
|
||||
// peerPostureChecksFromData mirrors getPeerPostureChecks on the twin store. The
|
||||
// sync response only encodes process-check file paths, so only ProcessCheck is
|
||||
// converted back to the server posture type.
|
||||
func peerPostureChecksFromData(nmData *networkmap.NetworkMapData, peerID string) []*posture.Checks {
|
||||
// peerPostureChecksFromData mirrors getPeerPostureChecks on the twin store.
|
||||
func peerPostureChecksFromData(nmData *networkmap.NetworkMapData, peerID string) []*nmdata.PostureChecks {
|
||||
if len(nmData.PostureChecks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
peerPostureChecks := make(map[string]*posture.Checks)
|
||||
peerPostureChecks := make(map[string]*nmdata.PostureChecks)
|
||||
for _, policy := range nmData.Policies {
|
||||
if policy == nil || !policy.Enabled || len(policy.SourcePostureChecks) == 0 {
|
||||
continue
|
||||
}
|
||||
if !isPeerInPolicySourceGroupsFromData(nmData, peerID, policy) {
|
||||
if !isPeerInPolicySourcesFromData(nmData, peerID, policy) {
|
||||
continue
|
||||
}
|
||||
for _, checkID := range policy.SourcePostureChecks {
|
||||
twin := nmData.PostureChecks[checkID]
|
||||
if twin == nil {
|
||||
continue
|
||||
if twin := nmData.PostureChecks[checkID]; twin != nil {
|
||||
peerPostureChecks[checkID] = twin
|
||||
}
|
||||
peerPostureChecks[checkID] = postureChecksFromTwin(twin)
|
||||
}
|
||||
}
|
||||
|
||||
return maps.Values(peerPostureChecks)
|
||||
}
|
||||
|
||||
func isPeerInPolicySourceGroupsFromData(nmData *networkmap.NetworkMapData, peerID string, policy *nmdata.Policy) bool {
|
||||
func isPeerInPolicySourcesFromData(nmData *networkmap.NetworkMapData, peerID string, policy *nmdata.Policy) bool {
|
||||
for _, rule := range policy.Rules {
|
||||
if rule == nil || !rule.Enabled {
|
||||
continue
|
||||
}
|
||||
if rule.SourceResource.Type == string(types.ResourceTypePeer) && rule.SourceResource.ID == peerID {
|
||||
return true
|
||||
}
|
||||
for _, groupID := range rule.Sources {
|
||||
if group := nmData.Groups[groupID]; group != nil && slices.Contains(group.Peers, peerID) {
|
||||
return true
|
||||
@@ -608,18 +608,6 @@ func isPeerInPolicySourceGroupsFromData(nmData *networkmap.NetworkMapData, peerI
|
||||
return false
|
||||
}
|
||||
|
||||
func postureChecksFromTwin(twin *nmdata.PostureChecks) *posture.Checks {
|
||||
checks := &posture.Checks{ID: twin.ID}
|
||||
if twin.Checks.ProcessCheck != nil {
|
||||
processes := make([]posture.Process, 0, len(twin.Checks.ProcessCheck.Processes))
|
||||
for _, p := range twin.Checks.ProcessCheck.Processes {
|
||||
processes = append(processes, posture.Process{LinuxPath: p.LinuxPath, MacPath: p.MacPath, WindowsPath: p.WindowsPath})
|
||||
}
|
||||
checks.Checks.ProcessCheck = &posture.ProcessCheck{Processes: processes}
|
||||
}
|
||||
return checks
|
||||
}
|
||||
|
||||
func (c *Controller) perAccountOrGlobalSupportedSyncMessageVersions(accountId string) sharedgrpc.SyncMessageVersion {
|
||||
if perAccount, ok := c.perAccountServerSupportedSyncMessageVersions[accountId]; ok {
|
||||
return perAccount
|
||||
@@ -967,7 +955,7 @@ func (c *Controller) BufferUpdateAccountPeers(ctx context.Context, accountID str
|
||||
// data the legacy server folds in via NetworkMap.Merge). The gRPC layer
|
||||
// encodes both into the wire envelope. Callers must gate on capability
|
||||
// themselves before dispatching here — this method does NOT branch on it.
|
||||
func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, peer *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, peer *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
if isRequiresApproval {
|
||||
network, err := c.repo.GetAccountNetwork(ctx, accountID)
|
||||
if err != nil {
|
||||
@@ -1032,7 +1020,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
|
||||
// getValidatedPeerWithComponentsFromData is the account-free variant of
|
||||
// GetValidatedPeerWithComponents. The proxy network map fragment is omitted
|
||||
// like on the other nmdata paths.
|
||||
func (c *Controller) getValidatedPeerWithComponentsFromData(ctx context.Context, accountID string, peer *nbpeer.Peer, nmData *networkmap.NetworkMapData) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (c *Controller) getValidatedPeerWithComponentsFromData(ctx context.Context, accountID string, peer *nbpeer.Peer, nmData *networkmap.NetworkMapData) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
postureChecks := peerPostureChecksFromData(nmData, peer.ID)
|
||||
|
||||
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
|
||||
@@ -1142,7 +1130,7 @@ func (b *bufferAffectedUpdate) setTimer(d time.Duration, f func()) {
|
||||
b.next.Reset(d)
|
||||
}
|
||||
|
||||
func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, peerID string) (*types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, peerID string) (*types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
if isRequiresApproval {
|
||||
network, err := c.repo.GetAccountNetwork(ctx, accountID)
|
||||
if err != nil {
|
||||
@@ -1209,7 +1197,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
|
||||
// getValidatedPeerWithMapFromData is the account-free variant of
|
||||
// GetValidatedPeerWithMap. The proxy network map fragment is omitted like on
|
||||
// the other nmdata paths.
|
||||
func (c *Controller) getValidatedPeerWithMapFromData(ctx context.Context, accountID string, peerID string, nmData *networkmap.NetworkMapData) (*types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (c *Controller) getValidatedPeerWithMapFromData(ctx context.Context, accountID string, peerID string, nmData *networkmap.NetworkMapData) (*types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
postureChecks := peerPostureChecksFromData(nmData, peerID)
|
||||
|
||||
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
|
||||
@@ -1234,7 +1222,7 @@ func (c *Controller) GetDNSDomain(settings *types.Settings) string {
|
||||
}
|
||||
|
||||
// getPeerPostureChecks returns the posture checks applied for a given peer.
|
||||
func (c *Controller) getPeerPostureChecks(account *types.Account, peerID string) ([]*posture.Checks, error) {
|
||||
func (c *Controller) getPeerPostureChecks(account *types.Account, peerID string) ([]*nmdata.PostureChecks, error) {
|
||||
peerPostureChecks := make(map[string]*posture.Checks)
|
||||
|
||||
if len(account.PostureChecks) == 0 {
|
||||
@@ -1251,7 +1239,7 @@ func (c *Controller) getPeerPostureChecks(account *types.Account, peerID string)
|
||||
}
|
||||
}
|
||||
|
||||
return maps.Values(peerPostureChecks), nil
|
||||
return types.TwinPostureChecksList(maps.Values(peerPostureChecks)), nil
|
||||
}
|
||||
|
||||
func (c *Controller) StartWarmup(ctx context.Context) {
|
||||
@@ -1330,7 +1318,7 @@ func computeForwarderPortFromVersions(wtVersions []string, requiredVersion strin
|
||||
|
||||
// addPolicyPostureChecks adds posture checks from a policy to the peer posture checks map if the peer is in the policy's source groups.
|
||||
func addPolicyPostureChecks(account *types.Account, peerID string, policy *types.Policy, peerPostureChecks map[string]*posture.Checks) error {
|
||||
isInGroup, err := isPeerInPolicySourceGroups(account, peerID, policy)
|
||||
isInGroup, err := isPeerInPolicySources(account, peerID, policy)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -1350,13 +1338,17 @@ func addPolicyPostureChecks(account *types.Account, peerID string, policy *types
|
||||
return nil
|
||||
}
|
||||
|
||||
// isPeerInPolicySourceGroups checks if a peer is present in any of the policy rule source groups.
|
||||
func isPeerInPolicySourceGroups(account *types.Account, peerID string, policy *types.Policy) (bool, error) {
|
||||
// isPeerInPolicySources checks if a peer is a source of the policy, directly or through a source group.
|
||||
func isPeerInPolicySources(account *types.Account, peerID string, policy *types.Policy) (bool, error) {
|
||||
for _, rule := range policy.Rules {
|
||||
if !rule.Enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID == peerID {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
for _, sourceGroup := range rule.Sources {
|
||||
group := account.GetGroup(sourceGroup)
|
||||
if group == nil {
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
|
||||
func postureSelectionData(policies ...*nmdata.Policy) *networkmap.NetworkMapData {
|
||||
return &networkmap.NetworkMapData{
|
||||
Groups: map[string]*nmdata.Group{"g-src": {ID: "g-src", Peers: []string{"peer-group"}}},
|
||||
Policies: policies,
|
||||
PostureChecks: map[string]*nmdata.PostureChecks{
|
||||
"pc1": {ID: "pc1", Checks: nmdata.ChecksDefinition{NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "0.30.0"}}},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func gatedPolicy(id string, rule *nmdata.PolicyRule, checkIDs ...string) *nmdata.Policy {
|
||||
return &nmdata.Policy{ID: id, Enabled: true, SourcePostureChecks: checkIDs, Rules: []*nmdata.PolicyRule{rule}}
|
||||
}
|
||||
|
||||
func checkIDs(checks []*nmdata.PostureChecks) []string {
|
||||
ids := make([]string, 0, len(checks))
|
||||
for _, c := range checks {
|
||||
ids = append(ids, c.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func TestPeerPostureChecksFromData_SelectsPolicySourcePeers(t *testing.T) {
|
||||
groupRule := &nmdata.PolicyRule{ID: "r-group", Enabled: true, Sources: []string{"g-src"}, Destinations: []string{"g-dst"}}
|
||||
directRule := &nmdata.PolicyRule{ID: "r-direct", Enabled: true, SourceResource: nmdata.Resource{ID: "peer-direct", Type: string(types.ResourceTypePeer)}, Destinations: []string{"g-dst"}}
|
||||
|
||||
t.Run("source group member and direct source peer both get the checks", func(t *testing.T) {
|
||||
nmData := postureSelectionData(gatedPolicy("p1", groupRule, "pc1"), gatedPolicy("p2", directRule, "pc1"))
|
||||
|
||||
assert.Equal(t, []string{"pc1"}, checkIDs(peerPostureChecksFromData(nmData, "peer-group")))
|
||||
assert.Equal(t, []string{"pc1"}, checkIDs(peerPostureChecksFromData(nmData, "peer-direct")))
|
||||
assert.Empty(t, peerPostureChecksFromData(nmData, "peer-elsewhere"))
|
||||
})
|
||||
|
||||
t.Run("source resource of a non-peer type never matches a peer", func(t *testing.T) {
|
||||
hostRule := &nmdata.PolicyRule{ID: "r-host", Enabled: true, SourceResource: nmdata.Resource{ID: "peer-direct", Type: string(types.ResourceTypeHost)}, Destinations: []string{"g-dst"}}
|
||||
nmData := postureSelectionData(gatedPolicy("p1", hostRule, "pc1"))
|
||||
|
||||
assert.Empty(t, peerPostureChecksFromData(nmData, "peer-direct"))
|
||||
})
|
||||
|
||||
t.Run("same check through two policies is returned once", func(t *testing.T) {
|
||||
nmData := postureSelectionData(gatedPolicy("p1", groupRule, "pc1"), gatedPolicy("p2", groupRule, "pc1"))
|
||||
|
||||
assert.Equal(t, []string{"pc1"}, checkIDs(peerPostureChecksFromData(nmData, "peer-group")))
|
||||
})
|
||||
|
||||
t.Run("disabled policy, disabled rule and dangling check are ignored", func(t *testing.T) {
|
||||
disabledPolicy := gatedPolicy("p-off", groupRule, "pc1")
|
||||
disabledPolicy.Enabled = false
|
||||
disabledRule := &nmdata.PolicyRule{ID: "r-off", Enabled: false, Sources: []string{"g-src"}}
|
||||
nmData := postureSelectionData(disabledPolicy, gatedPolicy("p-rule-off", disabledRule, "pc1"), gatedPolicy("p-dangling", groupRule, "pc-missing"))
|
||||
|
||||
assert.Empty(t, peerPostureChecksFromData(nmData, "peer-group"))
|
||||
})
|
||||
}
|
||||
@@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (*
|
||||
}
|
||||
|
||||
func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) {
|
||||
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
}
|
||||
|
||||
func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) {
|
||||
|
||||
@@ -7,8 +7,8 @@ import (
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -23,8 +23,8 @@ type Controller interface {
|
||||
BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) error
|
||||
UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error
|
||||
BufferUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error
|
||||
GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, peerID string) (*types.NetworkMap, []*posture.Checks, int64, error)
|
||||
GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, p *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*posture.Checks, int64, error)
|
||||
GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, peerID string) (*types.NetworkMap, []*nmdata.PostureChecks, int64, error)
|
||||
GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, p *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*nmdata.PostureChecks, int64, error)
|
||||
GetDNSDomain(settings *types.Settings) string
|
||||
StartWarmup(context.Context)
|
||||
GetNetworkMap(ctx context.Context, peerID string) (*types.NetworkMap, error)
|
||||
|
||||
@@ -14,8 +14,8 @@ import (
|
||||
reflect "reflect"
|
||||
|
||||
peer "github.com/netbirdio/netbird/management/server/peer"
|
||||
posture "github.com/netbirdio/netbird/management/server/posture"
|
||||
types "github.com/netbirdio/netbird/management/server/types"
|
||||
nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
@@ -127,13 +127,13 @@ func (mr *MockControllerMockRecorder) GetNetworkMap(ctx, peerID any) *gomock.Cal
|
||||
}
|
||||
|
||||
// GetValidatedPeerWithComponents mocks base method.
|
||||
func (m *MockController) GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, p *peer.Peer) (*peer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (m *MockController) GetValidatedPeerWithComponents(ctx context.Context, isRequiresApproval bool, accountID string, p *peer.Peer) (*peer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetValidatedPeerWithComponents", ctx, isRequiresApproval, accountID, p)
|
||||
ret0, _ := ret[0].(*peer.Peer)
|
||||
ret1, _ := ret[1].(*types.NetworkMapComponents)
|
||||
ret2, _ := ret[2].(*types.NetworkMap)
|
||||
ret3, _ := ret[3].([]*posture.Checks)
|
||||
ret3, _ := ret[3].([]*nmdata.PostureChecks)
|
||||
ret4, _ := ret[4].(int64)
|
||||
ret5, _ := ret[5].(error)
|
||||
return ret0, ret1, ret2, ret3, ret4, ret5
|
||||
@@ -146,11 +146,11 @@ func (mr *MockControllerMockRecorder) GetValidatedPeerWithComponents(ctx, isRequ
|
||||
}
|
||||
|
||||
// GetValidatedPeerWithMap mocks base method.
|
||||
func (m *MockController) GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID, peerID string) (*types.NetworkMap, []*posture.Checks, int64, error) {
|
||||
func (m *MockController) GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID, peerID string) (*types.NetworkMap, []*nmdata.PostureChecks, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetValidatedPeerWithMap", ctx, isRequiresApproval, accountID, peerID)
|
||||
ret0, _ := ret[0].(*types.NetworkMap)
|
||||
ret1, _ := ret[1].([]*posture.Checks)
|
||||
ret1, _ := ret[1].([]*nmdata.PostureChecks)
|
||||
ret2, _ := ret[2].(int64)
|
||||
ret3, _ := ret[3].(error)
|
||||
return ret0, ret1, ret2, ret3
|
||||
|
||||
@@ -183,6 +183,7 @@ func RunCase(t *testing.T, c Case) {
|
||||
ctx := context.Background()
|
||||
nmData := c.Data
|
||||
applyFixtureDefaults(nmData)
|
||||
nmData.BuildPrivateServiceCandidates()
|
||||
nmData.PrecomputePostureValidation()
|
||||
|
||||
dnsDomain := c.DNSDomain
|
||||
@@ -244,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma
|
||||
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
|
||||
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
|
||||
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
|
||||
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
|
||||
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false)
|
||||
require.NoError(t, err, "expand envelope")
|
||||
return res.NetworkMap
|
||||
default:
|
||||
|
||||
+5
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"description": "A peer named directly as a rule source or destination is subject to approval exactly like a group member: unvalidated peer-b is neither a source for peer-c nor a destination for peer-a, while the validated direct source peer-a reaches peer-c. Legacy mode is excluded: the frozen legacynmap copy still carries the direct-peer bypass.",
|
||||
"peers": ["peer-a", "peer-c"],
|
||||
"modes": ["full", "envelope"]
|
||||
}
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
{
|
||||
"Serial": "22",
|
||||
"peerConfig": {
|
||||
"address": "100.64.0.1/10",
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-a.netbird.test",
|
||||
"autoUpdate": {}
|
||||
},
|
||||
"remotePeers": [
|
||||
{
|
||||
"wgPubKey": "4deEImv8zGvsyBmmfC2G0eQkbyMzyGuz/YK7pcYETwM=",
|
||||
"allowedIps": [
|
||||
"100.64.0.3/32"
|
||||
],
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-c.netbird.test",
|
||||
"agentVersion": "0.60.0"
|
||||
}
|
||||
],
|
||||
"DNSConfig": {
|
||||
"ServiceEnable": true,
|
||||
"CustomZones": [
|
||||
{
|
||||
"Domain": "netbird.test.",
|
||||
"Records": [
|
||||
{
|
||||
"Name": "peer-a.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.1"
|
||||
},
|
||||
{
|
||||
"Name": "peer-c.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.3"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"ForwarderPort": "22054"
|
||||
},
|
||||
"FirewallRules": [
|
||||
{
|
||||
"PeerIP": "100.64.0.3",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
},
|
||||
{
|
||||
"PeerIP": "100.64.0.3",
|
||||
"Direction": "OUT",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
}
|
||||
],
|
||||
"routesFirewallRulesIsEmpty": true,
|
||||
"sshAuth": {
|
||||
"UserIDClaim": "sub"
|
||||
}
|
||||
}
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
{
|
||||
"Serial": "22",
|
||||
"peerConfig": {
|
||||
"address": "100.64.0.3/10",
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-c.netbird.test",
|
||||
"autoUpdate": {}
|
||||
},
|
||||
"remotePeers": [
|
||||
{
|
||||
"wgPubKey": "vblMc9U8RAI6cVopcKEMTVT6lVC3D9nTTMSwot5d3L4=",
|
||||
"allowedIps": [
|
||||
"100.64.0.1/32"
|
||||
],
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-a.netbird.test",
|
||||
"agentVersion": "0.60.0"
|
||||
}
|
||||
],
|
||||
"DNSConfig": {
|
||||
"ServiceEnable": true,
|
||||
"CustomZones": [
|
||||
{
|
||||
"Domain": "netbird.test.",
|
||||
"Records": [
|
||||
{
|
||||
"Name": "peer-a.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.1"
|
||||
},
|
||||
{
|
||||
"Name": "peer-c.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.3"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"ForwarderPort": "22054"
|
||||
},
|
||||
"FirewallRules": [
|
||||
{
|
||||
"PeerIP": "100.64.0.1",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
},
|
||||
{
|
||||
"PeerIP": "100.64.0.1",
|
||||
"Direction": "OUT",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
}
|
||||
],
|
||||
"routesFirewallRulesIsEmpty": true,
|
||||
"sshAuth": {
|
||||
"UserIDClaim": "sub"
|
||||
}
|
||||
}
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
{
|
||||
"Network": {"Serial": 22},
|
||||
"Peers": {
|
||||
"peer-a": {"IP": "100.64.0.1", "Meta": {"WtVersion": "0.60.0"}},
|
||||
"peer-b": {"IP": "100.64.0.2", "Meta": {"WtVersion": "0.60.0"}},
|
||||
"peer-c": {"IP": "100.64.0.3", "Meta": {"WtVersion": "0.60.0"}}
|
||||
},
|
||||
"ValidatedPeers": {"peer-a": {}, "peer-c": {}},
|
||||
"Groups": {
|
||||
"grp-dev": {"Peers": ["peer-a"]},
|
||||
"grp-ops": {"Peers": ["peer-c"]}
|
||||
},
|
||||
"Policies": [
|
||||
{
|
||||
"ID": "pol-direct-ok",
|
||||
"PublicID": "pol-direct-ok-pub",
|
||||
"Enabled": true,
|
||||
"Rules": [
|
||||
{
|
||||
"Enabled": true,
|
||||
"Action": "accept",
|
||||
"Protocol": "tcp",
|
||||
"Ports": ["443"],
|
||||
"Bidirectional": true,
|
||||
"SourceResource": {"ID": "peer-a", "Type": "peer"},
|
||||
"Destinations": ["grp-ops"]
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"ID": "pol-src-unval",
|
||||
"PublicID": "pol-src-unval-pub",
|
||||
"Enabled": true,
|
||||
"Rules": [
|
||||
{
|
||||
"Enabled": true,
|
||||
"Action": "accept",
|
||||
"Protocol": "tcp",
|
||||
"Ports": ["8443"],
|
||||
"Bidirectional": true,
|
||||
"SourceResource": {"ID": "peer-b", "Type": "peer"},
|
||||
"Destinations": ["grp-ops"]
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"ID": "pol-dst-unval",
|
||||
"PublicID": "pol-dst-unval-pub",
|
||||
"Enabled": true,
|
||||
"Rules": [
|
||||
{
|
||||
"Enabled": true,
|
||||
"Action": "accept",
|
||||
"Protocol": "tcp",
|
||||
"Ports": ["9443"],
|
||||
"Bidirectional": true,
|
||||
"Sources": ["grp-dev"],
|
||||
"DestinationResource": {"ID": "peer-b", "Type": "peer"}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
management/internals/controllers/network_map/nmaptest/testdata/cases/posture-direct-source/case.json
Vendored
+5
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"description": "A peer named directly as a rule source is gated by the policy's posture checks exactly like a group member: peer-b (0.40.0) fails the 0.45.0 minimum, so it gets no connectivity and peer-c must not see it, while the compliant direct source peer-a reaches peer-c. Legacy mode is excluded: the frozen legacynmap copy still carries the direct-peer bypass.",
|
||||
"peers": ["peer-b", "peer-c"],
|
||||
"modes": ["full", "envelope"]
|
||||
}
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"Serial": "21",
|
||||
"peerConfig": {
|
||||
"address": "100.64.0.2/10",
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-b.netbird.test",
|
||||
"autoUpdate": {}
|
||||
},
|
||||
"remotePeersIsEmpty": true,
|
||||
"DNSConfig": {
|
||||
"ServiceEnable": true,
|
||||
"CustomZones": [
|
||||
{
|
||||
"Domain": "netbird.test.",
|
||||
"Records": [
|
||||
{
|
||||
"Name": "peer-b.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.2"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"ForwarderPort": "5353"
|
||||
},
|
||||
"firewallRulesIsEmpty": true,
|
||||
"routesFirewallRulesIsEmpty": true,
|
||||
"sshAuth": {
|
||||
"UserIDClaim": "sub"
|
||||
}
|
||||
}
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
{
|
||||
"Serial": "21",
|
||||
"peerConfig": {
|
||||
"address": "100.64.0.3/10",
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-c.netbird.test",
|
||||
"autoUpdate": {}
|
||||
},
|
||||
"remotePeers": [
|
||||
{
|
||||
"wgPubKey": "vblMc9U8RAI6cVopcKEMTVT6lVC3D9nTTMSwot5d3L4=",
|
||||
"allowedIps": [
|
||||
"100.64.0.1/32"
|
||||
],
|
||||
"sshConfig": {},
|
||||
"fqdn": "peer-a.netbird.test",
|
||||
"agentVersion": "0.60.0"
|
||||
}
|
||||
],
|
||||
"DNSConfig": {
|
||||
"ServiceEnable": true,
|
||||
"CustomZones": [
|
||||
{
|
||||
"Domain": "netbird.test.",
|
||||
"Records": [
|
||||
{
|
||||
"Name": "peer-a.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.1"
|
||||
},
|
||||
{
|
||||
"Name": "peer-c.netbird.test",
|
||||
"Type": "1",
|
||||
"Class": "IN",
|
||||
"TTL": "300",
|
||||
"RData": "100.64.0.3"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"ForwarderPort": "5353"
|
||||
},
|
||||
"FirewallRules": [
|
||||
{
|
||||
"PeerIP": "100.64.0.1",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
},
|
||||
{
|
||||
"PeerIP": "100.64.0.1",
|
||||
"Direction": "OUT",
|
||||
"Protocol": "TCP",
|
||||
"Port": "443",
|
||||
"PolicyID": "cG9sLWRpcmVjdC1vaw=="
|
||||
}
|
||||
],
|
||||
"routesFirewallRulesIsEmpty": true,
|
||||
"sshAuth": {
|
||||
"UserIDClaim": "sub"
|
||||
}
|
||||
}
|
||||
+51
@@ -0,0 +1,51 @@
|
||||
{
|
||||
"Network": {"Serial": 21},
|
||||
"Peers": {
|
||||
"peer-a": {"IP": "100.64.0.1", "Meta": {"WtVersion": "0.60.0"}},
|
||||
"peer-b": {"IP": "100.64.0.2", "Meta": {"WtVersion": "0.40.0"}},
|
||||
"peer-c": {"IP": "100.64.0.3", "Meta": {"WtVersion": "0.60.0"}}
|
||||
},
|
||||
"Groups": {
|
||||
"grp-ops": {"Peers": ["peer-c"]}
|
||||
},
|
||||
"PostureChecks": {
|
||||
"chk-ver": {"Checks": {"NBVersionCheck": {"MinVersion": "0.45.0"}}}
|
||||
},
|
||||
"PostureCheckXIDToPublicID": {"chk-ver": "chk-ver-pub"},
|
||||
"Policies": [
|
||||
{
|
||||
"ID": "pol-direct-ok",
|
||||
"PublicID": "pol-direct-ok-pub",
|
||||
"Enabled": true,
|
||||
"SourcePostureChecks": ["chk-ver"],
|
||||
"Rules": [
|
||||
{
|
||||
"Enabled": true,
|
||||
"Action": "accept",
|
||||
"Protocol": "tcp",
|
||||
"Ports": ["443"],
|
||||
"Bidirectional": true,
|
||||
"SourceResource": {"ID": "peer-a", "Type": "peer"},
|
||||
"Destinations": ["grp-ops"]
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"ID": "pol-direct-denied",
|
||||
"PublicID": "pol-direct-denied-pub",
|
||||
"Enabled": true,
|
||||
"SourcePostureChecks": ["chk-ver"],
|
||||
"Rules": [
|
||||
{
|
||||
"Enabled": true,
|
||||
"Action": "accept",
|
||||
"Protocol": "tcp",
|
||||
"Ports": ["8443"],
|
||||
"Bidirectional": true,
|
||||
"SourceResource": {"ID": "peer-b", "Type": "peer"},
|
||||
"Destinations": ["grp-ops"]
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
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/store"
|
||||
nbtypes "github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
// TestCleanupAccessLogs_RealStore_DeletedAccount covers a deleted account's access logs.
|
||||
// The sweep is driven by settings rows, which go with the account, so without a fallback
|
||||
// those logs would never expire. They get the default retention instead. A live account
|
||||
// can delete its own settings row, so "no settings" must not be mistaken for "deleted":
|
||||
// that account's logs are left alone, as are those of an account that keeps logs forever.
|
||||
func TestCleanupAccessLogs_RealStore_DeletedAccount(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s, cleanup, err := store.NewTestStoreFromSQL(ctx, "", t.TempDir())
|
||||
require.NoError(t, err, "real sqlite test store must come up")
|
||||
defer cleanup()
|
||||
|
||||
const (
|
||||
deletedAccountID = "acc-deleted"
|
||||
keepAccountID = "acc-keep-forever"
|
||||
noSettingsAccountID = "acc-live-no-settings"
|
||||
)
|
||||
old := time.Now().UTC().AddDate(0, 0, -(types.DefaultAccessLogRetentionDays + 10))
|
||||
recent := time.Now().UTC().AddDate(0, 0, -1)
|
||||
|
||||
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: keepAccountID}))
|
||||
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: noSettingsAccountID}))
|
||||
|
||||
keepSettings := types.DefaultSettings(keepAccountID)
|
||||
keepSettings.Domain = "keep.gw.example.com"
|
||||
keepSettings.AccessLogRetentionDays = 0
|
||||
require.NoError(t, s.SaveAgentNetworkSettings(ctx, keepSettings))
|
||||
|
||||
mkLog := func(id, accountID string, ts time.Time) {
|
||||
t.Helper()
|
||||
entry := &types.AgentNetworkAccessLog{
|
||||
ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts, StatusCode: 200, Model: "gpt-4o",
|
||||
}
|
||||
groups := []types.AgentNetworkAccessLogGroup{{LogID: id, GroupID: "grp-eng", AccountID: accountID}}
|
||||
require.NoError(t, s.CreateAgentNetworkAccessLog(ctx, entry, groups))
|
||||
}
|
||||
mkLog("deleted-old", deletedAccountID, old)
|
||||
mkLog("deleted-recent", deletedAccountID, recent)
|
||||
mkLog("keep-old", keepAccountID, old)
|
||||
mkLog("no-settings-old", noSettingsAccountID, old)
|
||||
|
||||
m := &managerImpl{store: s}
|
||||
m.cleanupAccessLogsOnce(ctx)
|
||||
|
||||
logIDs := func(accountID string) []string {
|
||||
t.Helper()
|
||||
logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID,
|
||||
types.AgentNetworkAccessLogFilter{Page: 1, PageSize: 50})
|
||||
require.NoError(t, err)
|
||||
ids := make([]string, 0, len(logs))
|
||||
for _, l := range logs {
|
||||
ids = append(ids, l.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
assert.Equal(t, []string{"deleted-recent"}, logIDs(deletedAccountID),
|
||||
"a deleted account should have logs past the default retention swept")
|
||||
assert.Equal(t, []string{"keep-old"}, logIDs(keepAccountID),
|
||||
"an account with retention disabled should keep its old logs")
|
||||
assert.Equal(t, []string{"no-settings-old"}, logIDs(noSettingsAccountID),
|
||||
"a live account without a settings row should keep its old logs")
|
||||
}
|
||||
@@ -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
|
||||
@@ -116,8 +120,7 @@ type Provider struct {
|
||||
// Discovery, when non-nil, describes how to ask this vendor which
|
||||
// models the operator's own credential can actually reach, so the
|
||||
// provider form can offer a live list instead of only the hand-curated
|
||||
// Models above. Nil for entries with no listing endpoint (gateways
|
||||
// vary too much) — those keep free-text entry.
|
||||
// Models above. Nil entries keep free-text entry.
|
||||
Discovery *Discovery
|
||||
}
|
||||
|
||||
@@ -154,10 +157,13 @@ const (
|
||||
// one from the caller is also what keeps this from being an open proxy: the
|
||||
// only hosts management will dial are the ones written here.
|
||||
type Discovery struct {
|
||||
Host string
|
||||
Path string
|
||||
Query string
|
||||
Shape ListingShape
|
||||
Host string
|
||||
Path string
|
||||
Query string
|
||||
Shape ListingShape
|
||||
// ExactModelsOnly omits wildcard patterns from listings when NetBird's
|
||||
// provider model rows cannot represent the vendor's matching semantics.
|
||||
ExactModelsOnly bool
|
||||
// Headers are static headers the vendor requires beyond the credential
|
||||
// (Anthropic versions its API through one and rejects a request without
|
||||
// it). The auth header itself comes from AuthHeaderName/Template.
|
||||
@@ -635,6 +641,34 @@ var providers = []Provider{
|
||||
},
|
||||
Models: []Model{},
|
||||
},
|
||||
{
|
||||
ID: "agentgateway",
|
||||
Kind: KindGateway,
|
||||
Name: "agentgateway",
|
||||
Description: "Bring your own agentgateway with trusted NetBird identity stamped on every request",
|
||||
DefaultHost: "",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#8023C3",
|
||||
// Agentgateway accepts both OpenAI and Anthropic request shapes.
|
||||
// Leave ParserID empty so the proxy detects the shape from the URL.
|
||||
ParserID: "",
|
||||
RouterVendors: []string{"openai", "anthropic"},
|
||||
PricingSurfaces: []string{"openai", "anthropic"},
|
||||
Discovery: &Discovery{
|
||||
Path: "/v1/models",
|
||||
Shape: ShapeOpenAIData,
|
||||
ExactModelsOnly: true,
|
||||
},
|
||||
IdentityInjection: &IdentityInjection{
|
||||
HeaderPair: &HeaderPairInjection{
|
||||
EndUserIDHeader: "x-netbird-user-id",
|
||||
TagsHeader: "x-netbird-groups",
|
||||
},
|
||||
},
|
||||
Models: []Model{},
|
||||
},
|
||||
{
|
||||
ID: "portkey",
|
||||
Kind: KindGateway,
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// TestClaudeLineupSelectable pins the models Claude Code resolves to by
|
||||
@@ -34,3 +36,51 @@ func TestClaudeLineupSelectable(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentgatewayCatalogEntry(t *testing.T) {
|
||||
entry, ok := Lookup("agentgateway")
|
||||
require.True(t, ok, "agentgateway must be available in the provider catalog")
|
||||
|
||||
assert.Equal(t, KindGateway, entry.Kind, "agentgateway must be grouped with AI gateways")
|
||||
assert.Empty(t, entry.DefaultHost, "operators must provide their agentgateway proxy URL")
|
||||
assert.Equal(t, "Authorization", entry.AuthHeaderName)
|
||||
assert.Equal(t, "Bearer ${API_KEY}", entry.AuthHeaderTemplate)
|
||||
assert.Equal(t, "application/json", entry.DefaultContentType)
|
||||
assert.Empty(t, entry.ParserID, "URL detection must select the OpenAI or Anthropic parser")
|
||||
assert.Equal(t, []string{"openai", "anthropic"}, entry.RouterVendors,
|
||||
"agentgateway must accept both parser surfaces")
|
||||
assert.Equal(t, []string{"openai", "anthropic"}, entry.PricingSurfaces,
|
||||
"agentgateway models can use either pricing surface")
|
||||
assert.Empty(t, entry.Models, "an empty model list makes agentgateway a catch-all route")
|
||||
require.NotNil(t, entry.Discovery)
|
||||
assert.Empty(t, entry.Discovery.Host, "discovery must use the configured proxy URL")
|
||||
assert.Equal(t, "/v1/models", entry.Discovery.Path)
|
||||
assert.Equal(t, ShapeOpenAIData, entry.Discovery.Shape)
|
||||
assert.True(t, entry.Discovery.ExactModelsOnly,
|
||||
"wildcard model semantics are not supported by NetBird")
|
||||
|
||||
require.NotNil(t, entry.IdentityInjection)
|
||||
require.NotNil(t, entry.IdentityInjection.HeaderPair)
|
||||
assert.Nil(t, entry.IdentityInjection.JSONMetadata)
|
||||
assert.False(t, entry.IdentityInjection.HeaderPair.Customizable,
|
||||
"NetBird identity header names are part of the integration contract")
|
||||
assert.Equal(t, "x-netbird-user-id", entry.IdentityInjection.HeaderPair.EndUserIDHeader)
|
||||
assert.Equal(t, "x-netbird-groups", entry.IdentityInjection.HeaderPair.TagsHeader)
|
||||
assert.False(t, entry.IdentityInjection.HeaderPair.EndUserIDInBody)
|
||||
assert.False(t, entry.IdentityInjection.HeaderPair.TagsInBody)
|
||||
}
|
||||
|
||||
func TestAgentgatewayCatalogAPIResponse(t *testing.T) {
|
||||
entry, ok := Lookup("agentgateway")
|
||||
require.True(t, ok)
|
||||
|
||||
resp := entry.ToAPIResponse()
|
||||
assert.Equal(t, "agentgateway", resp.Id)
|
||||
assert.Equal(t, api.AgentNetworkCatalogProviderKindGateway, resp.Kind)
|
||||
assert.Empty(t, resp.Models)
|
||||
require.NotNil(t, resp.IdentityInjection)
|
||||
require.NotNil(t, resp.IdentityInjection.HeaderPair)
|
||||
assert.False(t, resp.IdentityInjection.HeaderPair.Customizable)
|
||||
assert.Equal(t, "x-netbird-user-id", resp.IdentityInjection.HeaderPair.EndUserIdHeader)
|
||||
assert.Equal(t, "x-netbird-groups", resp.IdentityInjection.HeaderPair.TagsHeader)
|
||||
}
|
||||
|
||||
@@ -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,6 +9,7 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
@@ -17,6 +18,7 @@ import (
|
||||
|
||||
"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")
|
||||
}
|
||||
|
||||
@@ -46,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) {
|
||||
@@ -339,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"
|
||||
@@ -81,10 +80,20 @@ type Manager interface {
|
||||
ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error)
|
||||
GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error)
|
||||
StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int)
|
||||
// RemoveAccountGateway drops the account's gateway mappings from the
|
||||
// proxies. It runs as an account deletion hook.
|
||||
RemoveAccountGateway(ctx context.Context, accountID string) error
|
||||
RecordConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds, tokensIn, tokensOut int64, costUSD float64) error
|
||||
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
|
||||
@@ -126,24 +135,34 @@ type managerImpl struct {
|
||||
proxyController proxy.Controller
|
||||
|
||||
// modelDiscovery queries vendors for the models a credential can reach.
|
||||
// A field rather than a package call so tests can drive it without
|
||||
// reaching the network.
|
||||
// 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 *modeldiscovery.Client
|
||||
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
|
||||
@@ -156,37 +175,149 @@ 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 and
|
||||
// upstream are used instead of the ones in req — so the dashboard can refresh
|
||||
// the list without ever holding the key.
|
||||
// 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
|
||||
@@ -207,11 +338,29 @@ func (m *managerImpl) DiscoverProviderModels(ctx context.Context, accountID, use
|
||||
// name a different one would run a provider's credential against
|
||||
// whichever vendor endpoint they picked.
|
||||
req.CatalogID = record.ProviderID
|
||||
req.UpstreamURL = record.UpstreamURL
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
return m.modelDiscovery.Fetch(ctx, req)
|
||||
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
|
||||
@@ -229,6 +378,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)
|
||||
@@ -264,11 +425,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.
|
||||
@@ -842,6 +1039,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
|
||||
@@ -860,6 +1069,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
|
||||
@@ -871,11 +1173,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
|
||||
@@ -919,6 +1233,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.
|
||||
@@ -945,8 +1294,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)
|
||||
@@ -954,18 +1306,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)
|
||||
@@ -975,10 +1332,29 @@ 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
|
||||
// interval defaults to 24h.
|
||||
// AccessLogRetentionDays, and the consumption counters of deleted accounts.
|
||||
// Usage records are never swept. A non-positive interval defaults to 24h.
|
||||
func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int) {
|
||||
if cleanupIntervalHours <= 0 {
|
||||
cleanupIntervalHours = 24
|
||||
@@ -989,21 +1365,40 @@ func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupInterval
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
m.cleanupAccessLogsOnce(ctx) // run once on startup
|
||||
m.cleanupOnce(ctx) // run once on startup
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
m.cleanupAccessLogsOnce(ctx)
|
||||
m.cleanupOnce(ctx)
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (m *managerImpl) cleanupOnce(ctx context.Context) {
|
||||
m.cleanupAccessLogsOnce(ctx)
|
||||
m.cleanupDeletedAccountConsumption(ctx)
|
||||
}
|
||||
|
||||
// cleanupDeletedAccountConsumption deletes the consumption counters of accounts
|
||||
// that no longer exist. Best-effort: a failure is logged and retried next sweep.
|
||||
func (m *managerImpl) cleanupDeletedAccountConsumption(ctx context.Context) {
|
||||
deleted, err := m.store.DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Warnf("agent-network consumption cleanup: %v", err)
|
||||
return
|
||||
}
|
||||
if deleted > 0 {
|
||||
log.WithContext(ctx).Infof("agent-network consumption cleanup: deleted %d counters of deleted accounts", deleted)
|
||||
}
|
||||
}
|
||||
|
||||
// cleanupAccessLogsOnce sweeps every account's expired access-log rows against
|
||||
// its configured retention. Best-effort: a per-account failure is logged and
|
||||
// the sweep continues.
|
||||
// its configured retention. Deleted accounts, whose settings rows went with
|
||||
// them, get the default retention. Best-effort: a per-account failure is
|
||||
// logged and the sweep continues.
|
||||
func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
|
||||
settings, err := m.store.GetAllAgentNetworkSettings(ctx, store.LockingStrengthNone)
|
||||
if err != nil {
|
||||
@@ -1011,18 +1406,31 @@ func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
|
||||
return
|
||||
}
|
||||
for _, s := range settings {
|
||||
if s.AccessLogRetentionDays <= 0 {
|
||||
continue // keep indefinitely
|
||||
}
|
||||
cutoff := time.Now().UTC().AddDate(0, 0, -s.AccessLogRetentionDays)
|
||||
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, s.AccountID, cutoff)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", s.AccountID, err)
|
||||
continue
|
||||
}
|
||||
if deleted > 0 {
|
||||
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, s.AccountID, s.AccessLogRetentionDays)
|
||||
}
|
||||
m.cleanupAccountAccessLogs(ctx, s.AccountID, s.AccessLogRetentionDays)
|
||||
}
|
||||
|
||||
deleted, err := m.store.GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("agent-network access-log cleanup: list deleted accounts: %v", err)
|
||||
return
|
||||
}
|
||||
for _, accountID := range deleted {
|
||||
m.cleanupAccountAccessLogs(ctx, accountID, types.DefaultAccessLogRetentionDays)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *managerImpl) cleanupAccountAccessLogs(ctx context.Context, accountID string, retentionDays int) {
|
||||
if retentionDays <= 0 {
|
||||
return // keep indefinitely
|
||||
}
|
||||
cutoff := time.Now().UTC().AddDate(0, 0, -retentionDays)
|
||||
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, accountID, cutoff)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", accountID, err)
|
||||
return
|
||||
}
|
||||
if deleted > 0 {
|
||||
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, accountID, retentionDays)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1172,6 +1580,8 @@ func (*mockManager) GetUsageOverview(_ context.Context, _, _ string, _ types.Age
|
||||
|
||||
func (*mockManager) StartAccessLogCleanup(_ context.Context, _ int) {}
|
||||
|
||||
func (*mockManager) RemoveAccountGateway(_ context.Context, _ string) error { return nil }
|
||||
|
||||
func (*mockManager) RecordConsumption(_ context.Context, _ string, _ types.ConsumptionDimension, _ string, _, _, _ int64, _ float64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -52,9 +52,8 @@ const (
|
||||
)
|
||||
|
||||
// ErrNoDiscovery is returned for a catalog entry that declares no listing
|
||||
// endpoint. Gateways vary too much to have one, and the caller should fall
|
||||
// back to the catalog list plus free-text entry rather than treating this as
|
||||
// a failure.
|
||||
// endpoint. The caller should fall back to the catalog list plus free-text
|
||||
// entry rather than treating this as a failure.
|
||||
var ErrNoDiscovery = errors.New("provider has no model-discovery endpoint")
|
||||
|
||||
// ErrInvalidRequest marks a discovery failure caused by the caller's own input
|
||||
@@ -124,14 +123,28 @@ func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
|
||||
return nil, ErrNoDiscovery
|
||||
}
|
||||
|
||||
endpoint, err := c.discoveryURL(entry, req)
|
||||
// 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
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build discovery request: %w", err)
|
||||
@@ -146,7 +159,7 @@ func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
|
||||
|
||||
resp, err := c.httpClient().Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reach %s: %w", entry.Name, err)
|
||||
return nil, &UnreachableError{Provider: entry.Name, Err: err}
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
@@ -157,7 +170,7 @@ func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
|
||||
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, fmt.Errorf("%s returned %d for its model listing", entry.Name, resp.StatusCode)
|
||||
return nil, &VendorStatusError{Provider: entry.Name, Status: resp.StatusCode}
|
||||
}
|
||||
|
||||
ids, err := parseListing(entry.Discovery.Shape, body)
|
||||
@@ -176,12 +189,16 @@ func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
|
||||
// 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(entry catalog.Provider, req Request) (string, error) {
|
||||
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 == "" {
|
||||
return "", fmt.Errorf("%w: provider upstream %q is not a usable URL", ErrInvalidRequest, req.UpstreamURL)
|
||||
// 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
|
||||
}
|
||||
@@ -194,19 +211,47 @@ func (c *Client) discoveryURL(entry catalog.Provider, req Request) (string, erro
|
||||
region = RegionFromUpstream(entry, req.UpstreamURL)
|
||||
}
|
||||
if region == "" {
|
||||
return "", fmt.Errorf("%w: %s discovery needs a region, and none could be read from the provider upstream",
|
||||
ErrInvalidRequest, entry.Name)
|
||||
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.checkPublicHost(target.Hostname()); err != nil {
|
||||
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
|
||||
@@ -244,7 +289,7 @@ func RegionFromUpstream(entry catalog.Provider, upstreamURL string) string {
|
||||
|
||||
// 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(host string) error {
|
||||
func (c *Client) checkPublicHost(ctx context.Context, host string) error {
|
||||
if c.AllowPrivateHosts {
|
||||
return nil
|
||||
}
|
||||
@@ -255,9 +300,6 @@ func (c *Client) checkPublicHost(host string) error {
|
||||
if resolver == nil {
|
||||
resolver = net.DefaultResolver
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), fetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
addrs, err := resolver.LookupNetIP(ctx, "ip", host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve discovery host %q: %w", host, err)
|
||||
@@ -266,7 +308,7 @@ func (c *Client) checkPublicHost(host string) error {
|
||||
// loopback address is still a way to reach loopback.
|
||||
for _, addr := range addrs {
|
||||
if !isPublic(addr) {
|
||||
return fmt.Errorf("%w: discovery host %q resolves to a non-public address", ErrInvalidRequest, host)
|
||||
return fmt.Errorf("%w: %w: discovery host %q resolves to a non-public address", ErrInvalidRequest, ErrPrivateHost, host)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
@@ -356,6 +398,9 @@ func decorate(entry catalog.Provider, ids []listedModel) []Model {
|
||||
if listed.id == "" {
|
||||
continue
|
||||
}
|
||||
if entry.Discovery.ExactModelsOnly && strings.Contains(listed.id, "*") {
|
||||
continue
|
||||
}
|
||||
if _, dup := seen[listed.id]; dup {
|
||||
continue
|
||||
}
|
||||
@@ -463,6 +508,13 @@ func guardDialAddress(address string) error {
|
||||
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
|
||||
|
||||
@@ -2,7 +2,9 @@ package modeldiscovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
@@ -59,6 +61,13 @@ const openAIListing = `{"object":"list","data":[
|
||||
{"id":"gpt-4o","object":"model","created":1715367049,"owned_by":"system"}
|
||||
]}`
|
||||
|
||||
const agentgatewayListing = `{"object":"list","data":[
|
||||
{"id":"gpt-4o-mini","object":"model","created":1785166485,"owned_by":"openai"},
|
||||
{"id":"claude-haiku-4-5","object":"model","created":1785166485,"owned_by":"anthropic"},
|
||||
{"id":"openai/*","object":"model","created":1785166485,"owned_by":"openai"},
|
||||
{"id":"*-latest","object":"model","created":1785166485,"owned_by":"openai"}
|
||||
]}`
|
||||
|
||||
const anthropicListing = `{"data":[
|
||||
{"type":"model","id":"claude-haiku-4-5-20251001","display_name":"Claude Haiku 4.5"},
|
||||
{"type":"model","id":"claude-sonnet-4-6","display_name":"Claude Sonnet 4.6"}
|
||||
@@ -97,6 +106,26 @@ func TestFetchOpenAIListing(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAgentgatewayListing(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, agentgatewayListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "agentgateway",
|
||||
UpstreamURL: "https://gateway.example.com",
|
||||
APIKey: "virtual-key",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "https://gateway.example.com/v1/models", tr.got.URL.String())
|
||||
assert.Equal(t, "Bearer virtual-key", tr.got.Header.Get("Authorization"),
|
||||
"agentgateway model discovery must use the configured virtual key")
|
||||
assert.Equal(t, []string{"gpt-4o-mini", "claude-haiku-4-5"}, ids(models),
|
||||
"model patterns must not be offered as exact NetBird authorization rows")
|
||||
for _, m := range models {
|
||||
assert.True(t, m.PricingKnown, "known upstream model must use NetBird catalog pricing: %s", m.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAnthropicSendsTheVersionHeader(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, anthropicListing)
|
||||
|
||||
@@ -288,7 +317,7 @@ func TestHostGuardRejectsNonPublicAddresses(t *testing.T) {
|
||||
|
||||
func TestHostGuardResolvesAndRejectsLocalhost(t *testing.T) {
|
||||
cl := &Client{}
|
||||
err := cl.checkPublicHost("localhost")
|
||||
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")
|
||||
}
|
||||
@@ -530,3 +559,120 @@ func TestBedrockProfilesFromAnyGeographyArrivePriced(t *testing.T) {
|
||||
// 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")
|
||||
@@ -43,7 +43,7 @@ func parseOpenAIData(body []byte) ([]listedModel, error) {
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode model listing: %w", err)
|
||||
return nil, fmt.Errorf("%w: decode model listing: %w", ErrUnparseableListing, err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Data))
|
||||
for _, entry := range doc.Data {
|
||||
@@ -71,7 +71,7 @@ func parseBedrockInferenceProfiles(body []byte) ([]listedModel, error) {
|
||||
} `json:"inferenceProfileSummaries"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode inference-profile listing: %w", err)
|
||||
return nil, fmt.Errorf("%w: decode inference-profile listing: %w", ErrUnparseableListing, err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Summaries))
|
||||
for _, entry := range doc.Summaries {
|
||||
@@ -98,7 +98,7 @@ func parseVertexPublisherModels(body []byte) ([]listedModel, error) {
|
||||
} `json:"publisherModels"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode publisher-model listing: %w", err)
|
||||
return nil, fmt.Errorf("%w: decode publisher-model listing: %w", ErrUnparseableListing, err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Models))
|
||||
for _, entry := range doc.Models {
|
||||
|
||||
@@ -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"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -2,8 +2,10 @@ package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
goproto "google.golang.org/protobuf/proto"
|
||||
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
@@ -81,18 +83,66 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) {
|
||||
}
|
||||
m.reconcileMu.Unlock()
|
||||
|
||||
for _, entry := range creates {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
m.sendMappings(ctx, accountID, creates, proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED)
|
||||
m.sendMappings(ctx, accountID, updates, proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED)
|
||||
m.sendMappings(ctx, accountID, deletes, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
|
||||
}
|
||||
|
||||
// sendMappings sends each entry as updateType. It sends a copy: the entries'
|
||||
// mappings are shared with reconcileCache, which another reconcile or
|
||||
// RemoveAccountGateway may be reading, so they are never written.
|
||||
func (m *managerImpl) sendMappings(ctx context.Context, accountID string, entries []syntheticMapping, updateType proto.ProxyMappingUpdateType) {
|
||||
for _, entry := range entries {
|
||||
update := goproto.Clone(entry.mapping).(*proto.ProxyMapping)
|
||||
update.Type = updateType
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, update, entry.cluster)
|
||||
}
|
||||
for _, entry := range updates {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
}
|
||||
|
||||
// RemoveAccountGateway tells the proxies to drop every mapping of the account's
|
||||
// gateway, so a deleted account's proxy config, provider API keys included, does
|
||||
// not linger in proxy memory until the next resync. It is an account deletion
|
||||
// hook: it runs before the account's data is removed, the last point at which
|
||||
// the mappings can be synthesised from the store. The cache alone would miss
|
||||
// them, since it is per instance and empty after a restart. If the deletion
|
||||
// then fails, the gateway stays down until the account's next change reconciles
|
||||
// it back.
|
||||
func (m *managerImpl) RemoveAccountGateway(ctx context.Context, accountID string) error {
|
||||
if m.proxyController == nil {
|
||||
return nil
|
||||
}
|
||||
for _, entry := range deletes {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
|
||||
services, err := SynthesizeServices(ctx, m.store, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("synthesise agent network services: %w", err)
|
||||
}
|
||||
oidcCfg := m.proxyController.GetOIDCValidationConfig()
|
||||
removed := make(map[string]syntheticMapping, len(services))
|
||||
for _, svc := range services {
|
||||
if svc == nil || svc.ID == "" {
|
||||
continue
|
||||
}
|
||||
removed[svc.ID] = syntheticMapping{
|
||||
mapping: svc.ToProtoMapping(rpservice.Delete, "", oidcCfg),
|
||||
cluster: svc.ProxyCluster,
|
||||
}
|
||||
}
|
||||
|
||||
m.reconcileMu.Lock()
|
||||
for id, entry := range m.reconcileCache[accountID] {
|
||||
if _, ok := removed[id]; !ok {
|
||||
removed[id] = entry
|
||||
}
|
||||
}
|
||||
delete(m.reconcileCache, accountID)
|
||||
m.reconcileMu.Unlock()
|
||||
|
||||
entries := make([]syntheticMapping, 0, len(removed))
|
||||
for _, entry := range removed {
|
||||
entries = append(entries, entry)
|
||||
}
|
||||
m.sendMappings(ctx, accountID, entries, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
|
||||
return nil
|
||||
}
|
||||
|
||||
// diffMappings classifies the previous→current transition for a single
|
||||
|
||||
@@ -2,6 +2,8 @@ package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
@@ -12,6 +14,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
func newReconcileMgr(t *testing.T, ctrl *gomock.Controller) (*managerImpl, *store.MockStore, *proxy.MockController) {
|
||||
@@ -287,3 +290,154 @@ func TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster(t *testing.T) {
|
||||
assert.Equal(t, "brave-otter.gateway.example.com", deletes[0].cluster)
|
||||
}
|
||||
}
|
||||
|
||||
// TestRemoveAccountGateway_EmitsRemovedFromStore — account deletion runs on an
|
||||
// instance that may never have reconciled the account, so its cache is empty.
|
||||
// The mappings are synthesised from the store, still intact before the delete,
|
||||
// and each is sent as REMOVED to the cluster that serves it.
|
||||
func TestRemoveAccountGateway_EmitsRemovedFromStore(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
|
||||
provider := newReconcileTestProvider()
|
||||
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
|
||||
|
||||
expectReconcileSynthInputs(mockStore, ctx, []*types.Provider{provider}, []*types.Policy{policy}, []*types.Guardrail{})
|
||||
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
|
||||
|
||||
var sent []*proto.ProxyMapping
|
||||
mockProxy.EXPECT().
|
||||
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "eu.proxy.netbird.io").
|
||||
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
|
||||
sent = append(sent, m)
|
||||
})
|
||||
|
||||
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
|
||||
|
||||
require.Len(t, sent, 1, "the account's one gateway mapping must be removed")
|
||||
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type, "the update must be a removal")
|
||||
assert.Equal(t, "agent-net-svc-acct-1", sent[0].Id, "the removal must name the account's gateway service")
|
||||
}
|
||||
|
||||
// TestRemoveAccountGateway_AlsoRemovesCachedMappings — a mapping this instance
|
||||
// last sent but the store no longer synthesises (here, one on another cluster)
|
||||
// is removed too, and the account's cache entry is cleared.
|
||||
func TestRemoveAccountGateway_AlsoRemovesCachedMappings(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
|
||||
mgr.reconcileCache["acct-1"] = map[string]syntheticMapping{
|
||||
"stale-svc": {mapping: &proto.ProxyMapping{Id: "stale-svc"}, cluster: "us.proxy.netbird.io"},
|
||||
}
|
||||
|
||||
// Settings but no providers: the store synthesises nothing.
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
|
||||
Return(newReconcileTestSettings(), nil)
|
||||
mockStore.EXPECT().
|
||||
GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").
|
||||
Return([]*types.Provider{}, nil)
|
||||
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
|
||||
|
||||
var sent []*proto.ProxyMapping
|
||||
mockProxy.EXPECT().
|
||||
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "us.proxy.netbird.io").
|
||||
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
|
||||
sent = append(sent, m)
|
||||
})
|
||||
|
||||
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
|
||||
|
||||
require.Len(t, sent, 1, "the cached mapping must be removed from its own cluster")
|
||||
assert.Equal(t, "stale-svc", sent[0].Id)
|
||||
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type)
|
||||
mgr.reconcileMu.Lock()
|
||||
_, present := mgr.reconcileCache["acct-1"]
|
||||
mgr.reconcileMu.Unlock()
|
||||
assert.False(t, present, "the deleted account's cache entry must be cleared")
|
||||
}
|
||||
|
||||
// TestRemoveAccountGateway_SynthFailureAbortsDeletion — if the mappings cannot
|
||||
// be read, nothing is sent and the error is returned, which as an account
|
||||
// deletion hook keeps the account rather than leaving its gateway running.
|
||||
func TestRemoveAccountGateway_SynthFailureAbortsDeletion(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mgr, mockStore, _ := newReconcileMgr(t, ctrl)
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
|
||||
Return(nil, status.Errorf(status.Internal, "store unavailable"))
|
||||
|
||||
assert.Error(t, mgr.RemoveAccountGateway(ctx, "acct-1"), "a failed synthesis must fail the hook")
|
||||
}
|
||||
|
||||
func TestRemoveAccountGateway_NilProxyController_NoOp(t *testing.T) {
|
||||
mgr := &managerImpl{reconcileCache: make(map[string]map[string]syntheticMapping)}
|
||||
// Must not panic and must not query the store.
|
||||
assert.NoError(t, mgr.RemoveAccountGateway(context.Background(), "acct-1"))
|
||||
}
|
||||
|
||||
// TestReconcile_ConcurrentWithGatewayChanges — while an account's gateway
|
||||
// flaps (its policy is removed and re-added between reads), concurrent
|
||||
// reconciles and RemoveAccountGateway share the cached mappings: one caches a
|
||||
// mapping and sends it, another finds it gone and sends its removal. Run under
|
||||
// -race: neither path may write a cached mapping, only copies of it.
|
||||
func TestReconcile_ConcurrentWithGatewayChanges(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
|
||||
// gomock serialises every call on the controller's mutex, which would give
|
||||
// the race detector the ordering the code under test lacks. The sends go
|
||||
// through a fake that takes no lock.
|
||||
mgr.proxyController = unsyncedSender{MockController: mockProxy}
|
||||
provider := newReconcileTestProvider()
|
||||
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
|
||||
|
||||
var reads atomic.Int64
|
||||
mockStore.EXPECT().GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").Return(newReconcileTestSettings(), nil).AnyTimes()
|
||||
mockStore.EXPECT().GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Provider{provider}, nil).AnyTimes()
|
||||
mockStore.EXPECT().GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, "acct-1").
|
||||
DoAndReturn(func(context.Context, store.LockingStrength, string) ([]*types.Policy, error) {
|
||||
if reads.Add(1)%2 == 0 {
|
||||
return []*types.Policy{}, nil
|
||||
}
|
||||
return []*types.Policy{policy}, nil
|
||||
}).AnyTimes()
|
||||
mockStore.EXPECT().GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Guardrail{}, nil).AnyTimes()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 8; i++ {
|
||||
wg.Add(1)
|
||||
go func(remove bool) {
|
||||
defer wg.Done()
|
||||
for j := 0; j < 50; j++ {
|
||||
if remove && j%10 == 0 {
|
||||
_ = mgr.RemoveAccountGateway(ctx, "acct-1")
|
||||
continue
|
||||
}
|
||||
mgr.reconcile(ctx, "acct-1")
|
||||
}
|
||||
}(i == 0)
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// unsyncedSender answers the calls reconcile makes on every pass without any
|
||||
// locking, so concurrent callers are not ordered by the fake itself.
|
||||
type unsyncedSender struct {
|
||||
*proxy.MockController
|
||||
}
|
||||
|
||||
func (unsyncedSender) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
|
||||
return proxy.OIDCValidationConfig{}
|
||||
}
|
||||
|
||||
func (unsyncedSender) SendServiceUpdateToCluster(context.Context, string, *proto.ProxyMapping, string) {}
|
||||
|
||||
@@ -5,12 +5,14 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/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": ¬Private,
|
||||
// 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", ¬Private)
|
||||
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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -211,18 +211,19 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
}
|
||||
|
||||
groupIndex := indexProviderGroups(enabledPolicies)
|
||||
catalogByProvider := catalogIDsByProvider(enabledProviders)
|
||||
|
||||
// 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)
|
||||
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)
|
||||
modelPolicies := buildModelPolicies(enabledPolicies, guardrailsByID, catalogByProvider)
|
||||
|
||||
routerCfgJSON, err := buildRouterConfigJSON(enabledProviders, groupIndex, modelPolicies)
|
||||
if err != nil {
|
||||
@@ -352,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"`
|
||||
@@ -461,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,
|
||||
@@ -525,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
|
||||
@@ -894,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
|
||||
@@ -918,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{}{}
|
||||
}
|
||||
}
|
||||
@@ -939,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
|
||||
@@ -959,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
|
||||
@@ -1167,23 +1224,25 @@ type routerModelPolicy struct {
|
||||
// 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) map[string][]routerModelPolicy {
|
||||
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)
|
||||
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.
|
||||
rule.Models = append([]string{}, models...)
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -33,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")
|
||||
})
|
||||
@@ -43,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")
|
||||
})
|
||||
@@ -52,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")
|
||||
})
|
||||
@@ -62,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")
|
||||
})
|
||||
@@ -71,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"])
|
||||
})
|
||||
@@ -80,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")
|
||||
})
|
||||
@@ -89,7 +89,7 @@ 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")
|
||||
})
|
||||
@@ -124,7 +124,7 @@ func TestBuildModelPolicies(t *testing.T) {
|
||||
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
|
||||
policyForGroups("p2", []string{"grp-sales"}, []string{"g-opus"}, "prov-x"),
|
||||
}
|
||||
got := buildModelPolicies(policies, byID)
|
||||
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"}},
|
||||
@@ -137,14 +137,14 @@ func TestBuildModelPolicies(t *testing.T) {
|
||||
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
|
||||
policyForGroups("p2", []string{"grp-admin"}, nil, "prov-x"),
|
||||
}
|
||||
got := buildModelPolicies(policies, byID)
|
||||
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)
|
||||
got := buildModelPolicies(policies, byID, nil)
|
||||
assert.Nil(t, got["prov-x"][0].Models,
|
||||
"a guardrail with the allowlist check off restricts nothing")
|
||||
})
|
||||
@@ -154,7 +154,7 @@ func TestBuildModelPolicies(t *testing.T) {
|
||||
"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)
|
||||
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)
|
||||
@@ -162,7 +162,71 @@ func TestBuildModelPolicies(t *testing.T) {
|
||||
|
||||
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),
|
||||
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,9 +6,9 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
@@ -497,6 +497,55 @@ func TestSynthesizeServices_IdentityInject_LiteLLM(t *testing.T) {
|
||||
assert.Equal(t, "x-litellm-tags", entry.HeaderPair.TagsHeader)
|
||||
}
|
||||
|
||||
func TestBuildIdentityInjectConfigJSON_Agentgateway(t *testing.T) {
|
||||
provider := &types.Provider{
|
||||
ID: "prov-agentgateway",
|
||||
ProviderID: "agentgateway",
|
||||
}
|
||||
|
||||
raw, err := buildIdentityInjectConfigJSON(
|
||||
[]*types.Provider{provider},
|
||||
map[string][]string{provider.ID: []string{"grp-eng"}},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
var cfg identityInjectConfig
|
||||
require.NoError(t, json.Unmarshal(raw, &cfg))
|
||||
require.Len(t, cfg.Providers, 1)
|
||||
|
||||
rule := cfg.Providers[0]
|
||||
assert.Equal(t, provider.ID, rule.ProviderID)
|
||||
require.NotNil(t, rule.HeaderPair)
|
||||
assert.Nil(t, rule.JSONMetadata)
|
||||
assert.Equal(t, "x-netbird-user-id", rule.HeaderPair.EndUserIDHeader)
|
||||
assert.Equal(t, "x-netbird-groups", rule.HeaderPair.TagsHeader)
|
||||
assert.False(t, rule.HeaderPair.EndUserIDInBody)
|
||||
assert.False(t, rule.HeaderPair.TagsInBody)
|
||||
}
|
||||
|
||||
func TestBuildRouterConfigJSON_AgentgatewayVendors(t *testing.T) {
|
||||
provider := &types.Provider{
|
||||
ID: "prov-agentgateway",
|
||||
ProviderID: "agentgateway",
|
||||
UpstreamURL: "https://gateway.example.com",
|
||||
APIKey: "virtual-key",
|
||||
}
|
||||
|
||||
raw, err := buildRouterConfigJSON(
|
||||
[]*types.Provider{provider},
|
||||
map[string][]string{provider.ID: {"grp-eng"}},
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
var cfg routerConfig
|
||||
require.NoError(t, json.Unmarshal(raw, &cfg))
|
||||
require.Len(t, cfg.Providers, 1)
|
||||
assert.Empty(t, cfg.Providers[0].Vendor,
|
||||
"the singular vendor remains empty for a multi-surface gateway")
|
||||
assert.Equal(t, []string{"openai", "anthropic"}, cfg.Providers[0].Vendors)
|
||||
}
|
||||
|
||||
// TestSynthesizeServices_IdentityInject_Bifrost_OperatorOverrides
|
||||
// covers the customizable HeaderPair contract. The Bifrost catalog
|
||||
// entry sets HeaderPair.Customizable=true with x-bf-dim-* defaults
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string)
|
||||
return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID)
|
||||
}
|
||||
|
||||
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (string, error) {
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
"github.com/netbirdio/netbird/management/server/geolocation"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
@@ -18,14 +19,16 @@ import (
|
||||
)
|
||||
|
||||
type managerImpl struct {
|
||||
repo accesslogs.Repository
|
||||
store store.Store
|
||||
permissionsManager permissions.Manager
|
||||
geo geolocation.Geolocation
|
||||
cleanupCancel context.CancelFunc
|
||||
}
|
||||
|
||||
func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
|
||||
func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
|
||||
return &managerImpl{
|
||||
repo: repo,
|
||||
store: store,
|
||||
permissionsManager: permissionsManager,
|
||||
geo: geo,
|
||||
@@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac
|
||||
}
|
||||
}
|
||||
|
||||
if err := m.store.CreateAccessLog(ctx, logEntry); err != nil {
|
||||
if err := m.repo.Create(ctx, logEntry); err != nil {
|
||||
log.WithContext(ctx).WithFields(log.Fields{
|
||||
"service_id": logEntry.ServiceID,
|
||||
"method": logEntry.Method,
|
||||
@@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st
|
||||
log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err)
|
||||
}
|
||||
|
||||
logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter)
|
||||
logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
@@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in
|
||||
}
|
||||
|
||||
cutoffTime := time.Now().AddDate(0, 0, -retentionDays)
|
||||
deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime)
|
||||
deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err)
|
||||
return 0, err
|
||||
|
||||
@@ -5,27 +5,27 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
)
|
||||
|
||||
func TestCleanupOldAccessLogs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
retentionDays int
|
||||
setupMock func(*store.MockStore)
|
||||
setupMock func(*accesslogs.MockRepository)
|
||||
expectedCount int64
|
||||
expectedError bool
|
||||
}{
|
||||
{
|
||||
name: "cleanup logs older than retention period",
|
||||
retentionDays: 30,
|
||||
setupMock: func(mockStore *store.MockStore) {
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
setupMock: func(mockRepo *accesslogs.MockRepository) {
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
|
||||
expectedCutoff := time.Now().AddDate(0, 0, -30)
|
||||
timeDiff := olderThan.Sub(expectedCutoff)
|
||||
@@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) {
|
||||
{
|
||||
name: "no logs to cleanup",
|
||||
retentionDays: 30,
|
||||
setupMock: func(mockStore *store.MockStore) {
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
setupMock: func(mockRepo *accesslogs.MockRepository) {
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(0), nil)
|
||||
},
|
||||
expectedCount: 0,
|
||||
@@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
|
||||
{
|
||||
name: "zero retention days skips cleanup",
|
||||
retentionDays: 0,
|
||||
setupMock: func(mockStore *store.MockStore) {
|
||||
// No expectations - DeleteOldAccessLogs should not be called
|
||||
setupMock: func(mockRepo *accesslogs.MockRepository) {
|
||||
// No expectations - DeleteOlderThan should not be called
|
||||
},
|
||||
expectedCount: 0,
|
||||
expectedError: false,
|
||||
@@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
|
||||
{
|
||||
name: "negative retention days skips cleanup",
|
||||
retentionDays: -10,
|
||||
setupMock: func(mockStore *store.MockStore) {
|
||||
// No expectations - DeleteOldAccessLogs should not be called
|
||||
setupMock: func(mockRepo *accesslogs.MockRepository) {
|
||||
// No expectations - DeleteOlderThan should not be called
|
||||
},
|
||||
expectedCount: 0,
|
||||
expectedError: false,
|
||||
@@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
tt.setupMock(mockStore)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
tt.setupMock(mockRepo)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
|
||||
expectedCutoff := time.Now().AddDate(0, 0, -30)
|
||||
timeDiff := olderThan.Sub(expectedCutoff)
|
||||
@@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) {
|
||||
})
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
// No expectations - cleanup should not run
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) {
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// If DeleteOldAccessLogs was called, the test will fail due to unexpected call
|
||||
// If DeleteOlderThan was called, the test will fail due to unexpected call
|
||||
})
|
||||
|
||||
t.Run("periodic cleanup runs immediately on start", func(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(2), nil).
|
||||
Times(1)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(1), nil).
|
||||
Times(1)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(0), nil).
|
||||
Times(1)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(3), nil).
|
||||
Times(1)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
@@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockRepo := accesslogs.NewMockRepository(ctrl)
|
||||
|
||||
mockStore.EXPECT().
|
||||
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
|
||||
mockRepo.EXPECT().
|
||||
DeleteOlderThan(gomock.Any(), gomock.Any()).
|
||||
Return(int64(1), nil).
|
||||
Times(1)
|
||||
|
||||
manager := &managerImpl{
|
||||
store: mockStore,
|
||||
repo: mockRepo,
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
type sqlRepository struct {
|
||||
conn *db.Conn
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
// NewRepository returns the access log repository backed by conn.
|
||||
func NewRepository(conn *db.Conn) accesslogs.Repository {
|
||||
return &sqlRepository{conn: conn, db: conn.DB(nil)}
|
||||
}
|
||||
|
||||
func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository {
|
||||
return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)}
|
||||
}
|
||||
|
||||
func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error {
|
||||
if err := r.db.Create(entry).Error; err != nil {
|
||||
log.WithContext(ctx).WithFields(log.Fields{
|
||||
"service_id": entry.ServiceID,
|
||||
"method": entry.Method,
|
||||
"host": entry.Host,
|
||||
"path": entry.Path,
|
||||
}).Errorf("failed to create access log entry in store: %v", err)
|
||||
return status.Errorf(status.Internal, "failed to create access log entry in store")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListByAccount returns one page of an account's access logs together with the
|
||||
// total number of entries matching the filter.
|
||||
func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) {
|
||||
var totalCount int64
|
||||
countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter)
|
||||
if err := countQuery.Count(&totalCount).Error; err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to count access logs: %v", err)
|
||||
return nil, 0, status.Errorf(status.Internal, "failed to count access logs")
|
||||
}
|
||||
|
||||
query := applyFilters(r.db.Where("account_id = ?", accountID), filter)
|
||||
sortOrder := strings.ToUpper(filter.GetSortOrder())
|
||||
for _, column := range strings.Split(filter.GetSortColumn(), ",") {
|
||||
if column = strings.TrimSpace(column); column != "" {
|
||||
query = query.Order(column + " " + sortOrder)
|
||||
}
|
||||
}
|
||||
query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset())
|
||||
if lockStrength != db.LockingStrengthNone {
|
||||
query = query.Clauses(clause.Locking{Strength: string(lockStrength)})
|
||||
}
|
||||
|
||||
var logs []*accesslogs.AccessLogEntry
|
||||
if err := query.Find(&logs).Error; err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err)
|
||||
return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store")
|
||||
}
|
||||
|
||||
return logs, totalCount, nil
|
||||
}
|
||||
|
||||
func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
|
||||
result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{})
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error)
|
||||
return 0, status.Errorf(status.Internal, "failed to delete old access logs")
|
||||
}
|
||||
return result.RowsAffected, nil
|
||||
}
|
||||
|
||||
func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB {
|
||||
if filter.Search != nil {
|
||||
searchPattern := "%" + *filter.Search + "%"
|
||||
query = query.Where(
|
||||
"id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)",
|
||||
searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern,
|
||||
)
|
||||
}
|
||||
|
||||
if filter.SourceIP != nil {
|
||||
query = query.Where("location_connection_ip = ?", *filter.SourceIP)
|
||||
}
|
||||
|
||||
if filter.Host != nil {
|
||||
query = query.Where("host = ?", *filter.Host)
|
||||
}
|
||||
|
||||
if filter.Path != nil {
|
||||
query = query.Where("path LIKE ?", "%"+*filter.Path+"%")
|
||||
}
|
||||
|
||||
if filter.UserID != nil {
|
||||
query = query.Where("user_id = ?", *filter.UserID)
|
||||
}
|
||||
|
||||
if filter.Method != nil {
|
||||
query = query.Where("method = ?", *filter.Method)
|
||||
}
|
||||
|
||||
if filter.Status != nil {
|
||||
switch *filter.Status {
|
||||
case "success":
|
||||
query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400)
|
||||
case "failed":
|
||||
query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400)
|
||||
}
|
||||
}
|
||||
|
||||
if filter.StatusCode != nil {
|
||||
query = query.Where("status_code = ?", *filter.StatusCode)
|
||||
}
|
||||
|
||||
if filter.StartDate != nil {
|
||||
query = query.Where("timestamp >= ?", *filter.StartDate)
|
||||
}
|
||||
|
||||
if filter.EndDate != nil {
|
||||
query = query.Where("timestamp <= ?", *filter.EndDate)
|
||||
}
|
||||
|
||||
return query
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db/dbtest"
|
||||
)
|
||||
|
||||
func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) {
|
||||
conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{})
|
||||
return NewRepository(conn), conn
|
||||
}
|
||||
|
||||
func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry {
|
||||
return &accesslogs.AccessLogEntry{
|
||||
ID: id,
|
||||
AccountID: accountID,
|
||||
Method: method,
|
||||
Host: "app.example.com",
|
||||
Path: "/",
|
||||
StatusCode: 200,
|
||||
Timestamp: time.Now().Add(-age),
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlRepository_ListByAccount(t *testing.T) {
|
||||
repo, _ := newTestRepository(t)
|
||||
ctx := context.Background()
|
||||
for _, entry := range []*accesslogs.AccessLogEntry{
|
||||
newEntry("a1", "acc-a", "GET", 3*time.Hour),
|
||||
newEntry("a2", "acc-a", "POST", 2*time.Hour),
|
||||
newEntry("a3", "acc-a", "GET", time.Hour),
|
||||
newEntry("b1", "acc-b", "GET", time.Hour),
|
||||
} {
|
||||
require.NoError(t, repo.Create(ctx, entry))
|
||||
}
|
||||
|
||||
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 3, total)
|
||||
require.Len(t, logs, 2)
|
||||
assert.Equal(t, "a3", logs[0].ID)
|
||||
assert.Equal(t, "a2", logs[1].ID)
|
||||
|
||||
method := "GET"
|
||||
logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"})
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 2, total)
|
||||
require.Len(t, logs, 2)
|
||||
assert.Equal(t, "a1", logs[0].ID)
|
||||
assert.Equal(t, "a3", logs[1].ID)
|
||||
}
|
||||
|
||||
func TestSqlRepository_DeleteOlderThan(t *testing.T) {
|
||||
repo, _ := newTestRepository(t)
|
||||
ctx := context.Background()
|
||||
require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour)))
|
||||
require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour)))
|
||||
|
||||
deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour))
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 1, deleted)
|
||||
|
||||
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 1, total)
|
||||
require.Len(t, logs, 1)
|
||||
assert.Equal(t, "new", logs[0].ID)
|
||||
}
|
||||
|
||||
func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) {
|
||||
repo, conn := newTestRepository(t)
|
||||
ctx := context.Background()
|
||||
failure := errors.New("abort")
|
||||
|
||||
err := conn.RunInTx(ctx, func(tx *db.Tx) error {
|
||||
txRepo := repo.WithTx(tx)
|
||||
require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0)))
|
||||
_, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 1, total)
|
||||
return failure
|
||||
})
|
||||
require.ErrorIs(t, err, failure)
|
||||
|
||||
_, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Zero(t, total)
|
||||
}
|
||||
|
||||
func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) {
|
||||
repo, _ := newTestRepository(t)
|
||||
ctx := context.Background()
|
||||
statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404}
|
||||
for id, code := range statusCodes {
|
||||
entry := newEntry(id, "acc", "GET", time.Hour)
|
||||
entry.StatusCode = code
|
||||
require.NoError(t, repo.Create(ctx, entry))
|
||||
}
|
||||
foreign := newEntry("foreign", "other", "GET", time.Hour)
|
||||
foreign.StatusCode = 500
|
||||
require.NoError(t, repo.Create(ctx, foreign))
|
||||
|
||||
listIDs := func(status string) []string {
|
||||
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, len(logs), total)
|
||||
ids := make([]string, 0, len(logs))
|
||||
for _, entry := range logs {
|
||||
ids = append(ids, entry.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
assert.Equal(t, []string{"info", "notfound"}, listIDs("failed"))
|
||||
assert.Equal(t, []string{"ok"}, listIDs("success"))
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package accesslogs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
)
|
||||
|
||||
//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
|
||||
|
||||
// Repository persists reverse proxy access log entries.
|
||||
type Repository interface {
|
||||
WithTx(tx *db.Tx) Repository
|
||||
Create(ctx context.Context, entry *AccessLogEntry) error
|
||||
ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error)
|
||||
DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error)
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./repository.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package accesslogs is a generated GoMock package.
|
||||
package accesslogs
|
||||
|
||||
import (
|
||||
context "context"
|
||||
reflect "reflect"
|
||||
time "time"
|
||||
|
||||
db "github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockRepository is a mock of Repository interface.
|
||||
type MockRepository struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockRepositoryMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockRepositoryMockRecorder is the mock recorder for MockRepository.
|
||||
type MockRepositoryMockRecorder struct {
|
||||
mock *MockRepository
|
||||
}
|
||||
|
||||
// NewMockRepository creates a new mock instance.
|
||||
func NewMockRepository(ctrl *gomock.Controller) *MockRepository {
|
||||
mock := &MockRepository{ctrl: ctrl}
|
||||
mock.recorder = &MockRepositoryMockRecorder{mock}
|
||||
return mock
|
||||
}
|
||||
|
||||
// EXPECT returns an object that allows the caller to indicate expected use.
|
||||
func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder {
|
||||
return m.recorder
|
||||
}
|
||||
|
||||
// Create mocks base method.
|
||||
func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Create", ctx, entry)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// Create indicates an expected call of Create.
|
||||
func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry)
|
||||
}
|
||||
|
||||
// DeleteOlderThan mocks base method.
|
||||
func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan)
|
||||
ret0, _ := ret[0].(int64)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// DeleteOlderThan indicates an expected call of DeleteOlderThan.
|
||||
func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan)
|
||||
}
|
||||
|
||||
// ListByAccount mocks base method.
|
||||
func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter)
|
||||
ret0, _ := ret[0].([]*AccessLogEntry)
|
||||
ret1, _ := ret[1].(int64)
|
||||
ret2, _ := ret[2].(error)
|
||||
return ret0, ret1, ret2
|
||||
}
|
||||
|
||||
// ListByAccount indicates an expected call of ListByAccount.
|
||||
func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter)
|
||||
}
|
||||
|
||||
// WithTx mocks base method.
|
||||
func (m *MockRepository) WithTx(tx *db.Tx) Repository {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "WithTx", tx)
|
||||
ret0, _ := ret[0].(Repository)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// WithTx indicates an expected call of WithTx.
|
||||
func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
)
|
||||
|
||||
func TestDeleteDomain_ServiceDependencies(t *testing.T) {
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
domainName string
|
||||
serviceHost string
|
||||
accountID string
|
||||
enabled bool
|
||||
protected bool
|
||||
}{
|
||||
{"exact", "example.com", "example.com", accountA, true, true},
|
||||
{"subdomain", "example.com", "deep.app.example.com", accountA, true, true},
|
||||
{"disabled", "example.com", "app.example.com", accountA, false, true},
|
||||
// A service is authorized by its own account's registration, so another
|
||||
// account's service under this namespace is not a dependency of it.
|
||||
{"other account", "example.com", "app.example.com", accountB, true, false},
|
||||
{"case and trailing dot", "example.com", "APP.EXAMPLE.COM.", accountA, true, true},
|
||||
{"suffix boundary", "example.com", "notexample.com", accountA, true, false},
|
||||
{"literal underscore", "a_b.example.com", "app.a_b.example.com", accountA, true, true},
|
||||
{"underscore wildcard", "a_b.example.com", "app.axb.example.com", accountA, true, false},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
events := captureDomainEvents(env)
|
||||
d, err := env.store.CreateCustomDomain(ctx, accountA, tt.domainName, testCluster, true)
|
||||
require.NoError(t, err)
|
||||
svc := &rpservice.Service{
|
||||
ID: "dependent", AccountID: tt.accountID, Domain: tt.serviceHost,
|
||||
Enabled: tt.enabled, ProxyCluster: testCluster,
|
||||
}
|
||||
require.NoError(t, env.store.CreateService(ctx, svc))
|
||||
router := mux.NewRouter()
|
||||
RegisterEndpoints(router, env.manager)
|
||||
deleteDomain := func() *httptest.ResponseRecorder {
|
||||
req := httptest.NewRequest(http.MethodDelete, "/domains/"+d.ID, nil)
|
||||
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: accountA, UserId: accountAUser})
|
||||
response := httptest.NewRecorder()
|
||||
router.ServeHTTP(response, req)
|
||||
return response
|
||||
}
|
||||
|
||||
response := deleteDomain()
|
||||
if tt.protected {
|
||||
require.Equal(t, http.StatusPreconditionFailed, response.Code, "dependent services must block deletion: %s", response.Body.String())
|
||||
assert.NotContains(t, response.Body.String(), tt.accountID, "the error must not reveal the service's account")
|
||||
assert.NotNil(t, storedDomain(t, env.store, accountA, d.Domain), "the namespace must remain reserved")
|
||||
assert.Empty(t, events.get(), "rejected deletion must not emit DomainDeleted")
|
||||
stored, err := env.store.GetServiceByID(ctx, nbstore.LockingStrengthNone, tt.accountID, svc.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, svc.Enabled, stored.Enabled, "rejected deletion must preserve the service")
|
||||
require.NoError(t, env.store.DeleteService(ctx, tt.accountID, svc.ID))
|
||||
response = deleteDomain()
|
||||
}
|
||||
require.Equal(t, http.StatusNoContent, response.Code, "deletion must succeed without dependencies: %s", response.Body.String())
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "the registration must be deleted")
|
||||
captured := events.get()
|
||||
require.Len(t, captured, 1, "only successful deletion may emit an event")
|
||||
assert.Equal(t, activity.DomainDeleted, captured[0].Activity, "the event must describe the successful deletion")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,14 +346,37 @@ 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)
|
||||
}
|
||||
|
||||
// ValidateServiceDomain holds custom domain authorization through a service write transaction.
|
||||
func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, accountID, serviceDomain, cluster string) error {
|
||||
if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok {
|
||||
return nil
|
||||
}
|
||||
name, err := nbdomain.FromString(serviceDomain)
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "invalid service domain: %v", err)
|
||||
}
|
||||
customDomains, err := tx.LockCustomDomains(ctx, accountID, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
target, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
|
||||
if match != customDomainValidated || target != cluster {
|
||||
return status.Errorf(status.PreconditionFailed, "custom domain authorization changed; retry the service operation")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]string, error) {
|
||||
byopAddresses, err := m.proxyManager.GetActiveClusterAddressesForAccount(ctx, accountID)
|
||||
if err != nil {
|
||||
@@ -363,19 +434,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")
|
||||
}
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
|
||||
// Manager defines the interface for proxy operations
|
||||
type Manager interface {
|
||||
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
|
||||
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error)
|
||||
Disconnect(ctx context.Context, proxyID, sessionID string) error
|
||||
Heartbeat(ctx context.Context, p *Proxy) error
|
||||
GetActiveClusterAddresses(ctx context.Context) ([]string, error)
|
||||
@@ -20,6 +20,8 @@ type Manager interface {
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
|
||||
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
|
||||
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
nbversion "github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// store defines the interface for proxy persistence operations
|
||||
@@ -22,6 +23,8 @@ type store interface {
|
||||
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -29,6 +32,8 @@ type store interface {
|
||||
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
|
||||
}
|
||||
|
||||
const minSessionCodeVersion = "0.81.0"
|
||||
|
||||
// Manager handles all proxy operations
|
||||
type Manager struct {
|
||||
store store
|
||||
@@ -50,7 +55,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) {
|
||||
|
||||
// Connect registers a new proxy connection in the database.
|
||||
// capabilities may be nil for old proxies that do not report them.
|
||||
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
|
||||
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
|
||||
now := time.Now()
|
||||
var caps proxy.Capabilities
|
||||
if capabilities != nil {
|
||||
@@ -61,6 +66,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
|
||||
SessionID: sessionID,
|
||||
ClusterAddress: clusterAddress,
|
||||
IPAddress: ipAddress,
|
||||
Version: truncateVersion(version),
|
||||
AccountID: accountID,
|
||||
LastSeen: now,
|
||||
ConnectedAt: &now,
|
||||
@@ -78,6 +84,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
|
||||
"sessionID": sessionID,
|
||||
"clusterAddress": clusterAddress,
|
||||
"ipAddress": ipAddress,
|
||||
"version": p.Version,
|
||||
}).Info("proxy connected")
|
||||
|
||||
return p, nil
|
||||
@@ -143,6 +150,27 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
|
||||
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported).
|
||||
func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode reports whether all active proxies support session codes.
|
||||
func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
|
||||
versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr)
|
||||
if err != nil || len(versions) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, version := range versions {
|
||||
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
|
||||
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
|
||||
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
|
||||
@@ -184,3 +212,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// truncateVersion cuts a proxy-reported version to the column width so an
|
||||
// oversized value cannot fail the save and block the connect.
|
||||
func truncateVersion(version string) string {
|
||||
runes := []rune(version)
|
||||
if len(runes) <= proxy.MaxVersionLength {
|
||||
return version
|
||||
}
|
||||
return string(runes[:proxy.MaxVersionLength])
|
||||
}
|
||||
|
||||
@@ -4,8 +4,10 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -20,6 +22,7 @@ type mockStore struct {
|
||||
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
|
||||
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
|
||||
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
|
||||
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
|
||||
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
|
||||
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -102,6 +105,15 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
|
||||
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) {
|
||||
if m.getActiveProxyVersionsFunc != nil {
|
||||
return m.getActiveProxyVersionsFunc(ctx, clusterAddress)
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func newTestManager(s store) *Manager {
|
||||
meter := noop.NewMeterProvider().Meter("test")
|
||||
@@ -112,6 +124,34 @@ func newTestManager(s store) *Manager {
|
||||
return m
|
||||
}
|
||||
|
||||
func TestClusterSupportsSessionCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
versions []string
|
||||
storeErr error
|
||||
want bool
|
||||
}{
|
||||
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
|
||||
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
|
||||
{name: "missing version", versions: []string{"0.81.0", ""}},
|
||||
{name: "no active proxies"},
|
||||
{name: "store error", storeErr: errors.New("db error")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
s := &mockStore{
|
||||
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return tt.versions, tt.storeErr
|
||||
},
|
||||
}
|
||||
|
||||
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnect_WithAccountID(t *testing.T) {
|
||||
accountID := "acc-123"
|
||||
|
||||
@@ -124,7 +164,7 @@ func TestConnect_WithAccountID(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
@@ -132,6 +172,7 @@ func TestConnect_WithAccountID(t *testing.T) {
|
||||
assert.Equal(t, "session-1", savedProxy.SessionID)
|
||||
assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress)
|
||||
assert.Equal(t, "10.0.0.1", savedProxy.IPAddress)
|
||||
assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored")
|
||||
assert.Equal(t, &accountID, savedProxy.AccountID)
|
||||
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
|
||||
assert.NotNil(t, savedProxy.ConnectedAt)
|
||||
@@ -147,7 +188,7 @@ func TestConnect_WithoutAccountID(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
@@ -155,6 +196,29 @@ func TestConnect_WithoutAccountID(t *testing.T) {
|
||||
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
|
||||
}
|
||||
|
||||
func TestConnect_TruncatesOversizedVersion(t *testing.T) {
|
||||
var savedProxy *proxy.Proxy
|
||||
s := &mockStore{
|
||||
saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error {
|
||||
savedProxy = p
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
// Multi-byte runes make sure the cut counts characters, as varchar does,
|
||||
// and never splits a rune into invalid UTF-8.
|
||||
version := strings.Repeat("ü", proxy.MaxVersionLength+10)
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width")
|
||||
assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8")
|
||||
assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one")
|
||||
}
|
||||
|
||||
func TestConnect_StoreError(t *testing.T) {
|
||||
s := &mockStore{
|
||||
saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error {
|
||||
@@ -163,7 +227,7 @@ func TestConnect_StoreError(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
|
||||
@@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate mocks base method.
|
||||
func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(*bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate.
|
||||
func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterRequireSubdomain mocks base method.
|
||||
func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -112,19 +126,33 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// Connect mocks base method.
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
// ClusterSupportsSessionCode mocks base method.
|
||||
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
|
||||
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// Connect mocks base method.
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
|
||||
ret0, _ := ret[0].(*Proxy)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// Connect indicates an expected call of Connect.
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, 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)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
|
||||
}
|
||||
|
||||
// CountAccountProxies mocks base method.
|
||||
|
||||
@@ -9,6 +9,9 @@ const (
|
||||
StatusDisconnected = "disconnected"
|
||||
)
|
||||
|
||||
// MaxVersionLength is the width of the Version column, in characters.
|
||||
const MaxVersionLength = 255
|
||||
|
||||
// Capabilities describes what a proxy can handle, as reported via gRPC.
|
||||
// Nil fields mean the proxy never reported this capability.
|
||||
type Capabilities struct {
|
||||
@@ -31,6 +34,7 @@ type Proxy struct {
|
||||
SessionID string `gorm:"type:varchar(36)"`
|
||||
ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"`
|
||||
IPAddress string `gorm:"type:varchar(45)"`
|
||||
Version string `gorm:"type:varchar(255)"`
|
||||
AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"`
|
||||
LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"`
|
||||
ConnectedAt *time.Time
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package proxytoken
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"time"
|
||||
@@ -18,13 +19,29 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// RevocationGuard vetoes the tenant-facing revocation of a proxy access
|
||||
// token. Implementations are supplied by integrations; none is installed by
|
||||
// default, so every token the caller's account owns may be revoked. It is
|
||||
// consulted after the ownership check and before the token is revoked. A
|
||||
// returned status error is written with util.WriteError: its type selects the
|
||||
// HTTP status and its message is shown to the caller, so it must not carry
|
||||
// internal detail. Any other error is reported as a generic internal error.
|
||||
type RevocationGuard interface {
|
||||
CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error
|
||||
}
|
||||
|
||||
type handler struct {
|
||||
store store.Store
|
||||
permissionsManager permissions.Manager
|
||||
// revocationGuard vetoes revocations. Optional — when nil every owned
|
||||
// token may be revoked.
|
||||
revocationGuard RevocationGuard
|
||||
}
|
||||
|
||||
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, router *mux.Router) {
|
||||
h := &handler{store: s, permissionsManager: permissionsManager}
|
||||
// RegisterEndpoints registers the proxy token endpoints. revocationGuard is
|
||||
// optional; pass nil for no revocation policy.
|
||||
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, revocationGuard RevocationGuard, router *mux.Router) {
|
||||
h := &handler{store: s, permissionsManager: permissionsManager, revocationGuard: revocationGuard}
|
||||
router.HandleFunc("/reverse-proxies/proxy-tokens", h.listTokens).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/reverse-proxies/proxy-tokens", h.createToken).Methods("POST", "OPTIONS")
|
||||
router.HandleFunc("/reverse-proxies/proxy-tokens/{tokenId}", h.revokeToken).Methods("DELETE", "OPTIONS")
|
||||
@@ -154,6 +171,13 @@ func (h *handler) revokeToken(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if h.revocationGuard != nil {
|
||||
if err := h.revocationGuard.CheckProxyAccessTokenRevocation(ctx, token); err != nil {
|
||||
util.WriteError(ctx, err, w)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := h.store.RevokeProxyAccessToken(ctx, tokenID); err != nil {
|
||||
util.WriteErrorResponse("failed to revoke token", http.StatusInternalServerError, w)
|
||||
return
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
@@ -22,6 +23,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
func authContext(accountID, userID string) context.Context {
|
||||
@@ -273,3 +275,152 @@ func TestRevokeToken_ManagementWideToken(t *testing.T) {
|
||||
h.revokeToken(w, req)
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
}
|
||||
|
||||
type revocationGuardFunc func(ctx context.Context, token *types.ProxyAccessToken) error
|
||||
|
||||
func (f revocationGuardFunc) CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error {
|
||||
return f(ctx, token)
|
||||
}
|
||||
|
||||
func TestRevokeToken_GuardRefuses(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
accountID := "acc-123"
|
||||
|
||||
// No RevokeProxyAccessToken expectation: a refused revocation must not
|
||||
// reach the store.
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
|
||||
ID: "tok-1",
|
||||
AccountID: &accountID,
|
||||
}, nil)
|
||||
|
||||
permsMgr := permissions.NewMockManager(ctrl)
|
||||
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
|
||||
|
||||
var checked *types.ProxyAccessToken
|
||||
h := &handler{
|
||||
store: mockStore,
|
||||
permissionsManager: permsMgr,
|
||||
revocationGuard: revocationGuardFunc(func(_ context.Context, token *types.ProxyAccessToken) error {
|
||||
checked = token
|
||||
return status.Errorf(status.PreconditionFailed, "token is in use")
|
||||
}),
|
||||
}
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
|
||||
req = req.WithContext(authContext(accountID, "user-1"))
|
||||
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.revokeToken(w, req)
|
||||
assert.Equal(t, http.StatusPreconditionFailed, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "token is in use")
|
||||
require.NotNil(t, checked)
|
||||
assert.Equal(t, "tok-1", checked.ID)
|
||||
}
|
||||
|
||||
func TestRevokeToken_GuardAllows(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
accountID := "acc-123"
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
|
||||
ID: "tok-1",
|
||||
AccountID: &accountID,
|
||||
}, nil)
|
||||
mockStore.EXPECT().RevokeProxyAccessToken(gomock.Any(), "tok-1").Return(nil)
|
||||
|
||||
permsMgr := permissions.NewMockManager(ctrl)
|
||||
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
|
||||
|
||||
h := &handler{
|
||||
store: mockStore,
|
||||
permissionsManager: permsMgr,
|
||||
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
|
||||
return nil
|
||||
}),
|
||||
}
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
|
||||
req = req.WithContext(authContext(accountID, "user-1"))
|
||||
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.revokeToken(w, req)
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func TestRevokeToken_GuardFailure(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
accountID := "acc-123"
|
||||
|
||||
// No RevokeProxyAccessToken expectation: a guard that cannot decide must
|
||||
// not let the revocation through.
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
|
||||
ID: "tok-1",
|
||||
AccountID: &accountID,
|
||||
}, nil)
|
||||
|
||||
permsMgr := permissions.NewMockManager(ctrl)
|
||||
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
|
||||
|
||||
h := &handler{
|
||||
store: mockStore,
|
||||
permissionsManager: permsMgr,
|
||||
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
|
||||
return errors.New("connection refused")
|
||||
}),
|
||||
}
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
|
||||
req = req.WithContext(authContext(accountID, "user-1"))
|
||||
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.revokeToken(w, req)
|
||||
assert.Equal(t, http.StatusInternalServerError, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "internal server error")
|
||||
assert.NotContains(t, w.Body.String(), "connection refused")
|
||||
}
|
||||
|
||||
func TestRevokeToken_GuardNotConsultedForForeignToken(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
defer ctrl.Finish()
|
||||
|
||||
otherAccount := "acc-other"
|
||||
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
|
||||
ID: "tok-1",
|
||||
AccountID: &otherAccount,
|
||||
}, nil)
|
||||
|
||||
permsMgr := permissions.NewMockManager(ctrl)
|
||||
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), "acc-123", "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
|
||||
|
||||
// A foreign token must read as not found, not reveal through the guard's
|
||||
// answer that it belongs to some account's managed proxy.
|
||||
h := &handler{
|
||||
store: mockStore,
|
||||
permissionsManager: permsMgr,
|
||||
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
|
||||
t.Fatal("guard consulted for a token the caller does not own")
|
||||
return nil
|
||||
}),
|
||||
}
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
|
||||
req = req.WithContext(authContext("acc-123", "user-1"))
|
||||
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.revokeToken(w, req)
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
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")
|
||||
}
|
||||
|
||||
func TestCreateService_DomainDeletedBeforeWrite(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
withRealDomainManager(t, mgr, testStore)
|
||||
|
||||
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
|
||||
require.NoError(t, err)
|
||||
svc := newTestService("app.proven.example.com")
|
||||
require.NoError(t, mgr.initializeServiceForCreate(ctx, testAccountID, svc))
|
||||
|
||||
// Delete after the initial authorization check, before the service transaction starts.
|
||||
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
|
||||
err = mgr.persistNewService(ctx, testAccountID, svc)
|
||||
require.Error(t, err, "an earlier validation result must not authorize a deleted registration")
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "the caller must receive a typed precondition error")
|
||||
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the service must require current domain authorization")
|
||||
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, services, "the failed write must not leave a service")
|
||||
}
|
||||
|
||||
func TestUpdateService_DomainDeletedBeforeWrite(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
withRealDomainManager(t, mgr, testStore)
|
||||
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "original.example.com", validationTestCluster, true)
|
||||
require.NoError(t, err)
|
||||
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "destination.example.com", validationTestCluster, true)
|
||||
require.NoError(t, err)
|
||||
svc, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.original.example.com"))
|
||||
require.NoError(t, err)
|
||||
moved := svc.Copy()
|
||||
moved.Domain = "app.destination.example.com"
|
||||
cluster, err := mgr.resolveEffectiveCluster(ctx, testAccountID, moved)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
|
||||
err = testStore.ExecuteInTransaction(ctx, func(tx store.Store) error {
|
||||
return mgr.executeServiceUpdate(ctx, tx, testAccountID, moved, &serviceUpdateInfo{}, nil, cluster)
|
||||
})
|
||||
require.Error(t, err, "a domain deleted after cluster resolution must reject the update")
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "the caller must receive a typed precondition error")
|
||||
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the move must require current domain authorization")
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, svc.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, svc.Domain, stored.Domain, "the service must retain its authorized domain")
|
||||
}
|
||||
@@ -74,6 +74,7 @@ const unknownHostPlaceholder = "unknown"
|
||||
// ClusterDeriver derives the proxy cluster from a domain.
|
||||
type ClusterDeriver interface {
|
||||
DeriveClusterFromDomain(ctx context.Context, accountID, domain string) (string, error)
|
||||
ValidateServiceDomain(ctx context.Context, tx store.Store, accountID, domain, cluster string) error
|
||||
GetClusterDomains() []string
|
||||
}
|
||||
|
||||
@@ -83,6 +84,7 @@ type CapabilityProvider interface {
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
}
|
||||
|
||||
type Manager struct {
|
||||
@@ -331,7 +333,14 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc *
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
if svc.Domain != "" {
|
||||
if err := m.checkDomainAvailable(ctx, transaction, svc.Domain, ""); err != nil {
|
||||
return err
|
||||
@@ -365,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service)
|
||||
return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster)
|
||||
}
|
||||
|
||||
// validatePrivateClusterTargets rejects cluster and direct upstream targets unless
|
||||
// every active proxy in the service's cluster reports the private capability. The
|
||||
// mapping reaches all proxies in the cluster, so one non-private proxy would serve
|
||||
// these targets too. An unreported capability is treated as unsupported. Must be
|
||||
// called outside a transaction, like clusterCustomPorts.
|
||||
func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error {
|
||||
target := firstPrivateClusterTarget(targets)
|
||||
if target == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private {
|
||||
return nil
|
||||
}
|
||||
|
||||
if target.TargetType == service.TargetTypeCluster {
|
||||
return status.Errorf(status.InvalidArgument,
|
||||
"target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it",
|
||||
service.TargetTypeCluster, cluster)
|
||||
}
|
||||
return status.Errorf(status.InvalidArgument,
|
||||
"direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster)
|
||||
}
|
||||
|
||||
// firstPrivateClusterTarget returns the first target that only a private cluster may serve.
|
||||
func firstPrivateClusterTarget(targets []*service.Target) *service.Target {
|
||||
for _, target := range targets {
|
||||
if target == nil {
|
||||
continue
|
||||
}
|
||||
if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream {
|
||||
return target
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureL4Port auto-assigns a listen port when needed and validates cluster support.
|
||||
// customPorts must be pre-computed via clusterCustomPorts before entering a transaction.
|
||||
func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error {
|
||||
@@ -460,7 +506,14 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.validateEphemeralPreconditions(ctx, transaction, accountID, peerID, svc); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -577,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Validate subdomain requirement *before* the transaction: the underlying
|
||||
// capability lookup talks to the main DB pool, and SQLite's single-connection
|
||||
// pool would self-deadlock if this ran while the tx already held the only
|
||||
@@ -606,19 +663,25 @@ 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 {
|
||||
if err := m.validateServiceDomain(ctx, transaction, accountID, service, effectiveCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
existingService, err := transaction.GetServiceByID(ctx, store.LockingStrengthUpdate, accountID, service.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -674,6 +737,13 @@ func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.St
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) validateServiceDomain(ctx context.Context, tx store.Store, accountID string, svc *service.Service, cluster string) error {
|
||||
if m.clusterDeriver == nil {
|
||||
return nil
|
||||
}
|
||||
return m.clusterDeriver.ValidateServiceDomain(ctx, tx, accountID, svc.Domain, cluster)
|
||||
}
|
||||
|
||||
// validateL4PortDiffOnClusterDiff checks if custom L4 ports are configured and validates port changes across clusters.
|
||||
// It ensures no port changes if custom ports are unsupported for a given cluster and protocol mode.
|
||||
// Returns an error if validation fails, otherwise returns nil.
|
||||
|
||||
@@ -7,11 +7,10 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
"go.uber.org/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{}
|
||||
|
||||
@@ -433,8 +433,8 @@ func TestDeletePeerService_SourcePeerValidation(t *testing.T) {
|
||||
newProxyServer := func(t *testing.T) *nbgrpc.ProxyServiceServer {
|
||||
t.Helper()
|
||||
tokenStore := nbgrpc.NewOneTimeTokenStore(context.Background(), testCacheStore(t))
|
||||
pkceStore := nbgrpc.NewPKCEVerifierStore(context.Background(), testCacheStore(t))
|
||||
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
singleUseStore := nbgrpc.NewSingleUseStore(context.Background(), testCacheStore(t))
|
||||
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
return srv
|
||||
}
|
||||
|
||||
@@ -655,6 +655,10 @@ func (d *testClusterDeriver) GetClusterDomains() []string {
|
||||
return d.domains
|
||||
}
|
||||
|
||||
func (d *testClusterDeriver) ValidateServiceDomain(context.Context, store.Store, string, string, string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
const (
|
||||
testAccountID = "test-account"
|
||||
testPeerID = "test-peer-1"
|
||||
@@ -722,8 +726,8 @@ func setupIntegrationTest(t *testing.T) (*Manager, store.Store) {
|
||||
}
|
||||
|
||||
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
|
||||
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
|
||||
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
|
||||
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
|
||||
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
@@ -1146,8 +1150,8 @@ func TestDeleteService_DeletesTargets(t *testing.T) {
|
||||
mockAcct := account.NewMockManager(ctrl)
|
||||
|
||||
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
|
||||
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
|
||||
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
|
||||
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
|
||||
|
||||
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"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"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// setupPrivateClusterTest wires the real proxy manager as the capability
|
||||
// provider and connects one proxy to testCluster reporting the given private
|
||||
// capability. A nil private connects no proxy, so the capability is unreported.
|
||||
func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) {
|
||||
t.Helper()
|
||||
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
|
||||
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
mgr.capabilities = proxyMgr
|
||||
|
||||
if private != nil {
|
||||
connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private})
|
||||
}
|
||||
|
||||
return mgr, testStore
|
||||
}
|
||||
|
||||
func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) {
|
||||
t.Helper()
|
||||
_, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func clusterTarget() *rpservice.Target {
|
||||
return &rpservice.Target{
|
||||
TargetId: testCluster,
|
||||
TargetType: rpservice.TargetTypeCluster,
|
||||
Host: "backend.lan",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
Enabled: true,
|
||||
Options: rpservice.TargetOptions{DirectUpstream: true},
|
||||
}
|
||||
}
|
||||
|
||||
func directUpstreamPeerTarget() *rpservice.Target {
|
||||
return &rpservice.Target{
|
||||
TargetId: testPeerID,
|
||||
TargetType: rpservice.TargetTypePeer,
|
||||
Host: "backend.lan",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
Enabled: true,
|
||||
Options: rpservice.TargetOptions{DirectUpstream: true},
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateService_PrivateClusterTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
private *bool
|
||||
target *rpservice.Target
|
||||
wantErr string
|
||||
}{
|
||||
{name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()},
|
||||
{name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()},
|
||||
{name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
{name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, tc.private)
|
||||
|
||||
svc := newTestService("app.test.netbird.io")
|
||||
svc.Targets = []*rpservice.Target{tc.target}
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
|
||||
|
||||
services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, listErr)
|
||||
|
||||
if tc.wantErr == "" {
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, services, 1, "the service should be persisted")
|
||||
return
|
||||
}
|
||||
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.wantErr)
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "the caller must receive a typed error")
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument")
|
||||
assert.Empty(t, services, "a rejected service must not be persisted")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A cluster where only some proxies run in private mode must not accept these
|
||||
// targets: the mapping is delivered to every proxy in the cluster, so the
|
||||
// non-private ones would serve the target from their host network as well.
|
||||
func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
secondCaps *proxy.Capabilities
|
||||
}{
|
||||
{name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}},
|
||||
{name: "second proxy predates capability reporting", secondCaps: nil},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} {
|
||||
t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
|
||||
connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps)
|
||||
|
||||
svc := newTestService("app.test.netbird.io")
|
||||
svc.Targets = []*rpservice.Target{target}
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
|
||||
require.Error(t, err, "a cluster with a non-private proxy must not accept the target")
|
||||
assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled")
|
||||
|
||||
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, services, "a rejected service must not be persisted")
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, _ := setupPrivateClusterTest(t, boolPtr(false))
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err, "a peer target without direct upstream must not need a private cluster")
|
||||
}
|
||||
|
||||
func TestUpdateService_PrivateClusterTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
target *rpservice.Target
|
||||
wantErr string
|
||||
}{
|
||||
{name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(false))
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := newTestService("app.test.netbird.io")
|
||||
updated.ID = created.ID
|
||||
updated.AccountID = testAccountID
|
||||
updated.Targets = []*rpservice.Target{tc.target}
|
||||
|
||||
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.wantErr)
|
||||
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, stored.Targets, 1)
|
||||
assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged")
|
||||
assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := newTestService("app.test.netbird.io")
|
||||
updated.ID = created.ID
|
||||
updated.AccountID = testAccountID
|
||||
updated.Targets = []*rpservice.Target{clusterTarget()}
|
||||
|
||||
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, stored.Targets, 1)
|
||||
assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored")
|
||||
}
|
||||
|
||||
func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
// No ClusterAllProxiesPrivate expectation: a lookup would fail the test.
|
||||
mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)}
|
||||
|
||||
targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}}
|
||||
require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster))
|
||||
}
|
||||
@@ -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"`
|
||||
@@ -388,6 +390,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 {
|
||||
@@ -961,8 +964,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 {
|
||||
@@ -985,6 +988,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)
|
||||
@@ -1019,6 +1050,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",
|
||||
|
||||
@@ -3,15 +3,16 @@ package manager
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/affectedpeers"
|
||||
"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"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups)
|
||||
var snap *affectedpeers.Snapshot
|
||||
change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain)
|
||||
if err != nil {
|
||||
@@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
if err = transaction.CreateZone(ctx, zone); err != nil {
|
||||
return fmt.Errorf("failed to create zone: %w", err)
|
||||
return fmt.Errorf("create zone: %w", err)
|
||||
}
|
||||
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return fmt.Errorf("increment network serial: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta())
|
||||
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return zone, nil
|
||||
}
|
||||
|
||||
@@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
return nil, status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get zone: %w", err)
|
||||
}
|
||||
|
||||
if zone.Domain != updatedZone.Domain {
|
||||
return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
|
||||
}
|
||||
|
||||
zone.Name = updatedZone.Name
|
||||
zone.Enabled = updatedZone.Enabled
|
||||
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
|
||||
zone.DistributionGroups = updatedZone.DistributionGroups
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get zone: %w", err)
|
||||
}
|
||||
|
||||
if zone.Domain != updatedZone.Domain {
|
||||
return status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
|
||||
}
|
||||
|
||||
oldGroups := zone.DistributionGroups
|
||||
zone.Name = updatedZone.Name
|
||||
zone.Enabled = updatedZone.Enabled
|
||||
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
|
||||
zone.DistributionGroups = updatedZone.DistributionGroups
|
||||
|
||||
for _, groupID := range zone.DistributionGroups {
|
||||
_, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID)
|
||||
if err != nil {
|
||||
@@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
if err = transaction.UpdateZone(ctx, zone); err != nil {
|
||||
return fmt.Errorf("failed to update zone: %w", err)
|
||||
return fmt.Errorf("update zone: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return fmt.Errorf("increment network serial: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta())
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return zone, nil
|
||||
}
|
||||
@@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
|
||||
return status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get zone: %w", err)
|
||||
}
|
||||
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
var eventsToStore []func()
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get zone: %w", err)
|
||||
}
|
||||
|
||||
// Load before delete: the post-delete state no longer references the groups.
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get records: %w", err)
|
||||
@@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
|
||||
event()
|
||||
}
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,11 +9,11 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/affectedpeers"
|
||||
"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"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
}
|
||||
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL)
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
@@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to create dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return record, nil
|
||||
}
|
||||
@@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
|
||||
var zone *zones.Zone
|
||||
var record *records.Record
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
@@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to update dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return record, nil
|
||||
}
|
||||
@@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
|
||||
var record *records.Record
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
@@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to delete dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -3,10 +3,7 @@ package networkmapdb
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"strings"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
@@ -48,7 +45,7 @@ func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network: %w", err))
|
||||
}
|
||||
peers, proxyPeers, err := tx.GetPeers(ctx, accountId)
|
||||
peers, _, err := tx.GetPeers(ctx, accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get peers: %w", err))
|
||||
}
|
||||
@@ -80,10 +77,6 @@ func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
services, err := tx.GetPrivateServices(ctx, accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
proxyTargetedDomainResourceIDs, err := tx.GetProxyTargetedDomainResourceIDs(ctx, accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get proxy targeted domain resources: %w", err))
|
||||
@@ -113,7 +106,7 @@ func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId
|
||||
GroupIDToUserIDs: groupsToUserIds,
|
||||
NetworkXIDToPublicID: networkXIDToPublicID, // TODO (dmitri) maybe we can switch to public ids everywhere?
|
||||
AppliedZoneCandidates: dnsZones,
|
||||
PrivateServiceCandidates: buildPrivateServiceCandidates(services, domains, proxyPeers),
|
||||
Domains: TwinProxyDomains(domains),
|
||||
PostureCheckXIDToPublicID: postureCheckXIDToPublicID,
|
||||
ProxyTargetedDomainResourceIDs: proxyTargetedDomainResourceIDs,
|
||||
}
|
||||
@@ -154,94 +147,6 @@ func toSliceOfPtrs[T any](all []T) []*T {
|
||||
return toret
|
||||
}
|
||||
|
||||
func serviceDomainZone(svc Service, ds []Domain) string {
|
||||
if domainFromSuffix(svc.Domain.String, svc.ProxyCluster.String) {
|
||||
return svc.ProxyCluster.String
|
||||
}
|
||||
|
||||
var zoneName string
|
||||
for _, domain := range ds {
|
||||
if domain.TargetCluster.String != svc.ProxyCluster.String {
|
||||
continue
|
||||
}
|
||||
if domainFromSuffix(svc.Domain.String, domain.Domain.String) && len(domain.Domain.String) > len(zoneName) {
|
||||
zoneName = domain.Domain.String
|
||||
}
|
||||
}
|
||||
|
||||
return zoneName
|
||||
}
|
||||
|
||||
func domainFromSuffix(domain, suffix string) bool {
|
||||
if suffix == "" {
|
||||
return false
|
||||
}
|
||||
return domain == suffix || strings.HasSuffix(domain, "."+suffix)
|
||||
}
|
||||
|
||||
func buildPrivateServiceCandidates(svcs []Service, domains []Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate {
|
||||
var out []networkmap.PrivateServiceCandidate
|
||||
|
||||
if len(proxyPeersByCluster) == 0 {
|
||||
return out
|
||||
}
|
||||
|
||||
for _, svc := range svcs {
|
||||
if !svc.Enabled.Bool || !svc.Private.Bool {
|
||||
continue
|
||||
}
|
||||
if len(svc.AccessGroups) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
domainZone := serviceDomainZone(svc, domains)
|
||||
if domainZone == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// this is implied when domainZone != "", but for maintainability's sake the check is explicit
|
||||
// TODO (dmitri) make this an invariant
|
||||
if svc.Domain.String == "" {
|
||||
continue
|
||||
}
|
||||
var records []nmdata.SimpleRecord
|
||||
for _, proxyPeer := range proxyPeersByCluster[svc.ProxyCluster.String] {
|
||||
if record, ok := recordForProxyPeer(svc.Domain.String, proxyPeer.IP); ok {
|
||||
records = append(records, record)
|
||||
}
|
||||
}
|
||||
if len(records) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
out = append(out, networkmap.PrivateServiceCandidate{
|
||||
AccessGroups: svc.AccessGroups,
|
||||
Zone: nmdata.CustomZone{
|
||||
Domain: dns.Fqdn(domainZone),
|
||||
Records: records,
|
||||
NonAuthoritative: true,
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
func recordForProxyPeer(fqdn string, ip netip.Addr) (nmdata.SimpleRecord, bool) {
|
||||
if !ip.IsValid() {
|
||||
return nmdata.SimpleRecord{}, false
|
||||
}
|
||||
|
||||
return nmdata.SimpleRecord{
|
||||
Name: dns.Fqdn(fqdn),
|
||||
Type: int(dns.TypeA),
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: ip.String(),
|
||||
}, true
|
||||
}
|
||||
|
||||
func buildResourcePolicies(networkResources []nmdata.NetworkResource,
|
||||
policies []nmdata.Policy,
|
||||
resourceToGroupIdx map[string]map[string]any,
|
||||
|
||||
@@ -1,253 +1,12 @@
|
||||
package networkmapdb
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestDomainFromSuffix(t *testing.T) {
|
||||
assert.False(t, domainFromSuffix("test", ""))
|
||||
assert.False(t, domainFromSuffix("test", "suffix")) // domain != suffix
|
||||
assert.True(t, domainFromSuffix("test", "test")) // domain == suffix
|
||||
assert.False(t, domainFromSuffix("test.anothersuffix", "suffix")) // domain doesn't contain suffix
|
||||
assert.True(t, domainFromSuffix("test.suffix", "suffix")) // domain contains suffix
|
||||
}
|
||||
|
||||
func TestServiceDomainZone(t *testing.T) {
|
||||
// shortcut -- service's domain is a subomain of proxy cluster
|
||||
assert.Equal(t, "cluster",
|
||||
serviceDomainZone(
|
||||
Service{
|
||||
Domain: sql.NullString{Valid: true, String: "test.cluster"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
[]Domain{}))
|
||||
assert.Equal(t, "a.b", serviceDomainZone(
|
||||
Service{
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
[]Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "a-cluster"}},
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "b"}},
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}}, // should return this domain, as it's the longest match
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "b-cluster"}},
|
||||
}))
|
||||
// service and domain clusters don't match
|
||||
assert.Empty(t, serviceDomainZone(
|
||||
Service{
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "c-cluster"}},
|
||||
[]Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
}))
|
||||
// service domain is empty
|
||||
assert.Empty(t, serviceDomainZone(
|
||||
Service{
|
||||
Domain: sql.NullString{Valid: false, String: ""},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
[]Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
}))
|
||||
}
|
||||
|
||||
func TestRecordForProxyPeer(t *testing.T) {
|
||||
record, ok := recordForProxyPeer("test.cluster", netip.MustParseAddr("127.0.0.1"))
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, nmdata.SimpleRecord{
|
||||
Name: "test.cluster.",
|
||||
Type: 1,
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: "127.0.0.1",
|
||||
}, record)
|
||||
|
||||
// invalid address
|
||||
var addr netip.Addr
|
||||
_, ok = recordForProxyPeer("test.cluster", addr)
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
var empty []networkmap.PrivateServiceCandidate
|
||||
|
||||
// empty proxyPeersByCluster results in empty []PrivateServiceCandidates
|
||||
func TestBuildPrivateServiceCandidates_EmptyProxyPeers(t *testing.T) {
|
||||
assert.Equal(t, empty, buildPrivateServiceCandidates([]Service{}, []Domain{}, nil))
|
||||
}
|
||||
|
||||
// disabled service returns an empty result
|
||||
func TestBuildPrivateServiceCandidates_DisabledService(t *testing.T) {
|
||||
assert.Equal(t, empty,
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: false},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
// non-private service results in empty []PrivateServiceCandidates
|
||||
func TestBuildPrivateServiceCandidates_PublicService(t *testing.T) {
|
||||
assert.Equal(t, empty,
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: false},
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
// empty AccessList results in empty []PrivateServiceCandidates
|
||||
func TestBuildPrivateServiceCandidates_EmptyAccessList(t *testing.T) {
|
||||
assert.Equal(t, empty,
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
// empty TragetCluster results in empty []PrivateServiceCandidates
|
||||
func TestBuildPrivateServiceCandidates_EmptyTargetCluster(t *testing.T) {
|
||||
assert.Equal(t, empty,
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: ""},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
func TestBuildPrivateServiceCandidates_EmptyServiceDomain(t *testing.T) {
|
||||
assert.Equal(t, empty,
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
Domain: sql.NullString{Valid: true, String: ""},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
func TestBuildPrivateServiceCandidates_HappyPath(t *testing.T) {
|
||||
assert.Equal(t, []networkmap.PrivateServiceCandidate{
|
||||
{
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Zone: nmdata.CustomZone{
|
||||
Domain: "a.b.",
|
||||
SearchDomainDisabled: true,
|
||||
NonAuthoritative: true,
|
||||
Records: []nmdata.SimpleRecord{
|
||||
{
|
||||
Name: "test.a.b.",
|
||||
Type: 1,
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
Name: "test.a.b.",
|
||||
Type: 1,
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: "127.0.0.2",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Zone: nmdata.CustomZone{
|
||||
Domain: "c.d.",
|
||||
SearchDomainDisabled: true,
|
||||
NonAuthoritative: true,
|
||||
Records: []nmdata.SimpleRecord{
|
||||
{
|
||||
Name: "test.c.d.",
|
||||
Type: 1,
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: "127.0.0.3",
|
||||
},
|
||||
{
|
||||
Name: "test.c.d.",
|
||||
Type: 1,
|
||||
Class: "IN",
|
||||
TTL: 5,
|
||||
RData: "127.0.0.4",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
buildPrivateServiceCandidates([]Service{
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Domain: sql.NullString{Valid: true, String: "test.a.b"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "cluster"}},
|
||||
{Enabled: sql.NullBool{Valid: true, Bool: true},
|
||||
Private: sql.NullBool{Valid: true, Bool: true},
|
||||
AccessGroups: []string{"group-1", "group-2"},
|
||||
Domain: sql.NullString{Valid: true, String: "test.c.d"},
|
||||
ProxyCluster: sql.NullString{Valid: true, String: "a-cluster"}},
|
||||
}, []Domain{
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "a.b"}},
|
||||
{TargetCluster: sql.NullString{Valid: true, String: "a-cluster"},
|
||||
Domain: sql.NullString{Valid: true, String: "c.d"}},
|
||||
},
|
||||
map[string][]*nmdata.Peer{
|
||||
"cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.1")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.2")}},
|
||||
"a-cluster": {&nmdata.Peer{IP: netip.MustParseAddr("127.0.0.3")}, &nmdata.Peer{IP: netip.MustParseAddr("127.0.0.4")}},
|
||||
}))
|
||||
}
|
||||
|
||||
// disabled network resource shouldn't be in the resulting map
|
||||
func TestBuildResourcePolicies_DisabledNetworkResource(t *testing.T) {
|
||||
networkResources := []nmdata.NetworkResource{
|
||||
|
||||
@@ -15,6 +15,7 @@ const (
|
||||
from zones
|
||||
left join records as r on r.zone_id = zones.id
|
||||
where zones.account_id=$1 and zones.enabled
|
||||
order by zones.id
|
||||
`
|
||||
)
|
||||
|
||||
|
||||
@@ -302,6 +302,7 @@ func ConvertToNmdataPeers(peers []Peer) ([]nmdata.Peer, map[string][]*nmdata.Pee
|
||||
}
|
||||
dp.ProxyMeta.Cluster = p.ProxyMetaCluster.String
|
||||
// This is only used to build private service candidates, not connected peers are skipped
|
||||
dp.Connected = p.PeerStatusConnected.Bool
|
||||
if dp.ProxyMeta.Embedded && p.PeerStatusConnected.Bool {
|
||||
clusterToPeerIdx[p.ProxyMetaCluster.String] = append(clusterToPeerIdx[p.ProxyMetaCluster.String], &dp)
|
||||
}
|
||||
@@ -470,3 +471,16 @@ func ConvertToNmdataPolicy(policies []Policy) ([]nmdata.Policy, map[string]map[s
|
||||
|
||||
return toret, policyToDestinationResourceIdx, policyToDestinationGroupIdx, nil
|
||||
}
|
||||
|
||||
// TwinProxyDomains converts registered reverse-proxy domain rows to their slim
|
||||
// twins, so private-service zone apex resolution runs on the twin.
|
||||
func TwinProxyDomains(domains []Domain) []nmdata.ProxyDomain {
|
||||
if len(domains) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]nmdata.ProxyDomain, 0, len(domains))
|
||||
for _, d := range domains {
|
||||
out = append(out, nmdata.ProxyDomain{Domain: d.Domain.String, TargetCluster: d.TargetCluster.String})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
@@ -11,9 +11,11 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
// Outer join: a groupless router must survive.
|
||||
GetNetworkRouterQuery = `
|
||||
select public_id, peer, network_id, masquerade, metric, enabled, peer_groups, group_peers.peer_id
|
||||
from network_routers, json_each(peer_groups)
|
||||
from network_routers
|
||||
left join json_each(network_routers.peer_groups) on true
|
||||
left join group_peers on group_peers.account_id=? and group_peers.group_id=json_each.value
|
||||
where network_routers.account_id=?
|
||||
`
|
||||
|
||||
@@ -42,17 +42,20 @@ func (sc *SqliteStoreConn) GetAllowedUsers(ctx context.Context, accountId string
|
||||
userIdIdx := make(map[string]struct{})
|
||||
groupIdToUserIds := make(map[string][]string)
|
||||
for _, user := range users {
|
||||
for _, allgid := range allGroupIds {
|
||||
groupIdToUserIds[allgid] = append(groupIdToUserIds[allgid], user.ID)
|
||||
}
|
||||
userIdIdx[user.ID] = struct{}{}
|
||||
autogroups := make([]string, 0)
|
||||
if user.AutoGroups == nil {
|
||||
continue
|
||||
}
|
||||
if err := json.Unmarshal(user.AutoGroups, &autogroups); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
userIdIdx[user.ID] = struct{}{}
|
||||
for _, groupId := range autogroups {
|
||||
groupIdToUserIds[groupId] = append(groupIdToUserIds[groupId], user.ID)
|
||||
}
|
||||
for _, allgid := range allGroupIds {
|
||||
groupIdToUserIds[allgid] = append(groupIdToUserIds[allgid], user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return userIdIdx, groupIdToUserIds, nil
|
||||
|
||||
@@ -21,8 +21,6 @@ import (
|
||||
"google.golang.org/grpc/credentials"
|
||||
"google.golang.org/grpc/keepalive"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
@@ -33,17 +31,19 @@ import (
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
nbContext "github.com/netbirdio/netbird/management/server/context"
|
||||
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/management/server/idp"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
mgmtProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
"github.com/netbirdio/netbird/util/crypt"
|
||||
)
|
||||
|
||||
@@ -75,8 +75,8 @@ func (s *BaseServer) Metrics() telemetry.AppMetrics {
|
||||
|
||||
// CacheStore returns a shared cache store backed by Redis or in-memory depending on the environment.
|
||||
// All consumers should reuse this store to avoid creating multiple Redis connections.
|
||||
func (s *BaseServer) CacheStore() cachestore.StoreInterface {
|
||||
return Create(s, func() cachestore.StoreInterface {
|
||||
func (s *BaseServer) CacheStore() nbcache.Store {
|
||||
return Create(s, func() nbcache.Store {
|
||||
cs, err := nbcache.NewStore(context.Background(), nbcache.DefaultStoreMaxTimeout, nbcache.DefaultStoreCleanupInterval, nbcache.DefaultStoreMaxConn)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to create shared cache store: %v", err)
|
||||
@@ -85,9 +85,20 @@ func (s *BaseServer) CacheStore() cachestore.StoreInterface {
|
||||
})
|
||||
}
|
||||
|
||||
// DBConn opens the database connection shared by the store and the domain repositories.
|
||||
func (s *BaseServer) DBConn() *db.Conn {
|
||||
return Create(s, func() *db.Conn {
|
||||
conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to open database connection: %v", err)
|
||||
}
|
||||
return conn
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) Store() store.Store {
|
||||
return Create(s, func() store.Store {
|
||||
store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false)
|
||||
store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to create store: %v", err)
|
||||
}
|
||||
@@ -113,7 +124,8 @@ func (s *BaseServer) NetworkMapStore() *networkmapdb.NetworkMapDBStoreImpl {
|
||||
s.Config.StoreConfig.Engine,
|
||||
s.Config.Datadir,
|
||||
s.IntegratedValidator(),
|
||||
s.SettingsManager())
|
||||
s.SettingsManager(),
|
||||
)
|
||||
// networkmap db store supports postgres and sqlite backends only
|
||||
// for other backends a fallback is used, so NotSupportedStoreEngineError
|
||||
// is not a fatal error
|
||||
@@ -147,7 +159,7 @@ func (s *BaseServer) EventStore() activity.Store {
|
||||
|
||||
func (s *BaseServer) APIHandler() http.Handler {
|
||||
return Create(s, func() http.Handler {
|
||||
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
|
||||
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager(), nil)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to create API handler: %v", err)
|
||||
}
|
||||
@@ -171,10 +183,10 @@ func (s *BaseServer) Router() *mux.Router {
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
|
||||
return Create(s, func() *middleware.APIRateLimiter {
|
||||
cfg, enabled := middleware.RateLimiterConfigFromEnv()
|
||||
limiter := middleware.NewAPIRateLimiter(cfg)
|
||||
func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter {
|
||||
return Create(s, func() *ratelimit.APIRateLimiter {
|
||||
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
|
||||
limiter := ratelimit.NewAPIRateLimiter(cfg)
|
||||
limiter.SetEnabled(enabled)
|
||||
return limiter
|
||||
})
|
||||
@@ -182,24 +194,7 @@ func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
|
||||
|
||||
func (s *BaseServer) GRPCServer() *grpc.Server {
|
||||
return Create(s, func() *grpc.Server {
|
||||
trustedPeers := s.Config.ReverseProxy.TrustedPeers
|
||||
defaultTrustedPeers := []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0"), netip.MustParsePrefix("::/0")}
|
||||
if len(trustedPeers) == 0 || slices.Equal[[]netip.Prefix](trustedPeers, defaultTrustedPeers) {
|
||||
log.WithContext(context.Background()).Warn("TrustedPeers are configured to default value '0.0.0.0/0', '::/0'. This allows connection IP spoofing.")
|
||||
trustedPeers = defaultTrustedPeers
|
||||
}
|
||||
trustedHTTPProxies := s.Config.ReverseProxy.TrustedHTTPProxies
|
||||
trustedProxiesCount := s.Config.ReverseProxy.TrustedHTTPProxiesCount
|
||||
if len(trustedHTTPProxies) > 0 && trustedProxiesCount > 0 {
|
||||
log.WithContext(context.Background()).Warn("TrustedHTTPProxies and TrustedHTTPProxiesCount both are configured. " +
|
||||
"This is not recommended way to extract X-Forwarded-For. Consider using one of these options.")
|
||||
}
|
||||
realipOpts := []realip.Option{
|
||||
realip.WithTrustedPeers(trustedPeers),
|
||||
realip.WithTrustedProxies(trustedHTTPProxies),
|
||||
realip.WithTrustedProxiesCount(trustedProxiesCount),
|
||||
realip.WithHeaders([]string{realip.XForwardedFor, realip.XRealIp}),
|
||||
}
|
||||
realipOpts := realIPOptions(s.Config.ReverseProxy)
|
||||
proxyUnary, proxyStream, proxyAuthClose := nbgrpc.NewProxyAuthInterceptors(s.Store())
|
||||
s.proxyAuthClose = proxyAuthClose
|
||||
gRPCOpts := []grpc.ServerOption{
|
||||
@@ -253,7 +248,7 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
|
||||
|
||||
func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer {
|
||||
return Create(s, func() *nbgrpc.ProxyServiceServer {
|
||||
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
|
||||
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.SingleUseStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
|
||||
s.AfterInit(func(s *BaseServer) {
|
||||
proxyService.SetServiceManager(s.ServiceManager())
|
||||
proxyService.SetActivityManager(s.ProxyActivityManager())
|
||||
@@ -310,9 +305,9 @@ func (s *BaseServer) ProxyTokenStore() *nbgrpc.OneTimeTokenStore {
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore {
|
||||
return Create(s, func() *nbgrpc.PKCEVerifierStore {
|
||||
return nbgrpc.NewPKCEVerifierStore(context.Background(), s.CacheStore())
|
||||
func (s *BaseServer) SingleUseStore() *nbgrpc.SingleUseStore {
|
||||
return Create(s, func() *nbgrpc.SingleUseStore {
|
||||
return nbgrpc.NewSingleUseStore(context.Background(), s.CacheStore())
|
||||
})
|
||||
}
|
||||
|
||||
@@ -325,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager {
|
||||
|
||||
func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
|
||||
return Create(s, func() accesslogs.Manager {
|
||||
accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager())
|
||||
accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager())
|
||||
accessLogManager.StartPeriodicCleanup(
|
||||
context.Background(),
|
||||
s.Config.ReverseProxy.AccessLogRetentionDays,
|
||||
@@ -335,7 +330,7 @@ func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
|
||||
})
|
||||
}
|
||||
|
||||
func loadTLSConfig(certFile string, certKey string) (*tls.Config, error) {
|
||||
func loadTLSConfig(certFile, certKey string) (*tls.Config, error) {
|
||||
// Load server's certificate and private key
|
||||
serverCert, err := tls.LoadX509KeyPair(certFile, certKey)
|
||||
if err != nil {
|
||||
@@ -382,3 +377,37 @@ func streamInterceptor(
|
||||
wrapped.WrappedContext = context.WithValue(ctx, nbContext.RequestIDKey, reqID)
|
||||
return handler(srv, wrapped)
|
||||
}
|
||||
|
||||
// realIPOptions builds the real-IP middleware options.
|
||||
//
|
||||
// Empty TrustedPeers trusts all IPv4 and IPv6 sources. Configure TrustedPeers
|
||||
// with the reverse proxy address or network.
|
||||
//
|
||||
// X-Forwarded-For takes precedence over X-Real-IP.
|
||||
func realIPOptions(cfg nbconfig.ReverseProxy) []realip.Option {
|
||||
trustedPeers := cfg.TrustedPeers
|
||||
if len(trustedPeers) == 0 {
|
||||
trustedPeers = []netip.Prefix{
|
||||
netip.MustParsePrefix("0.0.0.0/0"),
|
||||
netip.MustParsePrefix("::/0"),
|
||||
}
|
||||
}
|
||||
if idx := slices.IndexFunc(trustedPeers, func(p netip.Prefix) bool { return p.Bits() == 0 }); idx >= 0 {
|
||||
log.WithContext(context.Background()).Warnf("TrustedPeers contains the default route %s, which trusts "+
|
||||
"X-Forwarded-For from every client and allows connection IP spoofing. Set TrustedPeers to the address "+
|
||||
"of your reverse proxy.", trustedPeers[idx])
|
||||
}
|
||||
if cfg.TrustedHTTPProxiesCount > 0 {
|
||||
log.WithContext(context.Background()).Warn(
|
||||
"TrustedHTTPProxiesCount skips X-Forwarded-For entries by position before TrustedHTTPProxies filters by address. " +
|
||||
"An incorrect count may skip the real client IP and produce an incorrect source address.",
|
||||
)
|
||||
}
|
||||
|
||||
return []realip.Option{
|
||||
realip.WithTrustedPeers(trustedPeers),
|
||||
realip.WithTrustedProxies(cfg.TrustedHTTPProxies),
|
||||
realip.WithTrustedProxiesCount(cfg.TrustedHTTPProxiesCount),
|
||||
realip.WithHeaders([]string{realip.XForwardedFor, realip.XRealIp}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -103,6 +103,7 @@ func (s *BaseServer) AccountManager() account.Manager {
|
||||
|
||||
s.AfterInit(func(s *BaseServer) {
|
||||
accountManager.SetServiceManager(s.ServiceManager())
|
||||
accountManager.AddAccountDeletionHook(s.AgentNetworkManager().RemoveAccountGateway)
|
||||
})
|
||||
|
||||
return accountManager
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
)
|
||||
|
||||
const (
|
||||
realIPProbeMethod = "/netbird.test.RealIPProbe/Probe"
|
||||
realIPProbeStreamMethod = "/netbird.test.RealIPProbe/ProbeStream"
|
||||
)
|
||||
|
||||
// realIPProbe records the real IP the middleware derived for each call.
|
||||
type realIPProbe struct {
|
||||
got chan string
|
||||
}
|
||||
|
||||
func (p *realIPProbe) record(ctx context.Context) {
|
||||
addr, _ := realip.FromContext(ctx)
|
||||
p.got <- addr.String()
|
||||
}
|
||||
|
||||
func (p *realIPProbe) wait(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
select {
|
||||
case got := <-p.got:
|
||||
return got
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for probe")
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func startProbeServer(t *testing.T, cfg nbconfig.ReverseProxy) (*grpc.ClientConn, *realIPProbe) {
|
||||
t.Helper()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
probe := &realIPProbe{got: make(chan string, 1)}
|
||||
opts := realIPOptions(cfg)
|
||||
srv := grpc.NewServer(
|
||||
grpc.ChainUnaryInterceptor(realip.UnaryServerInterceptorOpts(opts...)),
|
||||
grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(opts...)),
|
||||
)
|
||||
srv.RegisterService(&grpc.ServiceDesc{
|
||||
ServiceName: "netbird.test.RealIPProbe",
|
||||
HandlerType: (*any)(nil),
|
||||
Methods: []grpc.MethodDesc{{
|
||||
MethodName: "Probe",
|
||||
Handler: func(_ any, ctx context.Context, dec func(any) error, interceptor grpc.UnaryServerInterceptor) (any, error) {
|
||||
req := new(emptypb.Empty)
|
||||
if err := dec(req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
handler := func(ctx context.Context, _ any) (any, error) {
|
||||
probe.record(ctx)
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
if interceptor == nil {
|
||||
return handler(ctx, req)
|
||||
}
|
||||
return interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: realIPProbeMethod}, handler)
|
||||
},
|
||||
}},
|
||||
Streams: []grpc.StreamDesc{{
|
||||
StreamName: "ProbeStream",
|
||||
ServerStreams: true,
|
||||
Handler: func(_ any, stream grpc.ServerStream) error {
|
||||
probe.record(stream.Context())
|
||||
return nil
|
||||
},
|
||||
}},
|
||||
}, probe)
|
||||
|
||||
go func() { _ = srv.Serve(listener) }()
|
||||
t.Cleanup(srv.Stop)
|
||||
|
||||
conn, err := grpc.NewClient(listener.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
return conn, probe
|
||||
}
|
||||
|
||||
func callUnary(t *testing.T, conn *grpc.ClientConn, probe *realIPProbe, kv ...string) string {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ctx = metadata.AppendToOutgoingContext(ctx, kv...)
|
||||
require.NoError(t, conn.Invoke(ctx, realIPProbeMethod, &emptypb.Empty{}, &emptypb.Empty{}))
|
||||
|
||||
return probe.wait(t)
|
||||
}
|
||||
|
||||
func callStream(t *testing.T, conn *grpc.ClientConn, probe *realIPProbe, kv ...string) string {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ctx = metadata.AppendToOutgoingContext(ctx, kv...)
|
||||
desc := &grpc.StreamDesc{StreamName: "ProbeStream", ServerStreams: true}
|
||||
stream, err := conn.NewStream(ctx, desc, realIPProbeStreamMethod)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, stream.CloseSend())
|
||||
require.ErrorIs(t, stream.RecvMsg(&emptypb.Empty{}), io.EOF)
|
||||
|
||||
return probe.wait(t)
|
||||
}
|
||||
|
||||
func assertRealIP(t *testing.T, cfg nbconfig.ReverseProxy, want string, kv ...string) {
|
||||
t.Helper()
|
||||
|
||||
conn, probe := startProbeServer(t, cfg)
|
||||
t.Run("unary", func(t *testing.T) {
|
||||
assert.Equal(t, want, callUnary(t, conn, probe, kv...))
|
||||
})
|
||||
t.Run("stream", func(t *testing.T) {
|
||||
assert.Equal(t, want, callStream(t, conn, probe, kv...))
|
||||
})
|
||||
}
|
||||
|
||||
func TestRealIPDefaultTrustsForwardedHeaders(t *testing.T) {
|
||||
assertRealIP(t, nbconfig.ReverseProxy{}, "203.0.113.44",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPUntrustedPeerIgnoresForwardedHeaders(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{TrustedPeers: []netip.Prefix{netip.MustParsePrefix("10.9.8.7/32")}}
|
||||
|
||||
assertRealIP(t, cfg, "127.0.0.1",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPTrustedPeerHonoursForwardedHeaders(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{TrustedPeers: []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")}}
|
||||
|
||||
assertRealIP(t, cfg, "203.0.113.44",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPReadsXRealIPWhenProxyCountSkipsForwardedFor(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{
|
||||
TrustedPeers: []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")},
|
||||
TrustedHTTPProxiesCount: 1,
|
||||
}
|
||||
|
||||
t.Run("no X-Forwarded-For", func(t *testing.T) {
|
||||
assertRealIP(t, cfg, "203.0.113.44", realip.XRealIp, "203.0.113.44")
|
||||
})
|
||||
t.Run("single-entry X-Forwarded-For", func(t *testing.T) {
|
||||
assertRealIP(t, cfg, "198.51.100.7",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "198.51.100.7",
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -23,6 +23,8 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/idp"
|
||||
"github.com/netbirdio/netbird/management/server/metrics"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/lifecycle"
|
||||
"github.com/netbirdio/netbird/shared/profiling"
|
||||
"github.com/netbirdio/netbird/util/wsproxy"
|
||||
wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
@@ -36,6 +38,8 @@ const (
|
||||
DefaultSelfHostedDomain = "netbird.selfhosted"
|
||||
|
||||
ContainerKeyBaseServer = "baseServer"
|
||||
|
||||
applicationName = "management"
|
||||
)
|
||||
|
||||
type Server interface {
|
||||
@@ -66,7 +70,8 @@ type BaseServer struct {
|
||||
disableLegacyManagementPort bool
|
||||
autoResolveDomains bool
|
||||
|
||||
proxyAuthClose func()
|
||||
proxyAuthClose func()
|
||||
domainCleanupStop func()
|
||||
|
||||
// grpcExtensions holds additional gRPC services, interceptors, and shutdown
|
||||
// hooks registered by external modules via RegisterGRPCExtension. Populated
|
||||
@@ -74,12 +79,15 @@ type BaseServer struct {
|
||||
grpcExtensions []GRPCExtension
|
||||
|
||||
listener net.Listener
|
||||
tlsConfig *tls.Config
|
||||
certManager *autocert.Manager
|
||||
update *version.Update
|
||||
|
||||
errCh chan error
|
||||
wg sync.WaitGroup
|
||||
cancel context.CancelFunc
|
||||
|
||||
lifecycle.StopHandlers
|
||||
}
|
||||
|
||||
// Config holds the configuration parameters for creating a new server
|
||||
@@ -94,6 +102,7 @@ type Config struct {
|
||||
DisableGeoliteUpdate bool
|
||||
UserDeleteFromIDPEnabled bool
|
||||
AutoResolveDomains bool
|
||||
TLSConfig *tls.Config
|
||||
}
|
||||
|
||||
// NewServer initializes and configures a new Server instance
|
||||
@@ -110,9 +119,13 @@ func NewServer(cfg *Config) *BaseServer {
|
||||
disableLegacyManagementPort: cfg.DisableLegacyManagementPort,
|
||||
mgmtMetricsPort: cfg.MgmtMetricsPort,
|
||||
autoResolveDomains: cfg.AutoResolveDomains,
|
||||
tlsConfig: cfg.TLSConfig,
|
||||
}
|
||||
s.container[ContainerKeyBaseServer] = s
|
||||
|
||||
stopProfiling := profiling.Start(applicationName)
|
||||
s.OnStop(stopProfiling)
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
@@ -122,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) {
|
||||
|
||||
// Start begins listening for HTTP requests on the configured address
|
||||
func (s *BaseServer) Start(ctx context.Context) error {
|
||||
if err := s.start(ctx); err != nil {
|
||||
s.RunStopHandlers()
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *BaseServer) start(ctx context.Context) error {
|
||||
srvCtx, cancel := context.WithCancel(ctx)
|
||||
s.cancel = cancel
|
||||
s.errCh = make(chan error, 4)
|
||||
@@ -139,21 +160,9 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
}
|
||||
s.EphemeralManager().LoadInitialPeers(srvCtx)
|
||||
|
||||
var tlsConfig *tls.Config
|
||||
tlsEnabled := false
|
||||
if s.Config.HttpConfig.LetsEncryptDomain != "" {
|
||||
s.certManager, err = encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed creating LetsEncrypt cert manager: %v", err)
|
||||
}
|
||||
tlsEnabled = true
|
||||
} else if s.Config.HttpConfig.CertFile != "" && s.Config.HttpConfig.CertKey != "" {
|
||||
tlsConfig, err = loadTLSConfig(s.Config.HttpConfig.CertFile, s.Config.HttpConfig.CertKey)
|
||||
if err != nil {
|
||||
log.WithContext(srvCtx).Errorf("cannot load TLS credentials: %v", err)
|
||||
return err
|
||||
}
|
||||
tlsEnabled = true
|
||||
tlsEnabled, err := s.setupTLS(srvCtx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
installationID, err := getInstallationID(srvCtx, s.Store())
|
||||
@@ -215,8 +224,8 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
log.WithContext(ctx).Infof("running HTTP server (LetsEncrypt challenge handler): %s", cml.Addr().String())
|
||||
s.serveHTTP(ctx, cml, s.certManager.HTTPHandler(nil))
|
||||
}
|
||||
case tlsConfig != nil:
|
||||
s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), tlsConfig)
|
||||
case s.tlsConfig != nil:
|
||||
s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), s.tlsConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed creating TLS listener on port %d: %v", s.mgmtPort, err)
|
||||
}
|
||||
@@ -236,14 +245,60 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
s.update.SetOnUpdateListener(func() {
|
||||
log.WithContext(ctx).Infof("your management version, \"%s\", is outdated, a new management version is available. Learn more here: https://github.com/netbirdio/netbird/releases", version.NetbirdVersion())
|
||||
})
|
||||
s.startDomainCleanup(srvCtx)
|
||||
|
||||
return nil
|
||||
}
|
||||
func (s *BaseServer) startDomainCleanup(ctx context.Context) {
|
||||
if s.domainCleanupStop != nil {
|
||||
return
|
||||
}
|
||||
mgr := s.ReverseProxyDomainManager()
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
done := make(chan struct{})
|
||||
s.domainCleanupStop = func() {
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
go func() {
|
||||
defer close(done)
|
||||
mgr.RunValidationCleanup(ctx)
|
||||
}()
|
||||
}
|
||||
|
||||
// setupTLS resolves the listener's TLS source: an injected config wins over the HttpConfig certificate settings
|
||||
func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) {
|
||||
switch {
|
||||
case s.tlsConfig != nil:
|
||||
return true, nil
|
||||
case s.Config.HttpConfig.LetsEncryptDomain != "":
|
||||
certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed creating LetsEncrypt cert manager: %v", err)
|
||||
}
|
||||
s.certManager = certManager
|
||||
return true, nil
|
||||
case s.Config.HttpConfig.CertFile != "" && s.Config.HttpConfig.CertKey != "":
|
||||
tlsConfig, err := loadTLSConfig(s.Config.HttpConfig.CertFile, s.Config.HttpConfig.CertKey)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("cannot load TLS credentials: %v", err)
|
||||
return false, err
|
||||
}
|
||||
s.tlsConfig = tlsConfig
|
||||
return true, nil
|
||||
default:
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Stop attempts a graceful shutdown, waiting up to 5 seconds for active connections to finish
|
||||
func (s *BaseServer) Stop() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
defer s.RunStopHandlers()
|
||||
if s.domainCleanupStop != nil {
|
||||
s.domainCleanupStop()
|
||||
}
|
||||
|
||||
s.IntegratedValidator().Stop(ctx)
|
||||
if s.GeoLocationManager() != nil {
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTransactionTimeout = 5 * time.Minute
|
||||
connMaxLifetime = time.Hour
|
||||
connMaxIdleTime = 3 * time.Minute
|
||||
)
|
||||
|
||||
// TxMetrics receives the duration of every committed top-level transaction.
|
||||
type TxMetrics interface {
|
||||
CountTransactionDuration(duration time.Duration)
|
||||
}
|
||||
|
||||
// Conn is the database connection shared by all repositories: one gorm handle,
|
||||
// the pgx pool of a Postgres deployment and the engine they talk to.
|
||||
type Conn struct {
|
||||
db *gorm.DB
|
||||
pool *pgxpool.Pool
|
||||
engine Engine
|
||||
txTimeout time.Duration
|
||||
metrics TxMetrics
|
||||
}
|
||||
|
||||
// NewConn takes ownership of an open gorm handle and pool once it returns
|
||||
// without error, applying the connection limits and transaction timeout
|
||||
// configured through the environment.
|
||||
func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) {
|
||||
sqlDB, err := gormDB.DB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
txTimeout := defaultTransactionTimeout
|
||||
if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" {
|
||||
if parsed, err := time.ParseDuration(v); err == nil {
|
||||
txTimeout = parsed
|
||||
}
|
||||
}
|
||||
log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout)
|
||||
|
||||
conns := runtime.NumCPU()
|
||||
configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS"))
|
||||
connsConfigured := err == nil
|
||||
if connsConfigured {
|
||||
conns = configuredConns
|
||||
}
|
||||
if engine == SqliteStoreEngine {
|
||||
if connsConfigured {
|
||||
log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1")
|
||||
}
|
||||
conns = 1
|
||||
}
|
||||
|
||||
sqlDB.SetMaxOpenConns(conns)
|
||||
sqlDB.SetMaxIdleConns(conns)
|
||||
sqlDB.SetConnMaxLifetime(connMaxLifetime)
|
||||
sqlDB.SetConnMaxIdleTime(connMaxIdleTime)
|
||||
|
||||
log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v",
|
||||
conns, conns, connMaxLifetime, connMaxIdleTime)
|
||||
|
||||
return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil
|
||||
}
|
||||
|
||||
// DB returns the handle a query must run on: the transaction when tx is set,
|
||||
// otherwise the shared connection.
|
||||
func (c *Conn) DB(tx *Tx) *gorm.DB {
|
||||
if tx != nil {
|
||||
return tx.db
|
||||
}
|
||||
return c.db
|
||||
}
|
||||
|
||||
// Pool returns the pgx pool for read paths that bypass gorm. It is nil on
|
||||
// engines other than Postgres and inside a transaction, where the pool would
|
||||
// not see the uncommitted writes.
|
||||
func (c *Conn) Pool(tx *Tx) *pgxpool.Pool {
|
||||
if tx != nil {
|
||||
return nil
|
||||
}
|
||||
return c.pool
|
||||
}
|
||||
|
||||
func (c *Conn) Engine() Engine {
|
||||
return c.engine
|
||||
}
|
||||
|
||||
// SetTxMetrics registers the sink that receives transaction durations.
|
||||
func (c *Conn) SetTxMetrics(metrics TxMetrics) {
|
||||
c.metrics = metrics
|
||||
}
|
||||
|
||||
// AutoMigrate creates or updates the tables of the given models.
|
||||
func (c *Conn) AutoMigrate(models ...any) error {
|
||||
return c.db.AutoMigrate(models...)
|
||||
}
|
||||
|
||||
// Close releases the gorm connection and the pgx pool.
|
||||
func (c *Conn) Close() error {
|
||||
if c.pool != nil {
|
||||
c.pool.Close()
|
||||
}
|
||||
sqlDB, err := c.db.DB()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get db: %w", err)
|
||||
}
|
||||
return sqlDB.Close()
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type testRow struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
Name string
|
||||
}
|
||||
|
||||
func openTestConn(t *testing.T) *Conn {
|
||||
t.Helper()
|
||||
conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { require.NoError(t, conn.Close()) })
|
||||
require.NoError(t, conn.AutoMigrate(&testRow{}))
|
||||
return conn
|
||||
}
|
||||
|
||||
func countRows(t *testing.T, conn *Conn) int64 {
|
||||
t.Helper()
|
||||
var count int64
|
||||
require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error)
|
||||
return count
|
||||
}
|
||||
|
||||
func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) {
|
||||
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s")
|
||||
conn := openTestConn(t)
|
||||
assert.Equal(t, time.Second, conn.txTimeout)
|
||||
assert.Equal(t, SqliteStoreEngine, conn.Engine())
|
||||
}
|
||||
|
||||
func TestRunInTx_CommitsOnSuccess(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
|
||||
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.EqualValues(t, 1, countRows(t, conn))
|
||||
}
|
||||
|
||||
func TestRunInTx_RollsBackOnError(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
failure := errors.New("boom")
|
||||
|
||||
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
|
||||
return failure
|
||||
})
|
||||
require.ErrorIs(t, err, failure)
|
||||
assert.EqualValues(t, 0, countRows(t, conn))
|
||||
}
|
||||
|
||||
func TestRunInTx_RollsBackOnPanic(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
|
||||
require.Panics(t, func() {
|
||||
_ = conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
|
||||
panic("boom")
|
||||
})
|
||||
})
|
||||
assert.EqualValues(t, 0, countRows(t, conn))
|
||||
}
|
||||
|
||||
func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) {
|
||||
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms")
|
||||
conn := openTestConn(t)
|
||||
|
||||
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
|
||||
})
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
assert.EqualValues(t, 0, countRows(t, conn))
|
||||
}
|
||||
|
||||
func TestRunInTx_ReportsDurationToMetrics(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
metrics := &recordingMetrics{}
|
||||
conn.SetTxMetrics(metrics)
|
||||
|
||||
require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil }))
|
||||
assert.Equal(t, 1, metrics.calls)
|
||||
}
|
||||
|
||||
func TestConn_DBSelectsTransactionHandle(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
assert.Same(t, conn.db, conn.DB(nil))
|
||||
|
||||
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
assert.Same(t, tx.db, conn.DB(tx))
|
||||
assert.NotSame(t, conn.db, conn.DB(tx))
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) {
|
||||
conn := openTestConn(t)
|
||||
conn.pool = &pgxpool.Pool{}
|
||||
defer func() { conn.pool = nil }()
|
||||
|
||||
assert.Same(t, conn.pool, conn.Pool(nil))
|
||||
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
|
||||
assert.Nil(t, conn.Pool(tx))
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
type recordingMetrics struct {
|
||||
calls int
|
||||
}
|
||||
|
||||
func (m *recordingMetrics) CountTransactionDuration(time.Duration) {
|
||||
m.calls++
|
||||
}
|
||||
|
||||
func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) {
|
||||
t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7")
|
||||
|
||||
gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig())
|
||||
require.NoError(t, err)
|
||||
conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
sqlDB, err := conn.DB(nil).DB()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections)
|
||||
|
||||
sqliteDB, err := openTestConn(t).DB(nil).DB()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections)
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package dbtest
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
)
|
||||
|
||||
// NewConn opens a fresh SQLite database in a temporary directory, migrates the
|
||||
// given models and closes the connection when the test ends. It ignores
|
||||
// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never
|
||||
// touched, and is safe to call from parallel tests.
|
||||
func NewConn(t testing.TB, models ...any) *db.Conn {
|
||||
t.Helper()
|
||||
conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
require.NoError(t, conn.AutoMigrate(models...))
|
||||
return conn
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package dbtest
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/shared/db"
|
||||
)
|
||||
|
||||
func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) {
|
||||
override := filepath.Join(t.TempDir(), "configured.db")
|
||||
t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override)
|
||||
|
||||
conn := NewConn(t)
|
||||
|
||||
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
|
||||
_, err := os.Stat(override)
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
}
|
||||
|
||||
func TestNewConn_Parallel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
conn := NewConn(t)
|
||||
|
||||
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
package db
|
||||
|
||||
// Engine identifies the SQL engine behind a Conn.
|
||||
type Engine string
|
||||
|
||||
const (
|
||||
SqliteStoreEngine Engine = "sqlite"
|
||||
PostgresStoreEngine Engine = "postgres"
|
||||
MysqlStoreEngine Engine = "mysql"
|
||||
)
|
||||
@@ -0,0 +1,12 @@
|
||||
package db
|
||||
|
||||
// LockingStrength is the row lock a query holds until its transaction ends.
|
||||
type LockingStrength string
|
||||
|
||||
const (
|
||||
LockingStrengthUpdate LockingStrength = "UPDATE"
|
||||
LockingStrengthShare LockingStrength = "SHARE"
|
||||
LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE"
|
||||
LockingStrengthKeyShare LockingStrength = "KEY SHARE"
|
||||
LockingStrengthNone LockingStrength = "NONE"
|
||||
)
|
||||
@@ -0,0 +1,176 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
// SqliteFileName is the default SQLite database file inside the data directory.
|
||||
const SqliteFileName = "store.db"
|
||||
|
||||
// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths
|
||||
// that bypass gorm.
|
||||
type PoolConfig struct {
|
||||
MaxConns int32
|
||||
MinConns int32
|
||||
MaxConnLifetime time.Duration
|
||||
HealthCheckPeriod time.Duration
|
||||
}
|
||||
|
||||
var DefaultPoolConfig = PoolConfig{
|
||||
MaxConns: 30,
|
||||
MinConns: 1,
|
||||
MaxConnLifetime: 60 * time.Minute,
|
||||
HealthCheckPeriod: time.Minute,
|
||||
}
|
||||
|
||||
// GormConfig is the configuration every engine is opened with.
|
||||
func GormConfig() *gorm.Config {
|
||||
return &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
CreateBatchSize: 400,
|
||||
}
|
||||
}
|
||||
|
||||
// OpenSqlite opens the SQLite database in dataDir, or the file named by
|
||||
// NB_STORE_ENGINE_SQLITE_FILE.
|
||||
func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) {
|
||||
storeFile := SqliteFileName
|
||||
if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" {
|
||||
storeFile = envFile
|
||||
}
|
||||
return OpenSqliteFile(ctx, dataDir, storeFile)
|
||||
}
|
||||
|
||||
// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir
|
||||
// when relative. storeFile may carry SQLite URI query parameters.
|
||||
func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) {
|
||||
// Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc")
|
||||
filePath, query, hasQuery := strings.Cut(storeFile, "?")
|
||||
|
||||
connStr := filePath
|
||||
if !filepath.IsAbs(filePath) {
|
||||
connStr = filepath.Join(dataDir, filePath)
|
||||
}
|
||||
|
||||
// Compose query parameters. User-provided ?_busy_timeout (or its mattn alias
|
||||
// ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at
|
||||
// most that long on a lock instead of blocking the only Go-side connection.
|
||||
// mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so
|
||||
// the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared
|
||||
// stays the default on non-Windows for the same reason as before.
|
||||
parsed, _ := url.ParseQuery(query)
|
||||
var defaults []string
|
||||
if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" {
|
||||
defaults = append(defaults, "_busy_timeout=30000")
|
||||
}
|
||||
if !hasQuery && runtime.GOOS != "windows" {
|
||||
// To avoid `The process cannot access the file because it is being used by another process` on Windows
|
||||
defaults = append(defaults, "cache=shared")
|
||||
}
|
||||
parts := defaults
|
||||
if hasQuery {
|
||||
parts = append(parts, query)
|
||||
}
|
||||
if len(parts) > 0 {
|
||||
connStr += "?" + strings.Join(parts, "&")
|
||||
}
|
||||
|
||||
gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := NewConn(ctx, gormDB, SqliteStoreEngine, nil)
|
||||
if err != nil {
|
||||
closeGorm(gormDB)
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool.
|
||||
func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) {
|
||||
gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pgxPool, err := newPgxPool(ctx, dsn, pool)
|
||||
if err != nil {
|
||||
closeGorm(gormDB)
|
||||
return nil, err
|
||||
}
|
||||
conn, err := NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool)
|
||||
if err != nil {
|
||||
pgxPool.Close()
|
||||
closeGorm(gormDB)
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// MysqlDSN adds the connection parameters every MySQL handle needs, keeping
|
||||
// the options already present in dsn.
|
||||
func MysqlDSN(dsn string) string {
|
||||
separator := "?"
|
||||
if strings.Contains(dsn, "?") {
|
||||
separator = "&"
|
||||
}
|
||||
return dsn + separator + "charset=utf8&parseTime=True&loc=Local"
|
||||
}
|
||||
|
||||
// OpenMysql opens a MySQL database through gorm.
|
||||
func OpenMysql(ctx context.Context, dsn string) (*Conn, error) {
|
||||
gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := NewConn(ctx, gormDB, MysqlStoreEngine, nil)
|
||||
if err != nil {
|
||||
closeGorm(gormDB)
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) {
|
||||
config, err := pgxpool.ParseConfig(dsn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to parse database config: %w", err)
|
||||
}
|
||||
|
||||
config.MaxConns = cfg.MaxConns
|
||||
config.MinConns = cfg.MinConns
|
||||
config.MaxConnLifetime = cfg.MaxConnLifetime
|
||||
config.HealthCheckPeriod = cfg.HealthCheckPeriod
|
||||
|
||||
pool, err := pgxpool.NewWithConfig(ctx, config)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to create connection pool: %w", err)
|
||||
}
|
||||
|
||||
if err := pool.Ping(ctx); err != nil {
|
||||
pool.Close()
|
||||
return nil, fmt.Errorf("unable to ping database: %w", err)
|
||||
}
|
||||
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
func closeGorm(gormDB *gorm.DB) {
|
||||
if sqlDB, err := gormDB.DB(); err == nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestMysqlDSN(t *testing.T) {
|
||||
assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db"))
|
||||
assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true"))
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Tx is an open transaction handed to repository calls; nil means autocommit.
|
||||
type Tx struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
// RunInTx runs fn in one transaction that commits when fn returns nil and rolls
|
||||
// back otherwise, bounded by the configured transaction timeout.
|
||||
func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error {
|
||||
timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout)
|
||||
defer cancel()
|
||||
|
||||
startTime := time.Now()
|
||||
tx := c.db.WithContext(timeoutCtx).Begin()
|
||||
if tx.Error != nil {
|
||||
return tx.Error
|
||||
}
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
tx.Rollback()
|
||||
panic(r)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := c.applyStatementTimeouts(tx); err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
err := c.withForeignKeyChecksDisabled(tx, func() error {
|
||||
return fn(&Tx{db: tx})
|
||||
})
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime)
|
||||
return err
|
||||
}
|
||||
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime)
|
||||
return err
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime))
|
||||
if c.metrics != nil {
|
||||
c.metrics.CountTransactionDuration(time.Since(startTime))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error {
|
||||
if c.engine != PostgresStoreEngine {
|
||||
return nil
|
||||
}
|
||||
if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil {
|
||||
return fmt.Errorf("failed to set statement timeout: %w", err)
|
||||
}
|
||||
if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil {
|
||||
return fmt.Errorf("failed to set lock timeout: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids
|
||||
// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is
|
||||
// session-scoped and survives a rollback, so it is turned back on whenever fn
|
||||
// returns or panics; otherwise the pooled connection would keep it disabled.
|
||||
func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) {
|
||||
if c.engine != MysqlStoreEngine {
|
||||
return fn()
|
||||
}
|
||||
if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil {
|
||||
return fmt.Errorf("failed to disable FK checks: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error
|
||||
if restoreErr == nil {
|
||||
return
|
||||
}
|
||||
if err == nil {
|
||||
err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr)
|
||||
return
|
||||
}
|
||||
log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr)
|
||||
}()
|
||||
return fn()
|
||||
}
|
||||
|
||||
func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) {
|
||||
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
|
||||
log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack())
|
||||
}
|
||||
}
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/ssh/auth"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
@@ -37,7 +36,7 @@ func ToComponentSyncResponse(
|
||||
components *types.NetworkMapComponents,
|
||||
proxyPatch *types.NetworkMap,
|
||||
dnsName string,
|
||||
checks []*posture.Checks,
|
||||
checks []*nmdata.PostureChecks,
|
||||
settings *nmdata.AccountSettingsInfo,
|
||||
extraSettings *types.ExtraSettings,
|
||||
peerGroups []string,
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
@@ -154,7 +153,7 @@ func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, se
|
||||
return peerConfig
|
||||
}
|
||||
|
||||
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
|
||||
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*nmdata.PostureChecks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
|
||||
// IPv6 data in AllowedIPs and SourcePrefixes wildcard expansion depends on
|
||||
// whether the target peer supports IPv6. Routes and firewall rules are already
|
||||
// filtered at the source (network map builder).
|
||||
|
||||
@@ -14,6 +14,7 @@ const (
|
||||
baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit
|
||||
reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban
|
||||
metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer
|
||||
maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one
|
||||
)
|
||||
|
||||
type lfConfig struct {
|
||||
@@ -21,6 +22,7 @@ type lfConfig struct {
|
||||
baseBlockDuration time.Duration
|
||||
reconnLimitForBan int
|
||||
metaChangeLimit int
|
||||
maxBanLevel int
|
||||
}
|
||||
|
||||
func initCfg() *lfConfig {
|
||||
@@ -29,6 +31,7 @@ func initCfg() *lfConfig {
|
||||
baseBlockDuration: baseBlockDuration,
|
||||
reconnLimitForBan: reconnLimitForBan,
|
||||
metaChangeLimit: metaChangeLimit,
|
||||
maxBanLevel: maxBanLevel,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
|
||||
return
|
||||
}
|
||||
|
||||
if state.isBanned && now.After(state.banExpiresAt) {
|
||||
if state.isBanned {
|
||||
if now.Before(state.banExpiresAt) {
|
||||
return
|
||||
}
|
||||
state.isBanned = false
|
||||
}
|
||||
|
||||
if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) {
|
||||
quietSince := state.lastSeen
|
||||
if state.banExpiresAt.After(quietSince) {
|
||||
quietSince = state.banExpiresAt
|
||||
}
|
||||
if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) {
|
||||
state.banLevel = 0
|
||||
}
|
||||
|
||||
@@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
|
||||
return
|
||||
}
|
||||
|
||||
if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold {
|
||||
state.sessionStart = now
|
||||
state.sessionCounter = 0
|
||||
}
|
||||
|
||||
state.sessionCounter++
|
||||
if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold {
|
||||
if state.sessionCounter > l.cfg.reconnLimitForBan {
|
||||
state.isBanned = true
|
||||
state.banLevel++
|
||||
if state.banLevel < l.cfg.maxBanLevel {
|
||||
state.banLevel++
|
||||
}
|
||||
|
||||
backoffFactor := math.Pow(2, float64(state.banLevel-1))
|
||||
duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor)
|
||||
|
||||
@@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig {
|
||||
baseBlockDuration: 100 * time.Millisecond,
|
||||
reconnLimitForBan: 3,
|
||||
metaChangeLimit: 2,
|
||||
maxBanLevel: 3,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() {
|
||||
s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset")
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
s.Require().Contains(s.filter.logged, pubKey)
|
||||
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
|
||||
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count")
|
||||
|
||||
for i := 1; i < limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
s.True(s.filter.allowLogin(pubKey, meta))
|
||||
s.False(s.filter.logged[pubKey].isBanned)
|
||||
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
|
||||
s.False(s.filter.allowLogin(pubKey, meta))
|
||||
s.True(s.filter.logged[pubKey].isBanned)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
s.Require().Contains(s.filter.logged, pubKey)
|
||||
s.Require().True(s.filter.logged[pubKey].isBanned)
|
||||
|
||||
expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second))
|
||||
s.filter.logged[pubKey].banExpiresAt = expired
|
||||
s.filter.logged[pubKey].sessionStart = expired
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
|
||||
s.True(s.filter.logged[pubKey].isBanned)
|
||||
s.Equal(2, s.filter.logged[pubKey].banLevel)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
|
||||
for i := 0; i < limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
s.Require().Contains(s.filter.logged, pubKey)
|
||||
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
|
||||
|
||||
for i := 0; i < limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
|
||||
s.True(s.filter.allowLogin(pubKey, meta))
|
||||
s.False(s.filter.logged[pubKey].isBanned)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration)
|
||||
|
||||
s.filter.logged[pubKey] = &peerState{
|
||||
currentHash: meta,
|
||||
isBanned: true,
|
||||
banLevel: 1,
|
||||
banExpiresAt: time.Now().Add(-time.Millisecond),
|
||||
sessionStart: banTime,
|
||||
lastSeen: banTime,
|
||||
}
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
|
||||
s.True(s.filter.logged[pubKey].isBanned)
|
||||
s.Equal(2, s.filter.logged[pubKey].banLevel)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
quiet := 2*s.filter.cfg.baseBlockDuration + time.Second
|
||||
|
||||
s.filter.logged[pubKey] = &peerState{
|
||||
currentHash: meta,
|
||||
banLevel: 2,
|
||||
banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration),
|
||||
lastSeen: time.Now().Add(-2 * quiet),
|
||||
}
|
||||
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period")
|
||||
|
||||
s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet)
|
||||
s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet)
|
||||
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
s.Equal(0, s.filter.logged[pubKey].banLevel)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
maxLevel := s.filter.cfg.maxBanLevel
|
||||
|
||||
s.filter.logged[pubKey] = &peerState{
|
||||
currentHash: meta,
|
||||
banLevel: maxLevel,
|
||||
sessionStart: time.Now(),
|
||||
lastSeen: time.Now(),
|
||||
}
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
|
||||
s.True(s.filter.logged[pubKey].isBanned)
|
||||
s.Equal(maxLevel, s.filter.logged[pubKey].banLevel)
|
||||
expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1)
|
||||
s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond))
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
longAgo := time.Now().Add(-time.Hour)
|
||||
|
||||
s.filter.logged[pubKey] = &peerState{
|
||||
currentHash: meta,
|
||||
sessionCounter: 1,
|
||||
sessionStart: longAgo,
|
||||
lastSeen: longAgo,
|
||||
metaChangeWindowStart: longAgo,
|
||||
metaChangeCounter: 1,
|
||||
}
|
||||
|
||||
s.True(s.filter.allowLogin(pubKey, meta))
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
|
||||
s.True(s.filter.allowLogin(pubKey, meta))
|
||||
s.False(s.filter.logged[pubKey].isBanned)
|
||||
s.Equal(1, s.filter.logged[pubKey].sessionCounter)
|
||||
}
|
||||
|
||||
func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() {
|
||||
pubKey := "PUB_KEY_A"
|
||||
meta := uint64(1)
|
||||
limit := s.filter.cfg.reconnLimitForBan
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
s.Require().Contains(s.filter.logged, pubKey)
|
||||
s.Require().True(s.filter.logged[pubKey].isBanned)
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
s.filter.logged[pubKey].banExpiresAt = expiresAt
|
||||
lastSeen := s.filter.logged[pubKey].lastSeen
|
||||
|
||||
for i := 0; i <= limit; i++ {
|
||||
s.filter.addLogin(pubKey, meta)
|
||||
}
|
||||
|
||||
s.True(s.filter.logged[pubKey].isBanned)
|
||||
s.Equal(1, s.filter.logged[pubKey].banLevel)
|
||||
s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt)
|
||||
s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen)
|
||||
s.Equal(0, s.filter.logged[pubKey].sessionCounter)
|
||||
}
|
||||
|
||||
func BenchmarkHashingMethods(b *testing.B) {
|
||||
meta := nbpeer.PeerSystemMeta{
|
||||
WtVersion: "1.25.1",
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
func PeerUpdateHandlerFactory(
|
||||
peerKey wgtypes.Key,
|
||||
updates chan *network_map.UpdateMessage,
|
||||
secretsManager SecretsManager,
|
||||
srv proto.ManagementService_SyncServer,
|
||||
cleanupfunc func()) *PeerUpdateHandler {
|
||||
return &PeerUpdateHandler{
|
||||
peerKey: peerKey,
|
||||
updates: updates,
|
||||
secretsManager: secretsManager,
|
||||
srv: srv,
|
||||
encrypter: encryption.DefaultEncrypter{},
|
||||
debouncer: NewUpdateDebouncer(1000 * time.Millisecond),
|
||||
cleanupFunc: cleanupfunc,
|
||||
}
|
||||
}
|
||||
|
||||
// PeerUpdateHandler sends updates to the connected peer until the updates channel is closed.
|
||||
// It implements a backpressure mechanism that sends the first update immediately,
|
||||
// then debounces subsequent rapid updates, ensuring only the latest update is sent
|
||||
// after a quiet period.
|
||||
type PeerUpdateHandler struct {
|
||||
peerKey wgtypes.Key
|
||||
updates chan *network_map.UpdateMessage
|
||||
appMetrics telemetry.AppMetrics
|
||||
secretsManager SecretsManager
|
||||
srv syncSender
|
||||
encrypter encryption.Encrypter
|
||||
debouncer Debouncer
|
||||
cleanupFunc func()
|
||||
}
|
||||
|
||||
func (pu *PeerUpdateHandler) WithMetrics(appMetrics telemetry.AppMetrics) *PeerUpdateHandler {
|
||||
pu.appMetrics = appMetrics
|
||||
return pu
|
||||
}
|
||||
|
||||
//go:generate go tool mockgen -source=./peer_update_handler.go -destination=./sync_sender_mock.go -package=grpc
|
||||
type syncSender interface {
|
||||
Send(*proto.EncryptedMessage) error
|
||||
Context() context.Context
|
||||
}
|
||||
|
||||
func (pu *PeerUpdateHandler) HandleUpdates(ctx context.Context) error {
|
||||
log.WithContext(ctx).Tracef("starting to handle updates for peer %s", pu.peerKey.String())
|
||||
|
||||
defer pu.debouncer.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
// condition when there are some updates
|
||||
// todo set the updates channel size to 1
|
||||
case update, open := <-pu.updates:
|
||||
if pu.appMetrics != nil {
|
||||
pu.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(pu.updates) + 1)
|
||||
}
|
||||
|
||||
if !open {
|
||||
log.WithContext(ctx).Debugf("updates channel for peer %s was closed", pu.peerKey.String())
|
||||
pu.cleanupFunc()
|
||||
return nil
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Tracef("received an update for peer %s", pu.peerKey.String())
|
||||
if pu.debouncer.ProcessUpdate(update) {
|
||||
// Send immediately (first update or after quiet period)
|
||||
if err := pu.SendUpdate(ctx, update); err != nil {
|
||||
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Timer expired - quiet period reached, send pending updates if any
|
||||
case <-pu.debouncer.TimerChannel():
|
||||
pendingUpdates := pu.debouncer.GetPendingUpdates()
|
||||
if len(pendingUpdates) == 0 {
|
||||
continue
|
||||
}
|
||||
log.WithContext(ctx).Debugf("sending %d debounced update(s) for peer %s", len(pendingUpdates), pu.peerKey.String())
|
||||
for _, pendingUpdate := range pendingUpdates {
|
||||
if err := pu.SendUpdate(ctx, pendingUpdate); err != nil {
|
||||
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// condition when client <-> server connection has been terminated
|
||||
case <-pu.srv.Context().Done():
|
||||
// happens when connection drops, e.g. client disconnects
|
||||
log.WithContext(ctx).Debugf("stream of peer %s has been closed", pu.peerKey.String())
|
||||
pu.cleanupFunc()
|
||||
return pu.srv.Context().Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (pu *PeerUpdateHandler) SendUpdate(ctx context.Context, update *network_map.UpdateMessage) error {
|
||||
key, err := pu.secretsManager.GetWGKey()
|
||||
if err != nil {
|
||||
pu.cleanupFunc()
|
||||
return status.Errorf(codes.Internal, "failed processing update message")
|
||||
}
|
||||
|
||||
encryptedResp, err := pu.encrypter.EncryptMessage(pu.peerKey, key, update.Update)
|
||||
if err != nil {
|
||||
pu.cleanupFunc()
|
||||
return status.Errorf(codes.Internal, "failed processing update message")
|
||||
}
|
||||
err = pu.srv.Send(&proto.EncryptedMessage{
|
||||
WgPubKey: key.PublicKey().String(),
|
||||
Body: encryptedResp,
|
||||
})
|
||||
if err != nil {
|
||||
pu.cleanupFunc()
|
||||
return status.Errorf(codes.Internal, "failed sending update message")
|
||||
}
|
||||
log.WithContext(ctx).Tracef("sent an update to peer %s", pu.peerKey.String())
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
pb "github.com/golang/protobuf/proto" //nolint
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
)
|
||||
|
||||
func TestSendPeerUpdates_FirstUpdate(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
secretsManager := NewMockSecretsManager(ctrl)
|
||||
updateDebouncer := NewMockDebouncer(ctrl)
|
||||
syncSender := NewMocksyncSender(ctrl)
|
||||
|
||||
pu := PeerUpdateHandler{
|
||||
peerKey: mustGenerateKey(t),
|
||||
updates: make(chan *network_map.UpdateMessage),
|
||||
secretsManager: secretsManager,
|
||||
encrypter: testEncrypter{},
|
||||
debouncer: updateDebouncer,
|
||||
srv: syncSender,
|
||||
cleanupFunc: func() {},
|
||||
}
|
||||
|
||||
msg := network_map.UpdateMessage{
|
||||
Update: &proto.SyncResponse{Version: 1},
|
||||
}
|
||||
|
||||
timeCh := make(chan time.Time)
|
||||
srvCtx := context.TODO()
|
||||
srvKey := mustGenerateKey(t)
|
||||
// mock a first update, should send it right away
|
||||
updateDebouncer.EXPECT().ProcessUpdate(gomock.Eq(&msg)).Return(true)
|
||||
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
|
||||
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
|
||||
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
|
||||
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
|
||||
updateDebouncer.EXPECT().Stop()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
|
||||
pu.updates <- &msg
|
||||
close(pu.updates)
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestSendPeerUpdates_TimerUpdate(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
secretsManager := NewMockSecretsManager(ctrl)
|
||||
updateDebouncer := NewMockDebouncer(ctrl)
|
||||
syncSender := NewMocksyncSender(ctrl)
|
||||
|
||||
pu := PeerUpdateHandler{
|
||||
peerKey: mustGenerateKey(t),
|
||||
updates: make(chan *network_map.UpdateMessage),
|
||||
secretsManager: secretsManager,
|
||||
encrypter: testEncrypter{},
|
||||
debouncer: updateDebouncer,
|
||||
srv: syncSender,
|
||||
cleanupFunc: func() {},
|
||||
}
|
||||
|
||||
msg := network_map.UpdateMessage{
|
||||
Update: &proto.SyncResponse{Version: 1},
|
||||
}
|
||||
|
||||
timeCh := make(chan time.Time)
|
||||
srvCtx := context.TODO()
|
||||
srvKey := mustGenerateKey(t)
|
||||
updateDebouncer.EXPECT().GetPendingUpdates().Return([]*network_map.UpdateMessage{&msg})
|
||||
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
|
||||
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
|
||||
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
|
||||
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
|
||||
updateDebouncer.EXPECT().Stop()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
|
||||
timeCh <- time.Now()
|
||||
close(pu.updates)
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestSendPeerUpdates_ServerContextDone(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
secretsManager := NewMockSecretsManager(ctrl)
|
||||
updateDebouncer := NewMockDebouncer(ctrl)
|
||||
syncSender := NewMocksyncSender(ctrl)
|
||||
|
||||
pu := PeerUpdateHandler{
|
||||
peerKey: mustGenerateKey(t),
|
||||
updates: make(chan *network_map.UpdateMessage),
|
||||
secretsManager: secretsManager,
|
||||
encrypter: testEncrypter{},
|
||||
debouncer: updateDebouncer,
|
||||
srv: syncSender,
|
||||
cleanupFunc: func() {},
|
||||
}
|
||||
|
||||
timeCh := make(chan time.Time)
|
||||
srvCtx, cancel := context.WithCancel(context.TODO())
|
||||
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
|
||||
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
|
||||
updateDebouncer.EXPECT().Stop()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
|
||||
cancel()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func mustGenerateKey(t *testing.T) wgtypes.Key {
|
||||
t.Helper()
|
||||
k, err := wgtypes.GenerateKey()
|
||||
assert.NoError(t, err)
|
||||
return k
|
||||
}
|
||||
|
||||
func mustMarshal(t *testing.T, msg *network_map.UpdateMessage) []byte {
|
||||
t.Helper()
|
||||
r, err := pb.Marshal(msg.Update)
|
||||
assert.NoError(t, err)
|
||||
return r
|
||||
}
|
||||
|
||||
type testEncrypter struct{}
|
||||
|
||||
func (testEncrypter) EncryptMessage(remotePubKey wgtypes.Key, ourPrivateKey wgtypes.Key, message pb.Message) ([]byte, error) {
|
||||
return pb.Marshal(message)
|
||||
}
|
||||
|
||||
type pbMatcher struct {
|
||||
x pb.Message
|
||||
}
|
||||
|
||||
func (pbm pbMatcher) Matches(x any) bool {
|
||||
msg, ok := x.(pb.Message)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return pb.Equal(pbm.x, msg)
|
||||
}
|
||||
|
||||
func (pbm pbMatcher) String() string {
|
||||
return fmt.Sprintf("is equal to %s (%T)", pbm.x, pbm.x)
|
||||
}
|
||||
@@ -1,54 +0,0 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/eko/gocache/lib/v4/cache"
|
||||
"github.com/eko/gocache/lib/v4/store"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// PKCEVerifierStore manages PKCE verifiers for OAuth flows.
|
||||
// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var.
|
||||
type PKCEVerifierStore struct {
|
||||
cache *cache.Cache[string]
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
|
||||
func NewPKCEVerifierStore(ctx context.Context, cacheStore store.StoreInterface) *PKCEVerifierStore {
|
||||
return &PKCEVerifierStore{
|
||||
cache: cache.New[string](cacheStore),
|
||||
ctx: ctx,
|
||||
}
|
||||
}
|
||||
|
||||
// Store saves a PKCE verifier associated with an OAuth state parameter.
|
||||
// The verifier is stored with the specified TTL and will be automatically deleted after expiration.
|
||||
func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error {
|
||||
if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil {
|
||||
return fmt.Errorf("failed to store PKCE verifier: %w", err)
|
||||
}
|
||||
|
||||
log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl)
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadAndDelete retrieves and removes a PKCE verifier for the given state.
|
||||
// Returns the verifier and true if found, or empty string and false if not found.
|
||||
// This enforces single-use semantics for PKCE verifiers.
|
||||
func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) {
|
||||
verifier, err := s.cache.Get(s.ctx, state)
|
||||
if err != nil {
|
||||
log.Debugf("PKCE verifier not found for state")
|
||||
return "", false
|
||||
}
|
||||
|
||||
if err := s.cache.Delete(s.ctx, state); err != nil {
|
||||
log.Warnf("Failed to delete PKCE verifier for state: %v", err)
|
||||
}
|
||||
|
||||
return verifier, true
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user