diff --git a/management/server/certificate_challenge.go b/management/server/certificate_challenge.go index 22b0fd35d..71514753d 100644 --- a/management/server/certificate_challenge.go +++ b/management/server/certificate_challenge.go @@ -4,6 +4,8 @@ import ( "context" "encoding/binary" "hash/fnv" + "maps" + "slices" "sync" "time" @@ -147,13 +149,13 @@ func offsetWithin(key string, period time.Duration) time.Duration { return time.Duration(binary.BigEndian.Uint64(h.Sum(nil)) % uint64(period)) } -// refreshCertificateChallenges pushes the account's peers an update so each one is -// stamped with a nonce for the current window, and reports whether the account still -// has a certificate posture check to refresh for. +// refreshCertificateChallenges pushes an update to the peers that answer a certificate +// challenge, so each is stamped with a nonce for the current window. It reports whether +// the account still has a certificate posture check to refresh for. func (am *DefaultAccountManager) refreshCertificateChallenges(ctx context.Context, accountID string) bool { - wanted, err := am.accountNeedsCertificateChallenges(ctx, accountID) + peerIDs, wanted, err := am.certificateChallengeTargets(ctx, accountID) if err != nil { - log.WithContext(ctx).Debugf("cannot tell whether account %s still needs certificate challenges: %v", accountID, err) + log.WithContext(ctx).Debugf("cannot resolve the certificate challenge targets of account %s: %v", accountID, err) // Keep the account tracked: a store error now says nothing about its checks. return true } @@ -161,15 +163,96 @@ func (am *DefaultAccountManager) refreshCertificateChallenges(ctx context.Contex log.WithContext(ctx).Debugf("account %s has no certificate posture check left, stopping challenge refresh", accountID) return false } + if len(peerIDs) == 0 { + log.WithContext(ctx).Tracef("account %s has a certificate posture check but no peer answers it yet", accountID) + return true + } - log.WithContext(ctx).Debugf("refreshing certificate challenges for account %s", accountID) - am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{ - Resource: types.UpdateResourcePostureCheck, - Operation: types.UpdateOperationRefresh, - }) + log.WithContext(ctx).Debugf("refreshing certificate challenges for %d peers of account %s", len(peerIDs), accountID) + if err := am.networkMapController.UpdateAffectedPeers(ctx, accountID, peerIDs); err != nil { + log.WithContext(ctx).Warnf("failed refreshing certificate challenges for account %s: %v", accountID, err) + } return true } +// certificateChallengeTargets returns the peers that are sent a certificate challenge, +// and whether the account asks for one at all. Only those peers hold a nonce, so only +// they need the update; pushing to the whole account would wake every peer that has +// nothing to do with certificates. +// +// A peer is sent a challenge when it is a source of an enabled policy whose posture +// checks include a certificate check. This is the inverse of processPeerPostureChecks, +// which decides the same thing one peer at a time, and the two are held together by +// TestCertificateChallengeTargets_MatchesThePerPeerRule. +func (am *DefaultAccountManager) certificateChallengeTargets(ctx context.Context, accountID string) ([]string, bool, error) { + certCheckIDs, err := am.certificatePostureCheckIDs(ctx, accountID) + if err != nil { + return nil, false, err + } + if len(certCheckIDs) == 0 { + return nil, false, nil + } + + policies, err := am.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, true, err + } + groups, err := am.Store.GetAccountGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, true, err + } + groupPeers := make(map[string][]string, len(groups)) + for _, g := range groups { + groupPeers[g.ID] = g.Peers + } + + return certificateChallengeTargets(policies, groupPeers, certCheckIDs), true, nil +} + +// certificateChallengeTargets collects the source peers of every enabled policy whose +// posture checks include a certificate check. +func certificateChallengeTargets(policies []*types.Policy, groupPeers map[string][]string, certCheckIDs map[string]struct{}) []string { + targets := map[string]struct{}{} + for _, policy := range policies { + if !policy.Enabled || !slices.ContainsFunc(policy.SourcePostureChecks, func(id string) bool { + _, ok := certCheckIDs[id] + return ok + }) { + continue + } + for _, rule := range policy.Rules { + if !rule.Enabled { + continue + } + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + targets[rule.SourceResource.ID] = struct{}{} + } + for _, groupID := range rule.Sources { + for _, peerID := range groupPeers[groupID] { + targets[peerID] = struct{}{} + } + } + } + } + return slices.Collect(maps.Keys(targets)) +} + +// certificatePostureCheckIDs returns the IDs of the account's posture checks that ask a +// peer to prove a certificate. +func (am *DefaultAccountManager) certificatePostureCheckIDs(ctx context.Context, accountID string) (map[string]struct{}, error) { + checks, err := am.Store.GetAccountPostureChecks(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, err + } + ids := map[string]struct{}{} + for _, check := range checks { + if check.Checks.CertificateCheck != nil { + ids[check.ID] = struct{}{} + } + } + return ids, nil +} + // trackCertificateChallenges starts refreshing the account's certificate challenges if // it has a posture check that asks for one. func (am *DefaultAccountManager) trackCertificateChallenges(ctx context.Context, accountID string) { @@ -177,28 +260,13 @@ func (am *DefaultAccountManager) trackCertificateChallenges(ctx context.Context, return } - wanted, err := am.accountNeedsCertificateChallenges(ctx, accountID) + certCheckIDs, err := am.certificatePostureCheckIDs(ctx, accountID) if err != nil { log.WithContext(ctx).Debugf("cannot tell whether account %s needs certificate challenges: %v", accountID, err) return } - if !wanted { + if len(certCheckIDs) == 0 { return } am.certChallenges.Track(ctx, accountID) } - -// accountNeedsCertificateChallenges reports whether any of the account's posture checks -// asks its peers to prove a certificate. -func (am *DefaultAccountManager) accountNeedsCertificateChallenges(ctx context.Context, accountID string) (bool, error) { - checks, err := am.Store.GetAccountPostureChecks(ctx, store.LockingStrengthNone, accountID) - if err != nil { - return false, err - } - for _, check := range checks { - if check.Checks.CertificateCheck != nil { - return true, nil - } - } - return false, nil -} diff --git a/management/server/certificate_challenge_test.go b/management/server/certificate_challenge_test.go index ba75528c0..ffe9902a0 100644 --- a/management/server/certificate_challenge_test.go +++ b/management/server/certificate_challenge_test.go @@ -2,6 +2,7 @@ package server import ( "context" + "slices" "sync" "testing" "time" @@ -9,6 +10,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/certposture" ) @@ -218,3 +220,84 @@ func TestCertChallengeRefresher_RefreshesWithoutHoldingTheLock(t *testing.T) { t.Fatal("the refresh never ran") } } + +// TestCertificateChallengeTargets_MatchesThePerPeerRule holds the renewal to the same +// rule management uses when it hands out the challenge. The refresher asks "which peers +// get one", processPeerPostureChecks asks "does this peer get one", and the two have to +// answer identically: a peer the refresher forgets stops being renewed and silently +// falls out of its policies, which is the failure the refresher exists to prevent. +func TestCertificateChallengeTargets_MatchesThePerPeerRule(t *testing.T) { + const certCheck, osCheck = "check-cert", "check-os" + certCheckIDs := map[string]struct{}{certCheck: {}} + + groupPeers := map[string][]string{ + "g-devs": {"p1", "p2"}, + "g-servers": {"p3"}, + "g-empty": {}, + "g-mixed": {"p2", "p4"}, + } + peerGroups := map[string][]string{ + "p1": {"g-devs"}, + "p2": {"g-devs", "g-mixed"}, + "p3": {"g-servers"}, + "p4": {"g-mixed"}, + "p5": {}, + } + + policies := []*types.Policy{ + { + ID: "gated", Enabled: true, SourcePostureChecks: []string{certCheck}, + Rules: []*types.PolicyRule{{Enabled: true, Sources: []string{"g-devs"}}}, + }, + { + // A disabled policy hands out nothing. + ID: "disabled-policy", Enabled: false, SourcePostureChecks: []string{certCheck}, + Rules: []*types.PolicyRule{{Enabled: true, Sources: []string{"g-servers"}}}, + }, + { + // A disabled rule inside an enabled policy likewise. + ID: "disabled-rule", Enabled: true, SourcePostureChecks: []string{certCheck}, + Rules: []*types.PolicyRule{{Enabled: false, Sources: []string{"g-servers"}}}, + }, + { + // Gated on something other than a certificate: no nonce to renew. + ID: "other-check", Enabled: true, SourcePostureChecks: []string{osCheck}, + Rules: []*types.PolicyRule{{Enabled: true, Sources: []string{"g-mixed"}}}, + }, + { + // A peer named directly rather than through a group. + ID: "by-resource", Enabled: true, SourcePostureChecks: []string{osCheck, certCheck}, + Rules: []*types.PolicyRule{{ + Enabled: true, + SourceResource: types.Resource{Type: types.ResourceTypePeer, ID: "p4"}, + }}, + }, + { + // Destinations are never filtered, so being one earns no challenge. + ID: "as-destination", Enabled: true, SourcePostureChecks: []string{certCheck}, + Rules: []*types.PolicyRule{{Enabled: true, Sources: []string{"g-empty"}, Destinations: []string{"g-servers"}}}, + }, + } + + got := certificateChallengeTargets(policies, groupPeers, certCheckIDs) + + // The same question, asked one peer at a time the way the gRPC layer asks it. + var want []string + for peerID, groupIDs := range peerGroups { + for _, policy := range policies { + if !policy.Enabled { + continue + } + ids := processPeerPostureChecks(policy, peerID, groupIDs) + if slices.ContainsFunc(ids, func(id string) bool { _, ok := certCheckIDs[id]; return ok }) { + want = append(want, peerID) + break + } + } + } + + slices.Sort(got) + slices.Sort(want) + assert.Equal(t, want, got, "the peers the refresher renews must be exactly the peers that are sent a challenge") + assert.Equal(t, []string{"p1", "p2", "p4"}, got, "p3 is only reachable through disabled rules, p5 is in no source group") +}