mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-31 12:01:29 +02:00
account for proxy peers and services
This commit is contained in:
@@ -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 != "" {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user