From 33954ea15e616b72499cf82ac450b8c07d8591c0 Mon Sep 17 00:00:00 2001 From: Dmitri Dolguikh Date: Wed, 24 Jun 2026 13:07:53 +0200 Subject: [PATCH] fixing tests + adding tests Signed-off-by: Dmitri Dolguikh --- .../server/affected_peers_coverage_test.go | 10 +- management/server/affected_peers_test.go | 104 ++++++++++++------ management/server/affectedpeers/resolver.go | 84 +++++++++----- 3 files changed, 130 insertions(+), 68 deletions(-) diff --git a/management/server/affected_peers_coverage_test.go b/management/server/affected_peers_coverage_test.go index 56917905f..661e89b2e 100644 --- a/management/server/affected_peers_coverage_test.go +++ b/management/server/affected_peers_coverage_test.go @@ -32,7 +32,7 @@ func TestAffectedPeers_DependencyCoverageMatrix(t *testing.T) { _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) require.NoError(t, err) return affectedpeers.Change{ChangedGroupIDs: []string{s.sourceGroupID}}, - []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} // TODO (dmitri) routerPeer is missing }, }, { @@ -106,12 +106,8 @@ func TestAffectedPeers_DependencyCoverageMatrix(t *testing.T) { change, mustContain, mustExclude := r.build(t, s, ctx) affected := resolveAffected(t, s.manager.Store, s.accountID, change) - for _, id := range mustContain { - assert.Contains(t, affected, id, "expected peer to be affected") - } - for _, id := range mustExclude { - assert.NotContains(t, affected, id, "peer must not be affected") - } + assert.ElementsMatch(t, affected, mustContain, "expected peer to be affected") + assert.NotElementsMatch(t, affected, mustExclude, "peer must not be affected") }) } } diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index a19f9e174..ed2f558dc 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -96,31 +96,54 @@ func affectedGroupID(i int) string { return fmt.Sprintf("affected-grp-%d", i) func affectedGroupName(i int) string { return fmt.Sprintf("AffectedGroup%d", i) } func TestCollectGroupChange_PolicyLinked(t *testing.T) { - manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ { - Enabled: true, - Sources: []string{groupIDs[0]}, - Destinations: []string{groupIDs[1]}, - Bidirectional: true, - Action: types.PolicyTrafficActionAccept, + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + SourceResource: types.Resource{ID: peerIDs[0], Type: types.ResourceTypePeer}, + DestinationResource: types.Resource{ID: peerIDs[1], Type: types.ResourceTypePeer}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + SourceResource: types.Resource{ID: peerIDs[2], Type: types.ResourceTypeHost}, + DestinationResource: types.Resource{ID: peerIDs[3], Type: types.ResourceTypeHost}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + SourceResource: types.Resource{ID: "", Type: types.ResourceTypePeer}, + DestinationResource: types.Resource{ID: "", Type: types.ResourceTypePeer}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, }, }, }, true) require.NoError(t, err) - groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) - assert.ElementsMatch(t, groups, []string{groupIDs[1]}) + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) + assert.ElementsMatch(t, directPeers, []string{peerIDs[1]}) - groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) - assert.ElementsMatch(t, groups, []string{groupIDs[0]}) + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) + assert.ElementsMatch(t, directPeers, []string{peerIDs[0]}) - groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) assert.Empty(t, groups) + assert.Empty(t, directPeers) } func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { @@ -138,13 +161,37 @@ func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { Destinations: []string{groupIDs[1]}, Action: types.PolicyTrafficActionAccept, }, + { + Enabled: true, + Sources: []string{groupIDs[0]}, + SourceResource: types.Resource{ID: peerIDs[1], Type: types.ResourceTypeHost}, + DestinationResource: types.Resource{ID: peerIDs[2], Type: types.ResourceTypeHost}, + Destinations: []string{groupIDs[1]}, + Action: types.PolicyTrafficActionAccept, + }, + { + Enabled: true, + Sources: []string{groupIDs[0]}, + SourceResource: types.Resource{ID: "", Type: types.ResourceTypePeer}, + DestinationResource: types.Resource{ID: "", Type: types.ResourceTypePeer}, + Destinations: []string{groupIDs[1]}, + Action: types.PolicyTrafficActionAccept, + }, }, }, true) require.NoError(t, err) groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) - assert.ElementsMatch(t, groups, []string{groupIDs[1]}) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) assert.ElementsMatch(t, directPeers, []string{peerIDs[4]}) + + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) + assert.ElementsMatch(t, directPeers, []string{peerIDs[3]}) + + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) + assert.Empty(t, groups) + assert.Empty(t, directPeers) } func TestCollectGroupChange_PolicyWithNonPeerResource_NoDirectPeers(t *testing.T) { @@ -166,7 +213,7 @@ func TestCollectGroupChange_PolicyWithNonPeerResource_NoDirectPeers(t *testing.T require.NoError(t, err) groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) - assert.ElementsMatch(t, groups, []string{groupIDs[1]}) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) assert.Empty(t, directPeers, "non-peer resources should not produce direct peer IDs") } @@ -370,17 +417,11 @@ func TestCollectGroupChange_MultipleEntities(t *testing.T) { require.NoError(t, err) groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) - assert.NotContains(t, groups, groupIDs[0]) - assert.Contains(t, groups, groupIDs[1]) - assert.NotContains(t, groups, groupIDs[2]) - assert.NotContains(t, groups, groupIDs[3]) + assert.ElementsMatch(t, groups, []string{groupIDs[0], groupIDs[1]}) assert.Empty(t, directPeers) groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[3]}) - assert.Contains(t, groups, groupIDs[2]) - assert.Contains(t, groups, groupIDs[3]) - assert.NotContains(t, groups, groupIDs[0]) - assert.NotContains(t, groups, groupIDs[1]) + assert.ElementsMatch(t, groups, []string{groupIDs[2], groupIDs[3]}) assert.Empty(t, directPeers) } @@ -444,10 +485,10 @@ func TestResolveAffectedPeers_PolicyBetweenTwoGroups(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.ElementsMatch(t, []string{peerIDs[1]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) - assert.ElementsMatch(t, []string{peerIDs[0]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) assert.Empty(t, result) @@ -471,7 +512,7 @@ func TestResolveAffectedPeers_PolicyThreeGroups(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.ElementsMatch(t, []string{peerIDs[2]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result) } func TestResolveAffectedPeers_RoutePeerGroups(t *testing.T) { @@ -658,7 +699,7 @@ func TestResolveAffectedPeers_PeerInMultipleGroups(t *testing.T) { // peer0 is in group0 AND group1, so both policies apply result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2], peerIDs[3]}, result) } func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { @@ -694,7 +735,7 @@ func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0], peerIDs[2]}) - assert.ElementsMatch(t, []string{peerIDs[1], peerIDs[3]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2], peerIDs[1], peerIDs[3]}, result) } func TestResolveAffectedPeers_SharedGroupAcrossPolicyAndRoute(t *testing.T) { @@ -842,12 +883,12 @@ func TestAffectedPeers_IsolatedPolicies(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.ElementsMatch(t, []string{peerIDs[1]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) assert.NotContains(t, result, peerIDs[2]) assert.NotContains(t, result, peerIDs[3]) result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) - assert.ElementsMatch(t, []string{peerIDs[3]}, result) + assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) assert.NotContains(t, result, peerIDs[0]) assert.NotContains(t, result, peerIDs[1]) @@ -893,7 +934,7 @@ func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.ElementsMatch(t, []string{peerIDs[1]}, result) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) assert.NotContains(t, result, peerIDs[2]) assert.NotContains(t, result, peerIDs[3]) @@ -948,14 +989,13 @@ func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { }) result := manager.resolveAffectedPeersForPeerChanges(ctx, manager.Store, accountID, []string{peer1.ID}) - assert.ElementsMatch(t, []string{peer2.ID}, result) + assert.ElementsMatch(t, []string{peer1.ID, peer2.ID}, result) t.Run("group change updates all peers in policy groups", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldNotReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg1) peerShouldReceiveUpdate(t, updMsg2) - // TODO (dmitri) what's going on here? peerShouldReceiveUpdate(t, updMsg3) close(done) }() diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go index ad278039a..5e63524b5 100644 --- a/management/server/affectedpeers/resolver.go +++ b/management/server/affectedpeers/resolver.go @@ -449,14 +449,13 @@ func (r *resolver) collectFromPolicies() { for _, policy := range r.policies() { // changed peer IDs have been mapped to changedGroupSet on resolver creation (see seedChangedGroupsFromPeers) // there's no change to the groupSet if the same policies have been changed directly - groupIds := groupsFromPolicyDirectionally(policy, r.changedGroupSet) + peerIds, groupIds := getGroupsAndPeersFromPolicyViaGroups(policy, r.changedGroupSet) addAll(r.groupSet, groupIds) + addAll(r.peerSet, peerIds) - var peerIds []string - if len(r.changedPeerSet) > 0 { - peerIds = peersFromPolicyDirectionally(policy, r.changedPeerSet) - addAll(r.peerSet, peerIds) - } + peerIds, groupIds = getGroupsAndPeersFromPolicyViaPeers(policy, r.changedPeerSet) + addAll(r.groupSet, groupIds) + addAll(r.peerSet, peerIds) if len(groupIds) == 0 && len(peerIds) == 0 { continue @@ -742,40 +741,57 @@ func collectPolicySources(policy *types.Policy, groupSet, peerSet map[string]str } } -// returns group IDs of groups on the opposite side of the policy: -// i.e. if a group is present in the policy rule sources, use group IDs from the rule's destinations +// returns group and peer IDs on the opposite side of the policy: +// i.e. if a group is present in the policy rule sources, return destination group IDs and the destinationResource from the rule // and vice-versa -func groupsFromPolicyDirectionally(policy *types.Policy, groupSet map[string]struct{}) []string { - groupIds := make([]string, 0) - for _, rule := range policy.Rules { - // TODO (dmitri) can a group to be present on both sides of a policy? - if anyInSet(rule.Sources, groupSet) { - groupIds = append(groupIds, rule.Destinations...) - } else if anyInSet(rule.Destinations, groupSet) { - groupIds = append(groupIds, rule.Sources...) - } +func getGroupsAndPeersFromPolicyViaGroups(policy *types.Policy, groupSet map[string]struct{}) ([]string, []string) { + var groupIds, peerIds []string + if len(groupSet) == 0 { + return peerIds, groupIds } - return groupIds -} - -// returns peer IDs of peers on the opposite side of the policy: -// i.e. if a peer is present in the policy rule sourceResources, use destinationResources of the policy -// and vice-versa -func peersFromPolicyDirectionally(policy *types.Policy, changedSet map[string]struct{}) []string { - peerIds := make([]string, 0) for _, rule := range policy.Rules { - // TODO (dmitri) can a peer to be present on both sides of a policy? - if isDirectPeerInSet(rule.SourceResource, changedSet) { + if matchedIds, ok := allInSet(rule.Sources, groupSet); ok { + groupIds = append(groupIds, matchedIds...) + groupIds = append(groupIds, rule.Destinations...) if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { peerIds = append(peerIds, rule.DestinationResource.ID) } - } else if isDirectPeerInSet(rule.DestinationResource, changedSet) { + } + if matchedIds, ok := allInSet(rule.Destinations, groupSet); ok { + groupIds = append(groupIds, matchedIds...) + groupIds = append(groupIds, rule.Sources...) if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { peerIds = append(peerIds, rule.SourceResource.ID) } } } - return peerIds + return peerIds, groupIds +} + +// returns group and peer IDs on the opposite side of the policy: +// i.e. if a peer is present in the policy rule sourceResources, return destination group IDs and the destinationResource from the rule +// and vice-versa +func getGroupsAndPeersFromPolicyViaPeers(policy *types.Policy, changedSet map[string]struct{}) ([]string, []string) { + var groupIds, peerIds []string + if len(changedSet) == 0 { + return peerIds, groupIds + } + for _, rule := range policy.Rules { + if isDirectPeerInSet(rule.SourceResource, changedSet) { + groupIds = append(groupIds, rule.Destinations...) + peerIds = append(peerIds, rule.SourceResource.ID) + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + peerIds = append(peerIds, rule.DestinationResource.ID) + } + } else if isDirectPeerInSet(rule.DestinationResource, changedSet) { + groupIds = append(groupIds, rule.Sources...) + peerIds = append(peerIds, rule.DestinationResource.ID) + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + peerIds = append(peerIds, rule.SourceResource.ID) + } + } + } + return peerIds, groupIds } func policyReferencesPostureChecks(policy *types.Policy, ids map[string]struct{}) bool { @@ -821,6 +837,16 @@ func anyInSet(ids []string, set map[string]struct{}) bool { return false } +func allInSet(ids []string, set map[string]struct{}) ([]string, bool) { + var matchedIds []string + for _, id := range ids { + if _, ok := set[id]; ok { + matchedIds = append(matchedIds, id) + } + } + return matchedIds, len(matchedIds) > 0 +} + func isInSet(id string, set map[string]struct{}) bool { _, ok := set[id] return ok