From 611a9291cd99e4a16f26a68ef28d48bedf0edf40 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Fri, 28 Aug 2026 15:39:48 +0200 Subject: [PATCH] [management] fix posture check flip evaluation for affected peers calc (#7347) --- .../network_map/controller/controller.go | 38 +--- .../controller/posture_twin_test.go | 38 ++++ .../controllers/network_map/interface.go | 6 +- .../controllers/network_map/interface_mock.go | 10 +- .../grpc/components_envelope_response.go | 3 +- .../internals/shared/grpc/conversion.go | 3 +- management/internals/shared/grpc/server.go | 12 +- management/server/account.go | 3 +- management/server/account/manager.go | 11 +- management/server/account/manager_mock.go | 17 +- management/server/account_test.go | 58 ++++- .../affected_peers_router_paths_test.go | 42 ++++ .../server/affected_peers_router_test.go | 6 + management/server/mock_server/account_mock.go | 17 +- management/server/peer.go | 25 ++- management/server/peer_posture_test.go | 183 ++++++++++++++++ management/server/peer_test.go | 10 +- .../server/posture/affects_posture_test.go | 202 ------------------ management/server/posture/checks.go | 41 ---- .../server/types/account_networkmapdata.go | 14 +- .../management/networkmap/nmdata/posture.go | 16 ++ .../networkmap/nmdata/posture_test.go | 54 +++++ 22 files changed, 476 insertions(+), 333 deletions(-) create mode 100644 management/internals/controllers/network_map/controller/posture_twin_test.go create mode 100644 management/server/peer_posture_test.go delete mode 100644 management/server/posture/affects_posture_test.go create mode 100644 shared/management/networkmap/nmdata/posture_test.go diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index e74b17638..f21f878a4 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -566,15 +566,13 @@ 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 @@ -583,11 +581,9 @@ func peerPostureChecksFromData(nmData *networkmap.NetworkMapData, peerID string) 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) } } @@ -608,18 +604,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 +951,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 +1016,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 +1126,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 +1193,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 +1218,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 +1235,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) { diff --git a/management/internals/controllers/network_map/controller/posture_twin_test.go b/management/internals/controllers/network_map/controller/posture_twin_test.go new file mode 100644 index 000000000..d5c9035e0 --- /dev/null +++ b/management/internals/controllers/network_map/controller/posture_twin_test.go @@ -0,0 +1,38 @@ +package controller + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/networkmap" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" +) + +func TestPeerPostureChecksFromData_ReturnsTwinsUnchanged(t *testing.T) { + check := &nmdata.PostureChecks{ + ID: "pc1", + Checks: nmdata.ChecksDefinition{ + NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "0.30.0"}, + OSVersionCheck: &nmdata.OSVersionCheck{Linux: &nmdata.MinKernelVersionCheck{MinKernelVersion: "6.1"}}, + }, + } + nmData := &networkmap.NetworkMapData{ + Groups: map[string]*nmdata.Group{"g1": {ID: "g1", Peers: []string{"peer1"}}}, + Policies: []*nmdata.Policy{{ + ID: "policy1", + Enabled: true, + SourcePostureChecks: []string{"pc1"}, + Rules: []*nmdata.PolicyRule{{ID: "rule1", Enabled: true, Sources: []string{"g1"}}}, + }}, + PostureChecks: map[string]*nmdata.PostureChecks{"pc1": check}, + } + + got := peerPostureChecksFromData(nmData, "peer1") + require.Len(t, got, 1) + assert.Same(t, check, got[0]) + assert.Len(t, got[0].GetChecks(), 2) + + assert.Empty(t, peerPostureChecksFromData(nmData, "peer-outside-source-group")) +} diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index b535321d1..1e8c219b3 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -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) diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index 42051f172..8b104dfa0 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -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 diff --git a/management/internals/shared/grpc/components_envelope_response.go b/management/internals/shared/grpc/components_envelope_response.go index c059b2248..cdd2a7f37 100644 --- a/management/internals/shared/grpc/components_envelope_response.go +++ b/management/internals/shared/grpc/components_envelope_response.go @@ -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, diff --git a/management/internals/shared/grpc/conversion.go b/management/internals/shared/grpc/conversion.go index 5640127ca..96bd9f1f4 100644 --- a/management/internals/shared/grpc/conversion.go +++ b/management/internals/shared/grpc/conversion.go @@ -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). diff --git a/management/internals/shared/grpc/server.go b/management/internals/shared/grpc/server.go index 4435f6706..240243497 100644 --- a/management/internals/shared/grpc/server.go +++ b/management/internals/shared/grpc/server.go @@ -42,10 +42,10 @@ import ( "github.com/netbirdio/netbird/management/server/auth" nbContext "github.com/netbirdio/netbird/management/server/context" nbpeer "github.com/netbirdio/netbird/management/server/peer" - "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/proto" internalStatus "github.com/netbirdio/netbird/shared/management/status" ) @@ -902,7 +902,7 @@ func (s *Server) ExtendAuthSession(ctx context.Context, req *proto.EncryptedMess }, nil } -func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, network *types.Network, postureChecks []*posture.Checks, enableSSH bool) (*proto.LoginResponse, error) { +func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, network *types.Network, postureChecks []*nmdata.PostureChecks, enableSSH bool) (*proto.LoginResponse, error) { var relayToken *Token var err error if s.config.Relay != nil && len(s.config.Relay.Addresses) > 0 { @@ -990,7 +990,7 @@ func (s *Server) IsHealthy(ctx context.Context, req *proto.Empty) (*proto.Empty, } // sendInitialSync sends initial proto.SyncResponse to the peer requesting synchronization -func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer *nbpeer.Peer, networkMap *types.NetworkMap, postureChecks []*posture.Checks, srv proto.ManagementService_SyncServer, dnsFwdPort int64) error { +func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer *nbpeer.Peer, networkMap *types.NetworkMap, postureChecks []*nmdata.PostureChecks, srv proto.ManagementService_SyncServer, dnsFwdPort int64) error { var err error var turnToken *Token @@ -1301,7 +1301,7 @@ func (s *Server) Logout(ctx context.Context, req *proto.EncryptedMessage) (*prot } // toProtocolChecks converts posture checks to protocol checks. -func toProtocolChecks(ctx context.Context, postureChecks []*posture.Checks) []*proto.Checks { +func toProtocolChecks(ctx context.Context, postureChecks []*nmdata.PostureChecks) []*proto.Checks { protoChecks := make([]*proto.Checks, 0, len(postureChecks)) for _, postureCheck := range postureChecks { check := toProtocolCheck(postureCheck) @@ -1313,8 +1313,8 @@ func toProtocolChecks(ctx context.Context, postureChecks []*posture.Checks) []*p return protoChecks } -// toProtocolCheck converts a posture.Checks to a proto.Checks. -func toProtocolCheck(postureCheck *posture.Checks) *proto.Checks { +// toProtocolCheck converts posture checks to a proto.Checks. +func toProtocolCheck(postureCheck *nmdata.PostureChecks) *proto.Checks { protoCheck := &proto.Checks{} if check := postureCheck.Checks.ProcessCheck; check != nil { diff --git a/management/server/account.go b/management/server/account.go index 700dfa04d..4fe0e5338 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -52,6 +52,7 @@ import ( "github.com/netbirdio/netbird/management/server/util" "github.com/netbirdio/netbird/route" nbdomain "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/status" ) @@ -1920,7 +1921,7 @@ func domainIsUpToDate(domain string, domainCategory string, userAuth auth.UserAu // derived from syncTime (the moment the gRPC stream opened). Any // concurrent stream that started earlier loses the optimistic-lock race // in MarkPeerConnected and bails without writing. -func (am *DefaultAccountManager) SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (am *DefaultAccountManager) SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { peer, netMap, postureChecks, dnsfwdPort, err := am.SyncPeer(ctx, types.PeerSync{WireGuardPubKey: peerPubKey, Meta: meta, RealIP: realIP}, accountID) if err != nil { return nil, nil, nil, 0, fmt.Errorf("error syncing peer: %w", err) diff --git a/management/server/account/manager.go b/management/server/account/manager.go index f4b0408cf..154c9ab18 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -23,6 +23,7 @@ import ( "github.com/netbirdio/netbird/management/server/users" "github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" ) type ExternalCacheManager nbcache.UserDataCache @@ -70,7 +71,7 @@ type Manager interface { UpdatePeerIPv6(ctx context.Context, accountID, userID, peerID string, newIPv6 netip.Addr) error GetNetworkMap(ctx context.Context, peerID string) (*types.NetworkMap, error) GetPeerNetwork(ctx context.Context, peerID string) (*types.Network, error) - AddPeer(ctx context.Context, accountID, setupKey, userID string, p *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) + AddPeer(ctx context.Context, accountID, setupKey, userID string, p *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) CreatePAT(ctx context.Context, accountID string, initiatorUserID string, targetUserID string, tokenName string, expiresIn int) (*types.PersonalAccessTokenGenerated, error) DeletePAT(ctx context.Context, accountID string, initiatorUserID string, targetUserID string, tokenID string) error GetPAT(ctx context.Context, accountID string, initiatorUserID string, targetUserID string, tokenID string) (*types.PersonalAccessToken, error) @@ -109,9 +110,9 @@ type Manager interface { GetPeer(ctx context.Context, accountID, peerID, userID string) (*nbpeer.Peer, error) UpdateAccountSettings(ctx context.Context, accountID, userID string, newSettings *types.Settings) (*types.Settings, error) UpdateAccountOnboarding(ctx context.Context, accountID, userID string, newOnboarding *types.AccountOnboarding) (*types.AccountOnboarding, error) - LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) // used by peer gRPC API - ExtendPeerSession(ctx context.Context, peerPubKey, userID string) (time.Time, error) // used by peer gRPC API for ExtendAuthSession - SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) // used by peer gRPC API + LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) // used by peer gRPC API + ExtendPeerSession(ctx context.Context, peerPubKey, userID string) (time.Time, error) // used by peer gRPC API for ExtendAuthSession + SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) // used by peer gRPC API GetExternalCacheManager() ExternalCacheManager GetPostureChecks(ctx context.Context, accountID, postureChecksID, userID string) (*posture.Checks, error) SavePostureChecks(ctx context.Context, accountID, userID string, postureChecks *posture.Checks, create bool) (*posture.Checks, error) @@ -121,7 +122,7 @@ type Manager interface { UpdateIntegratedValidator(ctx context.Context, accountID, userID, validator string, groups []string) error GroupValidation(ctx context.Context, accountId string, groups []string) (bool, error) GetValidatedPeers(ctx context.Context, accountID string) (map[string]struct{}, map[string]string, error) - SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) + SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) OnPeerDisconnected(ctx context.Context, accountID string, peerPubKey string, streamStartTime time.Time) error SyncPeerMeta(ctx context.Context, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP) error FindExistingPostureCheck(accountID string, checks *posture.ChecksDefinition) (*posture.Checks, error) diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index 9ac10cba0..f31f63d0e 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -29,6 +29,7 @@ import ( route "github.com/netbirdio/netbird/route" auth "github.com/netbirdio/netbird/shared/auth" domain "github.com/netbirdio/netbird/shared/management/domain" + nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" gomock "go.uber.org/mock/gomock" ) @@ -86,12 +87,12 @@ func (mr *MockManagerMockRecorder) AccountExists(ctx, accountID any) *gomock.Cal } // AddPeer mocks base method. -func (m *MockManager) AddPeer(ctx context.Context, accountID, setupKey, userID string, p *peer.Peer, temporary bool) (*peer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (m *MockManager) AddPeer(ctx context.Context, accountID, setupKey, userID string, p *peer.Peer, temporary bool) (*peer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "AddPeer", ctx, accountID, setupKey, userID, p, temporary) ret0, _ := ret[0].(*peer.Peer) ret1, _ := ret[1].(*types.Network) - ret2, _ := ret[2].([]*posture.Checks) + ret2, _ := ret[2].([]*nmdata.PostureChecks) ret3, _ := ret[3].(bool) ret4, _ := ret[4].(error) return ret0, ret1, ret2, ret3, ret4 @@ -1323,12 +1324,12 @@ func (mr *MockManagerMockRecorder) ListUsers(ctx, accountID any) *gomock.Call { } // LoginPeer mocks base method. -func (m *MockManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*peer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (m *MockManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*peer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "LoginPeer", ctx, login) ret0, _ := ret[0].(*peer.Peer) ret1, _ := ret[1].(*types.Network) - ret2, _ := ret[2].([]*posture.Checks) + ret2, _ := ret[2].([]*nmdata.PostureChecks) ret3, _ := ret[3].(bool) ret4, _ := ret[4].(error) return ret0, ret1, ret2, ret3, ret4 @@ -1568,12 +1569,12 @@ func (mr *MockManagerMockRecorder) StoreEvent(ctx, initiatorID, targetID, accoun } // SyncAndMarkPeer mocks base method. -func (m *MockManager) SyncAndMarkPeer(ctx context.Context, accountID, peerPubKey string, meta peer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*peer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (m *MockManager) SyncAndMarkPeer(ctx context.Context, accountID, peerPubKey string, meta peer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*peer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "SyncAndMarkPeer", ctx, accountID, peerPubKey, meta, realIP, syncTime) ret0, _ := ret[0].(*peer.Peer) ret1, _ := ret[1].(*types.NetworkMap) - ret2, _ := ret[2].([]*posture.Checks) + ret2, _ := ret[2].([]*nmdata.PostureChecks) ret3, _ := ret[3].(int64) ret4, _ := ret[4].(error) return ret0, ret1, ret2, ret3, ret4 @@ -1586,12 +1587,12 @@ func (mr *MockManagerMockRecorder) SyncAndMarkPeer(ctx, accountID, peerPubKey, m } // SyncPeer mocks base method. -func (m *MockManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*peer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (m *MockManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*peer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "SyncPeer", ctx, sync, accountID) ret0, _ := ret[0].(*peer.Peer) ret1, _ := ret[1].(*types.NetworkMap) - ret2, _ := ret[2].([]*posture.Checks) + ret2, _ := ret[2].([]*nmdata.PostureChecks) ret3, _ := ret[3].(int64) ret4, _ := ret[4].(error) return ret0, ret1, ret2, ret3, ret4 diff --git a/management/server/account_test.go b/management/server/account_test.go index a5a484c1a..b462cc2a6 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -10,16 +10,17 @@ import ( "os" "reflect" "strconv" + "strings" "sync" "testing" "time" - "go.uber.org/mock/gomock" "github.com/prometheus/client_golang/prometheus/push" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel/metric/noop" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" @@ -37,6 +38,8 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" reverseproxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager" "github.com/netbirdio/netbird/management/internals/modules/zones" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" + networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory" "github.com/netbirdio/netbird/management/internals/server/config" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" nbAccount "github.com/netbirdio/netbird/management/server/account" @@ -3293,13 +3296,33 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU if err != nil { return nil, nil, err } - eventStore := &activity.InMemoryEventStore{} + return buildTestManager(t, store, nil) +} - metrics, err := telemetry.NewDefaultAppMetrics(context.Background()) - if err != nil { - return nil, nil, err +// createManagerWithNetworkMapStore builds a manager whose network map controller +// reads the twin (nmdata) store, the production path on sqlite and postgres. +func createManagerWithNetworkMapStore(t testing.TB) (*DefaultAccountManager, *update_channel.PeersUpdateManager) { + t.Helper() + + if engine := os.Getenv("NETBIRD_STORE_ENGINE"); engine != "" && !strings.EqualFold(engine, string(types.SqliteStoreEngine)) { + t.Skipf("network map store test needs the sqlite engine, got %s", engine) } + dataDir := t.TempDir() + store, err := createStoreAt(t, dataDir) + require.NoError(t, err) + + nmdataStore, err := networkmapdbfactory.NewNetworkMapDBStore(context.Background(), types.SqliteStoreEngine, dataDir, MockIntegratedValidator{}, newSettingsMockManager(t)) + require.NoError(t, err) + + manager, updateManager, err := buildTestManager(t, store, nmdataStore) + require.NoError(t, err) + return manager, updateManager +} + +func newSettingsMockManager(t testing.TB) *settings.MockManager { + t.Helper() + ctrl := gomock.NewController(t) t.Cleanup(ctrl.Finish) @@ -3312,6 +3335,23 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU UpdateExtraSettings(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). Return(false, nil). AnyTimes() + return settingsMockManager +} + +func buildTestManager(t testing.TB, store store.Store, nmdataStore *networkmapdb.NetworkMapDBStoreImpl) (*DefaultAccountManager, *update_channel.PeersUpdateManager, error) { + t.Helper() + + eventStore := &activity.InMemoryEventStore{} + + metrics, err := telemetry.NewDefaultAppMetrics(context.Background()) + if err != nil { + return nil, nil, err + } + + ctrl := gomock.NewController(t) + t.Cleanup(ctrl.Finish) + + settingsMockManager := newSettingsMockManager(t) permissionsManager := permissions.NewManager(store) peersManager := peers.NewManager(store, permissionsManager) @@ -3331,7 +3371,7 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU updateManager := update_channel.NewPeersUpdateManager(metrics) requestBuffer := NewAccountRequestBuffer(ctx, store) - networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil) + networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nmdataStore) manager, err := BuildManager(ctx, &config.Config{}, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) if err != nil { return nil, nil, err @@ -3349,7 +3389,11 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU func createStore(t testing.TB) (store.Store, error) { t.Helper() - dataDir := t.TempDir() + return createStoreAt(t, t.TempDir()) +} + +func createStoreAt(t testing.TB, dataDir string) (store.Store, error) { + t.Helper() store, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", dataDir) if err != nil { return nil, err diff --git a/management/server/affected_peers_router_paths_test.go b/management/server/affected_peers_router_paths_test.go index 5d83367fd..7fef1ab35 100644 --- a/management/server/affected_peers_router_paths_test.go +++ b/management/server/affected_peers_router_paths_test.go @@ -12,6 +12,7 @@ import ( resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" ) @@ -338,3 +339,44 @@ func TestAffectedPeers_PeerChange_RouterInOtherNetworkNotAffected(t *testing.T) assert.NotContains(t, affected, second.routerPeerID, "a router in an unrelated network must not be affected by a source-peer change for another resource") } + +// TestAffectedPeers_E2E_PostureFlip_RefreshesRoutingPeer drives the customer path +// on the twin store: the source peer's metadata flips a posture verdict on sync, +// and the routing peer serving the gated resource must be refreshed in both +// directions. Without the flip detection the deny direction takes the nmap +// shortcut (the denied peer's map holds no router) and the allow direction +// depends on which meta field moved, leaving the routers with a stale map. +func TestAffectedPeers_E2E_PostureFlip_RefreshesRoutingPeer(t *testing.T) { + manager, updateManager := createManagerWithNetworkMapStore(t) + s := buildRouterScenario(t, manager, updateManager, true) + ctx := context.Background() + + s.createPostureCheckGatedPolicy(t, ctx) + + source, err := s.manager.Store.GetPeerByID(ctx, store.LockingStrengthNone, s.accountID, s.sourcePeerID) + require.NoError(t, err) + + syncWithVersion := func(version string) { + meta := source.Meta + meta.WtVersion = version + _, _, _, _, err := s.manager.SyncPeer(ctx, types.PeerSync{WireGuardPubKey: source.Key, Meta: meta}, s.accountID) + require.NoError(t, err) + } + syncWithVersion("0.31.0") + + routerCh := s.updateManager.CreateChannel(ctx, s.routerPeerID) + unrelatedCh := s.updateManager.CreateChannel(ctx, s.unrelatedPeerID) + t.Cleanup(func() { + s.updateManager.CloseChannel(ctx, s.routerPeerID) + s.updateManager.CloseChannel(ctx, s.unrelatedPeerID) + }) + settleAffectedUpdates(routerCh, unrelatedCh) + + syncWithVersion("0.29.0") + peerShouldReceiveUpdate(t, routerCh) + peerShouldNotReceiveUpdate(t, unrelatedCh) + + syncWithVersion("0.31.0") + peerShouldReceiveUpdate(t, routerCh) + peerShouldNotReceiveUpdate(t, unrelatedCh) +} diff --git a/management/server/affected_peers_router_test.go b/management/server/affected_peers_router_test.go index cc9df0a6a..9ecfaed69 100644 --- a/management/server/affected_peers_router_test.go +++ b/management/server/affected_peers_router_test.go @@ -60,6 +60,12 @@ func setupRouterScenario(t *testing.T, directRouterPeer bool) *routerScenario { manager, updateManager, err := createManager(t) require.NoError(t, err) + return buildRouterScenario(t, manager, updateManager, directRouterPeer) +} + +func buildRouterScenario(t *testing.T, manager *DefaultAccountManager, updateManager *update_channel.PeersUpdateManager, directRouterPeer bool) *routerScenario { + t.Helper() + ctx := context.Background() account, err := createAccount(manager, "router_scenario", userID, "") diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 071e3771b..2f871c3e2 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -24,6 +24,7 @@ import ( "github.com/netbirdio/netbird/management/server/users" "github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" ) var _ account.Manager = (*MockAccountManager)(nil) @@ -41,11 +42,11 @@ type MockAccountManager struct { GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) MarkPeerConnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error MarkPeerDisconnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error - SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) + SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) DeletePeerFunc func(ctx context.Context, accountID, peerKey, userID string) error GetNetworkMapFunc func(ctx context.Context, peerKey string) (*types.NetworkMap, error) GetPeerNetworkFunc func(ctx context.Context, peerKey string) (*types.Network, error) - AddPeerFunc func(ctx context.Context, accountID string, setupKey string, userId string, peer *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) + AddPeerFunc func(ctx context.Context, accountID string, setupKey string, userId string, peer *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) GetGroupFunc func(ctx context.Context, accountID, groupID, userID string) (*types.Group, error) GetAllGroupsFunc func(ctx context.Context, accountID, userID string) ([]*types.Group, error) GetGroupByNameFunc func(ctx context.Context, groupName, accountID, userID string) (*types.Group, error) @@ -98,9 +99,9 @@ type MockAccountManager struct { SaveDNSSettingsFunc func(ctx context.Context, accountID, userID string, dnsSettingsToSave *types.DNSSettings) error GetPeerFunc func(ctx context.Context, accountID, peerID, userID string) (*nbpeer.Peer, error) UpdateAccountSettingsFunc func(ctx context.Context, accountID, userID string, newSettings *types.Settings) (*types.Settings, error) - LoginPeerFunc func(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) + LoginPeerFunc func(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) ExtendPeerSessionFunc func(ctx context.Context, peerPubKey, userID string) (time.Time, error) - SyncPeerFunc func(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) + SyncPeerFunc func(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) InviteUserFunc func(ctx context.Context, accountID string, initiatorUserID string, targetUserEmail string) error ApproveUserFunc func(ctx context.Context, accountID, initiatorUserID, targetUserID string) (*types.UserInfo, error) RejectUserFunc func(ctx context.Context, accountID, initiatorUserID, targetUserID string) error @@ -230,7 +231,7 @@ func (am *MockAccountManager) DeleteSetupKey(ctx context.Context, accountID, use return status.Errorf(codes.Unimplemented, "method DeleteSetupKey is not implemented") } -func (am *MockAccountManager) SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (am *MockAccountManager) SyncAndMarkPeer(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { if am.SyncAndMarkPeerFunc != nil { return am.SyncAndMarkPeerFunc(ctx, accountID, peerPubKey, meta, realIP, syncTime) } @@ -424,7 +425,7 @@ func (am *MockAccountManager) AddPeer( userId string, peer *nbpeer.Peer, temporary bool, -) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) { +) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { if am.AddPeerFunc != nil { return am.AddPeerFunc(ctx, accountID, setupKey, userId, peer, temporary) } @@ -862,7 +863,7 @@ func (am *MockAccountManager) UpdateAccountSettings(ctx context.Context, account } // LoginPeer mocks LoginPeer of the AccountManager interface -func (am *MockAccountManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (am *MockAccountManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { if am.LoginPeerFunc != nil { return am.LoginPeerFunc(ctx, login) } @@ -878,7 +879,7 @@ func (am *MockAccountManager) ExtendPeerSession(ctx context.Context, peerPubKey, } // SyncPeer mocks SyncPeer of the AccountManager interface -func (am *MockAccountManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (am *MockAccountManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { if am.SyncPeerFunc != nil { return am.SyncPeerFunc(ctx, sync, accountID) } diff --git a/management/server/peer.go b/management/server/peer.go index 579ff2708..87ca57c2b 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -23,7 +23,6 @@ import ( "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" - "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" @@ -741,7 +740,7 @@ func (am *DefaultAccountManager) handleSetupKeyAddedPeer(ctx context.Context, en // to it. We also add the User ID to the peer metadata to identify registrant. If no userID provided, then fail with status.PermissionDenied // Each new Peer will be assigned a new next net.IP from the Account.Network and Account.Network.LastIP will be updated (IP's are not reused). // The peer property is just a placeholder for the Peer properties to pass further -func (am *DefaultAccountManager) AddPeer(ctx context.Context, accountID, setupKey, userID string, peer *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (am *DefaultAccountManager) AddPeer(ctx context.Context, accountID, setupKey, userID string, peer *nbpeer.Peer, temporary bool) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { if setupKey == "" && userID == "" && !peer.ProxyMeta.Embedded { // no auth method provided => reject access return nil, nil, nil, false, status.ErrNoAuthMethodProvided @@ -1001,7 +1000,7 @@ func getPeerIPDNSLabel(ip netip.Addr, peerHostName string) (string, error) { } // SyncPeer checks whether peer is eligible for receiving NetworkMap (authenticated) and returns its NetworkMap if eligible -func (am *DefaultAccountManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { +func (am *DefaultAccountManager) SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) { var peer *nbpeer.Peer var ipv6CapabilityChanged bool var metaDiff nbpeer.MetaDiff @@ -1065,7 +1064,7 @@ func (am *DefaultAccountManager) SyncPeer(ctx context.Context, sync types.PeerSy return nil, nil, nil, 0, err } - metaDiffAffectsPosture := posture.AffectsPosture(ctx, &metaDiff, resPostureChecks) + metaDiffAffectsPosture := metaDiffAffectsPosture(&metaDiff, resPostureChecks) if requiresPeerUpdate(ctx, isStatusChanged, sync.UpdateAccountPeers, ipv6CapabilityChanged, metaDiffAffectsPosture, metaDiff.VersionChanged(), metaDiff.HostnameChanged()) { changedPeerIDs := []string{peer.ID} affectedPeerIDs := am.syncPeerAffectedPeers(ctx, accountID, peer.ID, nmap, peerNotValid, metaDiffAffectsPosture) @@ -1077,6 +1076,14 @@ func (am *DefaultAccountManager) SyncPeer(ctx context.Context, sync types.PeerSy return peer, nmap, resPostureChecks, dnsFwdPort, nil } +// metaDiffAffectsPosture reports whether the meta change flips the verdict of any of +// the peer's posture checks, replaying them against the old and new state. +func metaDiffAffectsPosture(diff *nbpeer.MetaDiff, checks []*nmdata.PostureChecks) bool { + oldPeer := types.TwinPeer(&nbpeer.Peer{Meta: diff.OldMeta, Location: diff.OldLocation}) + newPeer := types.TwinPeer(&nbpeer.Peer{Meta: diff.NewMeta, Location: diff.NewLocation}) + return nmdata.PostureVerdictChanged(checks, oldPeer, newPeer) +} + func requiresPeerUpdate(ctx context.Context, isStatusChanged, updateAccountPeers, ipv6CapabilityChanged, metaDiffAffectsPosture, versionChanged, hostname bool) bool { var reason string switch { @@ -1128,7 +1135,7 @@ func (am *DefaultAccountManager) markConnectedAffectedPeers(ctx context.Context, return affectedPeerIDsFromNetworkMap(nmap, peerID) } -func (am *DefaultAccountManager) handlePeerLoginNotFound(ctx context.Context, login types.PeerLogin, err error) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (am *DefaultAccountManager) handlePeerLoginNotFound(ctx context.Context, login types.PeerLogin, err error) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { if errStatus, ok := status.FromError(err); ok && errStatus.Type() == status.NotFound { // we couldn't find this peer by its public key which can mean that peer hasn't been registered yet. // Try registering it. @@ -1149,7 +1156,7 @@ func (am *DefaultAccountManager) handlePeerLoginNotFound(ctx context.Context, lo // LoginPeer logs in or registers a peer. // If peer doesn't exist the function checks whether a setup key or a user is present and registers a new peer if so. -func (am *DefaultAccountManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*posture.Checks, bool, error) { +func (am *DefaultAccountManager) LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.Network, []*nmdata.PostureChecks, bool, error) { accountID, err := am.Store.GetAccountIDByPeerPubKey(ctx, login.WireGuardPubKey) if err != nil { return am.handlePeerLoginNotFound(ctx, login, err) @@ -1322,7 +1329,7 @@ func (am *DefaultAccountManager) ExtendPeerSession(ctx context.Context, peerPubK // getPeerLoginInfo computes the login/register response data (network, posture // checks, SSH) from the store without building the peer's full network map. -func getPeerLoginInfo(ctx context.Context, transaction store.Store, accountID string, peer *nbpeer.Peer, isValid bool) (*types.Network, []*posture.Checks, bool, error) { +func getPeerLoginInfo(ctx context.Context, transaction store.Store, accountID string, peer *nbpeer.Peer, isValid bool) (*types.Network, []*nmdata.PostureChecks, bool, error) { network, err := transaction.GetAccountNetwork(ctx, store.LockingStrengthNone, accountID) if err != nil { return nil, nil, false, fmt.Errorf("get account network: %w", err) @@ -1364,7 +1371,7 @@ func isPeerSSHEnabled(ctx context.Context, peer *nbpeer.Peer, policies []*types. } // getPeerPostureChecks returns the posture checks for the peer. -func getPeerPostureChecks(ctx context.Context, transaction store.Store, accountID string, peerGroupIDs []string, policies []*types.Policy) ([]*posture.Checks, error) { +func getPeerPostureChecks(ctx context.Context, transaction store.Store, accountID string, peerGroupIDs []string, policies []*types.Policy) ([]*nmdata.PostureChecks, error) { if len(policies) == 0 { return nil, nil } @@ -1385,7 +1392,7 @@ func getPeerPostureChecks(ctx context.Context, transaction store.Store, accountI return nil, err } - return maps.Values(peerPostureChecks), nil + return types.TwinPostureChecksList(maps.Values(peerPostureChecks)), nil } // processPeerPostureChecks checks if the peer is in the source group of the policy and returns the posture checks. diff --git a/management/server/peer_posture_test.go b/management/server/peer_posture_test.go new file mode 100644 index 000000000..6b298f5d1 --- /dev/null +++ b/management/server/peer_posture_test.go @@ -0,0 +1,183 @@ +package server + +import ( + "net" + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" +) + +func diffFrom(oldMeta, newMeta nbpeer.PeerSystemMeta, oldLoc, newLoc nbpeer.Location) *nbpeer.MetaDiff { + return &nbpeer.MetaDiff{ + OldMeta: oldMeta, + NewMeta: newMeta, + OldLocation: oldLoc, + NewLocation: newLoc, + } +} + +func postureBundle(def nmdata.ChecksDefinition) []*nmdata.PostureChecks { + return []*nmdata.PostureChecks{{Checks: def}} +} + +func TestMetaDiffAffectsPosture_NBVersion(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "1.2.0"}}) + + tests := []struct { + name string + oldVer, newVer string + want bool + }{ + {"both above min, no flip", "1.3.0", "1.4.0", false}, + {"both below min, no flip", "1.0.0", "1.1.0", false}, + {"crosses up below->above", "1.1.0", "1.3.0", true}, + {"crosses down above->below", "1.3.0", "1.1.0", true}, + {"unparsable old only -> flip", "garbage", "1.3.0", true}, + {"unparsable both -> no flip", "garbage", "junk", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + diff := diffFrom( + nbpeer.PeerSystemMeta{WtVersion: tt.oldVer}, + nbpeer.PeerSystemMeta{WtVersion: tt.newVer}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.Equal(t, tt.want, metaDiffAffectsPosture(diff, c)) + }) + } +} + +func TestMetaDiffAffectsPosture_OSVersion_KernelBumpWithinMin(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{OSVersionCheck: &nmdata.OSVersionCheck{ + Linux: &nmdata.MinKernelVersionCheck{MinKernelVersion: "5.0.0"}, + }}) + + withinMin := diffFrom( + nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.10.0-arch1"}, + nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.15.0-arch2"}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.False(t, metaDiffAffectsPosture(withinMin, c)) + + crossesDown := diffFrom( + nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.10.0-arch1"}, + nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "4.19.0-arch1"}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.True(t, metaDiffAffectsPosture(crossesDown, c)) +} + +func TestMetaDiffAffectsPosture_OSVersion_GoOSSwitchFlipsVerdict(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{OSVersionCheck: &nmdata.OSVersionCheck{ + Linux: &nmdata.MinKernelVersionCheck{MinKernelVersion: "6.0.0"}, + }}) + + diff := diffFrom( + nbpeer.PeerSystemMeta{GoOS: "freebsd"}, + nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "4.19.0"}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.True(t, metaDiffAffectsPosture(diff, c)) +} + +func TestMetaDiffAffectsPosture_Process_GoOSSwitchFlipsVerdict(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{ProcessCheck: &nmdata.ProcessCheck{ + Processes: []nmdata.Process{{LinuxPath: "/usr/bin/foo"}}, + }}) + + files := []nbpeer.File{{Path: "/usr/bin/foo", ProcessIsRunning: true}} + diff := diffFrom( + nbpeer.PeerSystemMeta{GoOS: "linux", Files: files}, + nbpeer.PeerSystemMeta{GoOS: "windows", Files: files}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.True(t, metaDiffAffectsPosture(diff, c)) +} + +func TestMetaDiffAffectsPosture_Process_UnrelatedFileChange(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{ProcessCheck: &nmdata.ProcessCheck{ + Processes: []nmdata.Process{{LinuxPath: "/usr/bin/foo"}}, + }}) + + diff := diffFrom( + nbpeer.PeerSystemMeta{GoOS: "linux", Files: []nbpeer.File{ + {Path: "/usr/bin/foo", ProcessIsRunning: true}, + }}, + nbpeer.PeerSystemMeta{GoOS: "linux", Files: []nbpeer.File{ + {Path: "/usr/bin/foo", ProcessIsRunning: true}, + {Path: "/usr/bin/bar", ProcessIsRunning: true}, + }}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.False(t, metaDiffAffectsPosture(diff, c)) +} + +func TestMetaDiffAffectsPosture_GeoLocation(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{GeoLocationCheck: &nmdata.GeoLocationCheck{ + Action: posture.CheckActionAllow, + Locations: []nmdata.GeoLocation{{CountryCode: "DE"}}, + }}) + + stayAllowed := diffFrom( + nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, + nbpeer.Location{CountryCode: "DE", CityName: "Berlin"}, + nbpeer.Location{CountryCode: "DE", CityName: "Munich"}, + ) + assert.False(t, metaDiffAffectsPosture(stayAllowed, c)) + + moveOut := diffFrom( + nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, + nbpeer.Location{CountryCode: "DE"}, + nbpeer.Location{CountryCode: "FR"}, + ) + assert.True(t, metaDiffAffectsPosture(moveOut, c)) +} + +func TestMetaDiffAffectsPosture_PeerNetworkRange_ConnectionIP(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{PeerNetworkRangeCheck: &nmdata.PeerNetworkRangeCheck{ + Action: posture.CheckActionAllow, + Ranges: []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}, + }}) + + movesOutOfRange := diffFrom( + nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, + nbpeer.Location{ConnectionIP: net.ParseIP("10.1.2.3")}, + nbpeer.Location{ConnectionIP: net.ParseIP("8.8.8.8")}, + ) + assert.True(t, metaDiffAffectsPosture(movesOutOfRange, c)) + + staysInRange := diffFrom( + nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, + nbpeer.Location{ConnectionIP: net.ParseIP("10.1.2.3")}, + nbpeer.Location{ConnectionIP: net.ParseIP("10.9.9.9")}, + ) + assert.False(t, metaDiffAffectsPosture(staysInRange, c)) +} + +func TestMetaDiffAffectsPosture_IrrelevantFieldChange(t *testing.T) { + c := postureBundle(nmdata.ChecksDefinition{ + NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "1.0.0"}, + GeoLocationCheck: &nmdata.GeoLocationCheck{Action: posture.CheckActionAllow, Locations: []nmdata.GeoLocation{{CountryCode: "DE"}}}, + }) + + diff := diffFrom( + nbpeer.PeerSystemMeta{Hostname: "old", WtVersion: "1.5.0"}, + nbpeer.PeerSystemMeta{Hostname: "new", WtVersion: "1.5.0"}, + nbpeer.Location{CountryCode: "DE"}, nbpeer.Location{CountryCode: "DE"}, + ) + assert.False(t, metaDiffAffectsPosture(diff, c)) +} + +func TestMetaDiffAffectsPosture_NoChecks(t *testing.T) { + diff := diffFrom( + nbpeer.PeerSystemMeta{WtVersion: "1.0.0"}, + nbpeer.PeerSystemMeta{WtVersion: "2.0.0"}, + nbpeer.Location{}, nbpeer.Location{}, + ) + assert.False(t, metaDiffAffectsPosture(diff, nil)) +} diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 9a662bdbf..22f2b9b6f 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -16,11 +16,11 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/rs/xid" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" "golang.org/x/exp/maps" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -1170,11 +1170,11 @@ func TestToSyncResponse(t *testing.T) { }, } dnsName := "example.com" - checks := []*posture.Checks{ + checks := []*nmdata.PostureChecks{ { - Checks: posture.ChecksDefinition{ - ProcessCheck: &posture.ProcessCheck{ - Processes: []posture.Process{{LinuxPath: "/usr/bin/netbird"}}, + Checks: nmdata.ChecksDefinition{ + ProcessCheck: &nmdata.ProcessCheck{ + Processes: []nmdata.Process{{LinuxPath: "/usr/bin/netbird"}}, }, }, }, diff --git a/management/server/posture/affects_posture_test.go b/management/server/posture/affects_posture_test.go deleted file mode 100644 index 6aa54d892..000000000 --- a/management/server/posture/affects_posture_test.go +++ /dev/null @@ -1,202 +0,0 @@ -package posture - -import ( - "context" - "net" - "net/netip" - "testing" - - "github.com/stretchr/testify/assert" - - nbpeer "github.com/netbirdio/netbird/management/server/peer" -) - -// diffFrom builds a MetaDiff from the old/new snapshots AffectsPosture replays against. -func diffFrom(oldMeta, newMeta nbpeer.PeerSystemMeta, oldLoc, newLoc nbpeer.Location) *nbpeer.MetaDiff { - return &nbpeer.MetaDiff{ - OldMeta: oldMeta, - NewMeta: newMeta, - OldLocation: oldLoc, - NewLocation: newLoc, - } -} - -func checks(def ChecksDefinition) []*Checks { - return []*Checks{{Checks: def}} -} - -func TestAffectsPosture_NilDiff(t *testing.T) { - assert.False(t, AffectsPosture(context.Background(), nil, checks(ChecksDefinition{ - NBVersionCheck: &NBVersionCheck{MinVersion: "1.0.0"}, - }))) -} - -func TestAffectsPosture_NBVersion(t *testing.T) { - c := checks(ChecksDefinition{NBVersionCheck: &NBVersionCheck{MinVersion: "1.2.0"}}) - - tests := []struct { - name string - oldVer, newVer string - want bool - }{ - {"both above min, no flip", "1.3.0", "1.4.0", false}, - {"both below min, no flip", "1.0.0", "1.1.0", false}, - {"crosses up below->above", "1.1.0", "1.3.0", true}, - {"crosses down above->below", "1.3.0", "1.1.0", true}, - {"unparsable old only -> flip", "garbage", "1.3.0", true}, - {"unparsable both -> no flip", "garbage", "junk", false}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - diff := diffFrom( - nbpeer.PeerSystemMeta{WtVersion: tt.oldVer}, - nbpeer.PeerSystemMeta{WtVersion: tt.newVer}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.Equal(t, tt.want, AffectsPosture(context.Background(), diff, c)) - }) - } -} - -func TestAffectsPosture_OSVersion_KernelBumpWithinMin(t *testing.T) { - c := checks(ChecksDefinition{OSVersionCheck: &OSVersionCheck{ - Linux: &MinKernelVersionCheck{MinKernelVersion: "5.0.0"}, - }}) - - // Kernel moves but stays above the minimum: verdict stays pass -> not affected. - withinMin := diffFrom( - nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.10.0-arch1"}, - nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.15.0-arch2"}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.False(t, AffectsPosture(context.Background(), withinMin, c)) - - // Kernel drops below the minimum: verdict flips pass -> fail -> affected. - crossesDown := diffFrom( - nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "5.10.0-arch1"}, - nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "4.19.0-arch1"}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.True(t, AffectsPosture(context.Background(), crossesDown, c)) -} - -func TestAffectsPosture_OSVersion_GoOSSwitchFlipsVerdict(t *testing.T) { - // Only Linux is constrained. An OS outside the switch (freebsd) passes; switching to a - // failing linux kernel flips the verdict pass -> fail. - c := checks(ChecksDefinition{OSVersionCheck: &OSVersionCheck{ - Linux: &MinKernelVersionCheck{MinKernelVersion: "6.0.0"}, - }}) - - diff := diffFrom( - nbpeer.PeerSystemMeta{GoOS: "freebsd"}, - nbpeer.PeerSystemMeta{GoOS: "linux", KernelVersion: "4.19.0"}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.True(t, AffectsPosture(context.Background(), diff, c)) -} - -func TestAffectsPosture_Process_GoOSSwitchFlipsVerdict(t *testing.T) { - // Process runs at a linux path. Switching GoOS to windows (no WindowsPath configured) - // flips the verdict. - c := checks(ChecksDefinition{ProcessCheck: &ProcessCheck{ - Processes: []Process{{LinuxPath: "/usr/bin/foo"}}, - }}) - - files := []nbpeer.File{{Path: "/usr/bin/foo", ProcessIsRunning: true}} - diff := diffFrom( - nbpeer.PeerSystemMeta{GoOS: "linux", Files: files}, - nbpeer.PeerSystemMeta{GoOS: "windows", Files: files}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.True(t, AffectsPosture(context.Background(), diff, c)) -} - -func TestAffectsPosture_Process_UnrelatedFileChange(t *testing.T) { - // A tracked process stays running while an unrelated file is added: the verdict does - // not move, so posture is not affected. - c := checks(ChecksDefinition{ProcessCheck: &ProcessCheck{ - Processes: []Process{{LinuxPath: "/usr/bin/foo"}}, - }}) - - diff := diffFrom( - nbpeer.PeerSystemMeta{GoOS: "linux", Files: []nbpeer.File{ - {Path: "/usr/bin/foo", ProcessIsRunning: true}, - }}, - nbpeer.PeerSystemMeta{GoOS: "linux", Files: []nbpeer.File{ - {Path: "/usr/bin/foo", ProcessIsRunning: true}, - {Path: "/usr/bin/bar", ProcessIsRunning: true}, - }}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.False(t, AffectsPosture(context.Background(), diff, c)) -} - -func TestAffectsPosture_GeoLocation(t *testing.T) { - c := checks(ChecksDefinition{GeoLocationCheck: &GeoLocationCheck{ - Action: CheckActionAllow, - Locations: []Location{{CountryCode: "DE"}}, - }}) - - // Moving within allowed countries keeps the verdict; moving out flips it. - stayAllowed := diffFrom( - nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, - nbpeer.Location{CountryCode: "DE", CityName: "Berlin"}, - nbpeer.Location{CountryCode: "DE", CityName: "Munich"}, - ) - assert.False(t, AffectsPosture(context.Background(), stayAllowed, c)) - - moveOut := diffFrom( - nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, - nbpeer.Location{CountryCode: "DE"}, - nbpeer.Location{CountryCode: "FR"}, - ) - assert.True(t, AffectsPosture(context.Background(), moveOut, c)) -} - -func TestAffectsPosture_PeerNetworkRange_ConnectionIP(t *testing.T) { - // The check reads the connection IP. Moving out of the allowed range flips the verdict; - // moving within it does not. - _, allowed, _ := net.ParseCIDR("10.0.0.0/8") - c := checks(ChecksDefinition{PeerNetworkRangeCheck: &PeerNetworkRangeCheck{ - Action: CheckActionAllow, - Ranges: []netip.Prefix{netip.MustParsePrefix(allowed.String())}, - }}) - - movesOutOfRange := diffFrom( - nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, - nbpeer.Location{ConnectionIP: net.ParseIP("10.1.2.3")}, - nbpeer.Location{ConnectionIP: net.ParseIP("8.8.8.8")}, - ) - assert.True(t, AffectsPosture(context.Background(), movesOutOfRange, c)) - - staysInRange := diffFrom( - nbpeer.PeerSystemMeta{}, nbpeer.PeerSystemMeta{}, - nbpeer.Location{ConnectionIP: net.ParseIP("10.1.2.3")}, - nbpeer.Location{ConnectionIP: net.ParseIP("10.9.9.9")}, - ) - assert.False(t, AffectsPosture(context.Background(), staysInRange, c)) -} - -func TestAffectsPosture_IrrelevantFieldChange(t *testing.T) { - // Hostname changes but no check reads it: not affected even with checks present. - c := checks(ChecksDefinition{ - NBVersionCheck: &NBVersionCheck{MinVersion: "1.0.0"}, - GeoLocationCheck: &GeoLocationCheck{Action: CheckActionAllow, Locations: []Location{{CountryCode: "DE"}}}, - }) - - diff := diffFrom( - nbpeer.PeerSystemMeta{Hostname: "old", WtVersion: "1.5.0"}, - nbpeer.PeerSystemMeta{Hostname: "new", WtVersion: "1.5.0"}, - nbpeer.Location{CountryCode: "DE"}, nbpeer.Location{CountryCode: "DE"}, - ) - assert.False(t, AffectsPosture(context.Background(), diff, c)) -} - -func TestAffectsPosture_NoChecks(t *testing.T) { - diff := diffFrom( - nbpeer.PeerSystemMeta{WtVersion: "1.0.0"}, - nbpeer.PeerSystemMeta{WtVersion: "2.0.0"}, - nbpeer.Location{}, nbpeer.Location{}, - ) - assert.False(t, AffectsPosture(context.Background(), diff, nil)) -} diff --git a/management/server/posture/checks.go b/management/server/posture/checks.go index 72b719252..c38136d1c 100644 --- a/management/server/posture/checks.go +++ b/management/server/posture/checks.go @@ -7,7 +7,6 @@ import ( "regexp" "github.com/hashicorp/go-version" - log "github.com/sirupsen/logrus" nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/shared/management/http/api" @@ -55,46 +54,6 @@ type Checks struct { Checks ChecksDefinition `gorm:"serializer:json"` } -// AffectsPosture reports whether the change in diff flips the verdict of any check. It -// replays each check against the peer's old and new state and compares verdicts, so a -// change that moves a field but stays the right side of a threshold (e.g. a kernel bump -// still above the minimum) does not force a re-evaluation. See verdictChanged for how an -// evaluation error counts. -func AffectsPosture(ctx context.Context, diff *nbpeer.MetaDiff, checks []*Checks) bool { - if diff == nil { - return false - } - - oldPeer := nbpeer.Peer{Meta: diff.OldMeta, Location: diff.OldLocation} - newPeer := nbpeer.Peer{Meta: diff.NewMeta, Location: diff.NewLocation} - - for _, c := range checks { - for _, check := range c.GetChecks() { - if verdictChanged(ctx, check, oldPeer, newPeer) { - return true - } - } - } - return false -} - -// verdictChanged replays check against old and new state and reports whether the verdict -// differs. Like callers, it treats an evaluation error as deny: two errors are the same -// verdict (no change), an error on one side only is a flip. -func verdictChanged(ctx context.Context, check Check, oldPeer, newPeer nbpeer.Peer) bool { - oldPass, oldErr := check.Check(ctx, oldPeer) - newPass, newErr := check.Check(ctx, newPeer) - - oldVerdict := oldPass && (oldErr == nil) - newVerdict := newPass && (newErr == nil) - changed := oldVerdict != newVerdict - - log.WithContext(ctx).Tracef("posture check %s replay: verdict %t -> %t (changed=%t), errs: %v -> %v", - check.Name(), oldVerdict, newVerdict, changed, oldErr, newErr) - - return changed -} - // ChecksDefinition contains definition of actual check type ChecksDefinition struct { NBVersionCheck *NBVersionCheck `json:",omitempty"` diff --git a/management/server/types/account_networkmapdata.go b/management/server/types/account_networkmapdata.go index 8f2e03a10..d554bfe80 100644 --- a/management/server/types/account_networkmapdata.go +++ b/management/server/types/account_networkmapdata.go @@ -93,7 +93,7 @@ func (a *Account) toNetworkMapData( } for _, pc := range a.PostureChecks { if pc != nil { - nmd.PostureChecks[pc.ID] = twinPostureChecks(pc) + nmd.PostureChecks[pc.ID] = TwinPostureChecks(pc) nmd.PostureCheckXIDToPublicID[pc.ID] = pc.PublicID } } @@ -391,7 +391,17 @@ func TwinNetwork(n *Network) *nmdata.Network { } } -func twinPostureChecks(pc *posture.Checks) *nmdata.PostureChecks { +// TwinPostureChecksList converts posture checks to their slim nmdata twins. +func TwinPostureChecksList(checks []*posture.Checks) []*nmdata.PostureChecks { + out := make([]*nmdata.PostureChecks, 0, len(checks)) + for _, pc := range checks { + out = append(out, TwinPostureChecks(pc)) + } + return out +} + +// TwinPostureChecks converts posture checks to their slim nmdata twin. +func TwinPostureChecks(pc *posture.Checks) *nmdata.PostureChecks { if pc == nil { return nil } diff --git a/shared/management/networkmap/nmdata/posture.go b/shared/management/networkmap/nmdata/posture.go index dc1753791..6a6b028c7 100644 --- a/shared/management/networkmap/nmdata/posture.go +++ b/shared/management/networkmap/nmdata/posture.go @@ -45,6 +45,22 @@ func PassesChecks(checks []Check, peer *Peer) bool { return true } +// PostureVerdictChanged reports whether any check in the bundles gives a different +// verdict for newPeer than for oldPeer. Checks are replayed one by one, so a change +// that moves a field but stays on the same side of a threshold does not count. An +// evaluation error is a deny, like in PassesChecks. +func PostureVerdictChanged(checks []*PostureChecks, oldPeer, newPeer *Peer) bool { + for _, pc := range checks { + for _, c := range pc.GetChecks() { + single := []Check{c} + if PassesChecks(single, oldPeer) != PassesChecks(single, newPeer) { + return true + } + } + } + return false +} + // GetChecks returns the initialized checks in the same order as posture.Checks.GetChecks. func (pc *PostureChecks) GetChecks() []Check { var checks []Check diff --git a/shared/management/networkmap/nmdata/posture_test.go b/shared/management/networkmap/nmdata/posture_test.go new file mode 100644 index 000000000..13e5f268e --- /dev/null +++ b/shared/management/networkmap/nmdata/posture_test.go @@ -0,0 +1,54 @@ +package nmdata + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func bundle(def ChecksDefinition) []*PostureChecks { + return []*PostureChecks{{Checks: def}} +} + +func TestPostureVerdictChanged_ErrorCountsAsDeny(t *testing.T) { + c := bundle(ChecksDefinition{NBVersionCheck: &NBVersionCheck{MinVersion: "1.2.0"}}) + + tests := []struct { + name string + oldVer, newVer string + want bool + }{ + {"both above min, no flip", "1.3.0", "1.4.0", false}, + {"crosses up below->above", "1.1.0", "1.3.0", true}, + {"unparsable old only -> flip", "garbage", "1.3.0", true}, + {"unparsable both -> no flip", "garbage", "junk", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + oldPeer := &Peer{Meta: PeerSystemMeta{WtVersion: tt.oldVer}} + newPeer := &Peer{Meta: PeerSystemMeta{WtVersion: tt.newVer}} + assert.Equal(t, tt.want, PostureVerdictChanged(c, oldPeer, newPeer)) + }) + } +} + +func TestPostureVerdictChanged_ReplaysEachCheck(t *testing.T) { + // Old fails the version check, new fails the kernel check: the bundle denies on + // both sides, yet every single check flipped, so the posture must be re-evaluated. + c := bundle(ChecksDefinition{ + NBVersionCheck: &NBVersionCheck{MinVersion: "1.0.0"}, + OSVersionCheck: &OSVersionCheck{Linux: &MinKernelVersionCheck{MinKernelVersion: "5.0.0"}}, + }) + oldPeer := &Peer{Meta: PeerSystemMeta{WtVersion: "0.9.0", GoOS: "linux", KernelVersion: "6.0.0"}} + newPeer := &Peer{Meta: PeerSystemMeta{WtVersion: "1.1.0", GoOS: "linux", KernelVersion: "4.0.0"}} + + assert.False(t, c[0].Passes(oldPeer)) + assert.False(t, c[0].Passes(newPeer)) + assert.True(t, PostureVerdictChanged(c, oldPeer, newPeer)) +} + +func TestPostureVerdictChanged_NoChecks(t *testing.T) { + oldPeer := &Peer{Meta: PeerSystemMeta{WtVersion: "1.0.0"}} + newPeer := &Peer{Meta: PeerSystemMeta{WtVersion: "2.0.0"}} + assert.False(t, PostureVerdictChanged(nil, oldPeer, newPeer)) +}