account for proxy peers and services

This commit is contained in:
pascal
2026-05-26 17:55:18 +02:00
parent 7af7630e5b
commit 68b942722c
5 changed files with 292 additions and 0 deletions

View File

@@ -5,6 +5,7 @@ import (
log "github.com/sirupsen/logrus"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
@@ -30,6 +31,7 @@ func collectPeerChangeAffectedGroups(ctx context.Context, transaction store.Stor
collectAffectedFromNameServers(ctx, transaction, accountID, changedGroupSet, groupSet)
collectAffectedFromDNSSettings(ctx, transaction, accountID, changedGroupSet, groupSet)
collectAffectedFromNetworkRouters(ctx, transaction, accountID, changedGroupSet, changedPeerSet, groupSet, peerSet)
collectAffectedFromProxyServices(ctx, transaction, accountID, changedGroupSet, changedPeerSet, peerSet)
allGroupIDs = setToSlice(groupSet)
directPeerIDs = setToSlice(peerSet)
@@ -139,6 +141,113 @@ func collectAffectedFromNetworkRouters(ctx context.Context, transaction store.St
}
}
// collectAffectedFromProxyServices handles policies that are synthesized at
// network-map computation time by Account.InjectProxyPolicies. Those policies
// connect proxy peers (peer.ProxyMeta.Embedded == true) to service targets and
// never reach the database, so the other collectors cannot see them.
func collectAffectedFromProxyServices(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}, peerSet map[string]struct{}) {
if len(changedGroupSet) == 0 && len(changedPeerSet) == 0 {
return
}
services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID)
if err != nil {
log.WithContext(ctx).Errorf("failed to get services for affected group resolution: %v", err)
return
}
if len(services) == 0 {
return
}
proxyByCluster, err := transaction.GetEmbeddedProxyPeerIDsByCluster(ctx, accountID)
if err != nil {
log.WithContext(ctx).Errorf("failed to get embedded proxy peers for affected group resolution: %v", err)
return
}
if len(proxyByCluster) == 0 {
return
}
expandedPeerSet := changedPeerSet
expanded := false
expand := func() {
if expanded {
return
}
expanded = true
if len(changedGroupSet) == 0 {
return
}
ids, err := transaction.GetPeerIDsByGroups(ctx, accountID, setToSlice(changedGroupSet))
if err != nil {
log.WithContext(ctx).Errorf("failed to expand changed groups to peers for service resolution: %v", err)
return
}
if len(ids) == 0 {
return
}
merged := make(map[string]struct{}, len(changedPeerSet)+len(ids))
for id := range changedPeerSet {
merged[id] = struct{}{}
}
for _, id := range ids {
merged[id] = struct{}{}
}
expandedPeerSet = merged
}
for _, svc := range services {
if svc == nil {
continue
}
proxyPeers := proxyByCluster[svc.ProxyCluster]
if len(proxyPeers) == 0 {
continue
}
expand()
matched := false
for _, pid := range proxyPeers {
if _, ok := expandedPeerSet[pid]; ok {
matched = true
break
}
}
if !matched {
for _, target := range svc.Targets {
if target.TargetType != rpservice.TargetTypePeer || target.TargetId == "" {
continue
}
if _, ok := expandedPeerSet[target.TargetId]; ok {
matched = true
break
}
}
}
if !matched {
continue
}
log.WithContext(ctx).Tracef("collectAffectedFromProxyServices: service %s (cluster=%s) matched; folding %d proxy peers and target peers",
svc.ID, svc.ProxyCluster, len(proxyPeers))
for _, pid := range proxyPeers {
peerSet[pid] = struct{}{}
}
for _, target := range svc.Targets {
if target.TargetType == rpservice.TargetTypePeer && target.TargetId != "" {
peerSet[target.TargetId] = struct{}{}
}
}
}
}
func collectPolicyDirectPeers(policy *types.Policy, peerSet map[string]struct{}) {
for _, rule := range policy.Rules {
if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" {

View File

@@ -12,6 +12,7 @@ import (
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
nbdns "github.com/netbirdio/netbird/dns"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
@@ -1976,3 +1977,141 @@ func addPeerToAccount(t *testing.T, manager *DefaultAccountManager, _, setupKeyK
require.NoError(t, err)
return peer
}
// markPeerAsProxy flips an existing peer's ProxyMeta to mark it as an embedded
// proxy peer in the given cluster.
func markPeerAsProxy(t *testing.T, s store.Store, accountID, peerID, cluster string) {
t.Helper()
ctx := context.Background()
peer, err := s.GetPeerByID(ctx, store.LockingStrengthNone, accountID, peerID)
require.NoError(t, err)
peer.ProxyMeta = nbpeer.ProxyMeta{Embedded: true, Cluster: cluster}
require.NoError(t, s.SavePeer(ctx, accountID, peer))
}
// createServiceWithTargets persists a service with the given cluster and targets
// directly in the store, bypassing the proxy-service manager (which would also
// run cluster derivation and trigger UpdateAccountPeers).
func createServiceWithTargets(t *testing.T, s store.Store, accountID, cluster string, targets []*rpservice.Target) *rpservice.Service {
t.Helper()
svc := &rpservice.Service{
AccountID: accountID,
Name: fmt.Sprintf("svc-%s", cluster),
Domain: fmt.Sprintf("%s.example.com", cluster),
ProxyCluster: cluster,
Enabled: true,
Mode: "tcp",
Targets: targets,
}
svc.InitNewRecord()
for _, target := range targets {
target.AccountID = accountID
target.ServiceID = svc.ID
}
require.NoError(t, s.CreateService(context.Background(), svc))
return svc
}
func TestCollectAffectedFromProxyServices_TargetPeerChanged(t *testing.T) {
manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t)
ctx := context.Background()
cluster := "cluster-a"
markPeerAsProxy(t, s, accountID, peerIDs[0], cluster)
createServiceWithTargets(t, s, accountID, cluster, []*rpservice.Target{
{TargetType: rpservice.TargetTypePeer, TargetId: peerIDs[1], Enabled: true, Port: 80, Protocol: "tcp"},
})
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, nil, []string{peerIDs[1]})
assert.Contains(t, directPeers, peerIDs[0], "proxy peer must be refreshed when its target peer changes")
assert.Contains(t, directPeers, peerIDs[1], "target peer must be refreshed")
}
func TestCollectAffectedFromProxyServices_ProxyPeerChanged(t *testing.T) {
manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t)
ctx := context.Background()
cluster := "cluster-a"
markPeerAsProxy(t, s, accountID, peerIDs[0], cluster)
createServiceWithTargets(t, s, accountID, cluster, []*rpservice.Target{
{TargetType: rpservice.TargetTypePeer, TargetId: peerIDs[1], Enabled: true, Port: 80, Protocol: "tcp"},
{TargetType: rpservice.TargetTypePeer, TargetId: peerIDs[2], Enabled: true, Port: 80, Protocol: "tcp"},
})
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, nil, []string{peerIDs[0]})
assert.Contains(t, directPeers, peerIDs[0], "changed proxy peer is itself refreshed")
assert.Contains(t, directPeers, peerIDs[1], "target peer 1 must be refreshed when proxy peer changes")
assert.Contains(t, directPeers, peerIDs[2], "target peer 2 must be refreshed when proxy peer changes")
}
func TestCollectAffectedFromProxyServices_GroupContainingTargetPeerChanged(t *testing.T) {
manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
ctx := context.Background()
cluster := "cluster-a"
markPeerAsProxy(t, s, accountID, peerIDs[0], cluster)
createServiceWithTargets(t, s, accountID, cluster, []*rpservice.Target{
{TargetType: rpservice.TargetTypePeer, TargetId: peerIDs[1], Enabled: true, Port: 80, Protocol: "tcp"},
})
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, []string{groupIDs[1]}, nil)
assert.Contains(t, directPeers, peerIDs[0], "proxy peer must be refreshed when a group containing its target peer changes")
assert.Contains(t, directPeers, peerIDs[1], "target peer must be refreshed")
}
func TestCollectAffectedFromProxyServices_NoServices(t *testing.T) {
manager, _, accountID, peerIDs, _ := setupAffectedPeersTest(t)
ctx := context.Background()
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, nil, []string{peerIDs[0]})
assert.NotContains(t, directPeers, peerIDs[0], "no services means no proxy contribution")
}
func TestCollectAffectedFromProxyServices_DisabledServiceStillMatches(t *testing.T) {
manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t)
ctx := context.Background()
cluster := "cluster-a"
markPeerAsProxy(t, s, accountID, peerIDs[0], cluster)
svc := &rpservice.Service{
AccountID: accountID,
Name: "disabled-svc",
Domain: "disabled.example.com",
ProxyCluster: cluster,
Enabled: false,
Mode: "tcp",
Targets: []*rpservice.Target{
{TargetType: rpservice.TargetTypePeer, TargetId: peerIDs[1], Enabled: false, Port: 80, Protocol: "tcp"},
},
}
svc.InitNewRecord()
for _, target := range svc.Targets {
target.AccountID = accountID
target.ServiceID = svc.ID
}
require.NoError(t, s.CreateService(ctx, svc))
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, nil, []string{peerIDs[1]})
assert.Contains(t, directPeers, peerIDs[0], "disabled service should still trigger a refresh so peers are ready when re-enabled")
assert.Contains(t, directPeers, peerIDs[1], "disabled target should still trigger a refresh")
}
func TestCollectAffectedFromProxyServices_NonPeerTargetType(t *testing.T) {
manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t)
ctx := context.Background()
cluster := "cluster-a"
markPeerAsProxy(t, s, accountID, peerIDs[0], cluster)
createServiceWithTargets(t, s, accountID, cluster, []*rpservice.Target{
{TargetType: rpservice.TargetTypeHost, TargetId: "10.0.0.1", Host: "10.0.0.1", Enabled: true, Port: 80, Protocol: "tcp"},
})
_, directPeers := collectPeerChangeAffectedGroups(ctx, manager.Store, accountID, nil, []string{peerIDs[0]})
assert.Contains(t, directPeers, peerIDs[0], "host target service still refreshes its proxy peer when the proxy peer changes")
assert.NotContains(t, directPeers, "10.0.0.1", "non-peer target ids must not appear as affected peer IDs")
}

View File

@@ -4883,6 +4883,34 @@ func (s *SqlStore) GetGroupIDsByPeerIDs(ctx context.Context, accountID string, p
return groupIDs, nil
}
// GetEmbeddedProxyPeerIDsByCluster returns peer IDs of all embedded proxy peers
// in the account, grouped by their ProxyCluster. The map is nil when no embedded
// proxy peers exist.
func (s *SqlStore) GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error) {
type row struct {
ID string
Cluster string
}
var rows []row
result := s.db.Model(&nbpeer.Peer{}).
Select("id, proxy_meta_cluster AS cluster").
Where("account_id = ? AND proxy_meta_embedded = ?", accountID, true).
Scan(&rows)
if result.Error != nil {
return nil, status.Errorf(status.Internal, "failed to get embedded proxy peers: %s", result.Error)
}
if len(rows) == 0 {
return nil, nil
}
out := make(map[string][]string, len(rows))
for _, r := range rows {
out[r.Cluster] = append(out[r.Cluster], r.ID)
}
return out, nil
}
func (s *SqlStore) GetUserIDByPeerKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (string, error) {
tx := s.db
if lockStrength != LockingStrengthNone {

View File

@@ -164,6 +164,7 @@ type Store interface {
GetPeersByGroupIDs(ctx context.Context, accountID string, groupIDs []string) ([]*nbpeer.Peer, error)
GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error)
GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error)
GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error)
GetAccountPeersWithExpiration(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error)
GetAccountPeersWithInactivity(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error)
GetAllEphemeralPeers(ctx context.Context, lockStrength LockingStrength) ([]*nbpeer.Peer, error)

View File

@@ -1941,6 +1941,21 @@ func (mr *MockStoreMockRecorder) GetGroupIDsByPeerIDs(ctx, accountID, peerIDs in
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetGroupIDsByPeerIDs", reflect.TypeOf((*MockStore)(nil).GetGroupIDsByPeerIDs), ctx, accountID, peerIDs)
}
// GetEmbeddedProxyPeerIDsByCluster mocks base method.
func (m *MockStore) GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetEmbeddedProxyPeerIDsByCluster", ctx, accountID)
ret0, _ := ret[0].(map[string][]string)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetEmbeddedProxyPeerIDsByCluster indicates an expected call of GetEmbeddedProxyPeerIDsByCluster.
func (mr *MockStoreMockRecorder) GetEmbeddedProxyPeerIDsByCluster(ctx, accountID interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEmbeddedProxyPeerIDsByCluster", reflect.TypeOf((*MockStore)(nil).GetEmbeddedProxyPeerIDsByCluster), ctx, accountID)
}
// GetPeersByIDs mocks base method.
func (m *MockStore) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*peer.Peer, error) {
m.ctrl.T.Helper()