diff --git a/management/server/account.go b/management/server/account.go index 6ccf673f5..ea27f0e0e 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -102,6 +102,8 @@ type DefaultAccountManager struct { peerInactivityExpiry Scheduler + certChallenges *certChallengeRefresher + // userDeleteFromIDPEnabled allows to delete user from IDP when user is deleted from account userDeleteFromIDPEnabled bool @@ -230,6 +232,10 @@ func BuildManager( disableDefaultPolicy: disableDefaultPolicy, } + am.certChallenges = newCertChallengeRefresher(am.refreshCertificateChallenges) + // The loop outlives this call, so it must not inherit its cancellation. + am.certChallenges.Start(context.WithoutCancel(ctx)) + am.networkMapController.StartWarmup(ctx) accountsCounter, err := store.GetAccountsCounter(ctx) @@ -900,6 +906,7 @@ func (am *DefaultAccountManager) DeleteAccount(ctx context.Context, accountID, u } // cancel peer login expiry job am.peerLoginExpiry.Cancel(ctx, []string{account.Id}) + am.certChallenges.Forget(account.Id) meta := map[string]any{"account_id": account.Id, "domain": account.Domain, "created_at": account.CreatedAt} am.StoreEvent(ctx, userID, accountID, accountID, activity.AccountDeleted, meta) diff --git a/management/server/certificate_challenge.go b/management/server/certificate_challenge.go new file mode 100644 index 000000000..fcedd056e --- /dev/null +++ b/management/server/certificate_challenge.go @@ -0,0 +1,195 @@ +package server + +import ( + "context" + "encoding/binary" + "hash/fnv" + "sync" + "time" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/certposture" +) + +const ( + // certChallengePeriod is how often an account whose policies carry a certificate + // posture check is pushed a fresh challenge. A nonce is accepted for its own window + // and the one before it, so one issued at the very end of a window lives only + // certposture.Window. A third of that leaves a missed run well clear of the edge, + // where a half would put it exactly on it. + certChallengePeriod = certposture.Window / 3 + + // certChallengeTick is how often the refresher looks for accounts that are due. The + // period is measured in hours, so this only has to be fine enough to spread the + // accounts over it. + certChallengeTick = 15 * time.Minute +) + +// certChallengeRefresher pushes a fresh certificate challenge to the peers of every +// account that needs one, from a single goroutine. +// +// A nonce only reaches a peer attached to a network map, and a quiet account sends no +// map. Without this the peer re-sends an expired nonce on its next sync, management +// rejects its whole proof set and drops its certificates, and it loses every policy +// gated on the check until something else changes. +// +// Accounts are held in a map rather than a queue ordered by due time: one pass over +// them every certChallengeTick costs nothing next to a period measured in hours, and it +// avoids having to re-arm a timer whenever an account that falls due sooner is added. +type certChallengeRefresher struct { + mu sync.Mutex + due map[string]time.Time + + period time.Duration + tick time.Duration + now func() time.Time + // refresh pushes the account's peers an update, and reports whether the account + // still wants challenges at all. + refresh func(ctx context.Context, accountID string) bool +} + +func newCertChallengeRefresher(refresh func(ctx context.Context, accountID string) bool) *certChallengeRefresher { + return &certChallengeRefresher{ + due: map[string]time.Time{}, + period: certChallengePeriod, + tick: certChallengeTick, + now: time.Now, + refresh: refresh, + } +} + +// Start runs the refresh loop until ctx is done. +func (r *certChallengeRefresher) Start(ctx context.Context) { + go r.run(ctx) +} + +// Track starts refreshing accountID, spreading its first run over one period so that a +// global window rollover does not fan out to every account in the same moment. An +// account already tracked keeps the schedule it has. +func (r *certChallengeRefresher) Track(ctx context.Context, accountID string) { + r.mu.Lock() + defer r.mu.Unlock() + if _, ok := r.due[accountID]; ok { + return + } + r.due[accountID] = r.now().Add(offsetWithin(accountID, r.period)) + log.WithContext(ctx).Debugf("tracking certificate challenge refresh for account %s", accountID) +} + +// Forget stops refreshing accountID. +func (r *certChallengeRefresher) Forget(accountID string) { + r.mu.Lock() + defer r.mu.Unlock() + delete(r.due, accountID) +} + +func (r *certChallengeRefresher) tracked(accountID string) bool { + r.mu.Lock() + defer r.mu.Unlock() + _, ok := r.due[accountID] + return ok +} + +func (r *certChallengeRefresher) run(ctx context.Context) { + ticker := time.NewTicker(r.tick) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + for _, accountID := range r.takeDue() { + if !r.refresh(ctx, accountID) { + r.Forget(accountID) + } + } + } + } +} + +// takeDue returns the accounts due now and books their next run straight away, so a +// slow refresh cannot make an account fall due twice, and so the refresh itself runs +// without the lock. +func (r *certChallengeRefresher) takeDue() []string { + r.mu.Lock() + defer r.mu.Unlock() + + now := r.now() + var due []string + for accountID, at := range r.due { + if at.After(now) { + continue + } + due = append(due, accountID) + r.due[accountID] = now.Add(r.period) + } + return due +} + +// offsetWithin maps a key to a stable duration in [0, period). +func offsetWithin(key string, period time.Duration) time.Duration { + h := fnv.New64a() + //nolint:errcheck // hash.Write never returns an error + h.Write([]byte(key)) + 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. +func (am *DefaultAccountManager) refreshCertificateChallenges(ctx context.Context, accountID string) bool { + wanted, err := am.accountNeedsCertificateChallenges(ctx, accountID) + if err != nil { + log.WithContext(ctx).Debugf("cannot tell whether account %s still needs certificate challenges: %v", accountID, err) + // Keep the account tracked: a store error now says nothing about its checks. + return true + } + if !wanted { + log.WithContext(ctx).Debugf("account %s has no certificate posture check left, stopping challenge refresh", accountID) + return false + } + + log.WithContext(ctx).Debugf("refreshing certificate challenges for account %s", accountID) + am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{ + Resource: types.UpdateResourcePostureCheck, + Operation: types.UpdateOperationRefresh, + }) + return true +} + +// 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) { + if am.certChallenges.tracked(accountID) { + return + } + + wanted, err := am.accountNeedsCertificateChallenges(ctx, accountID) + if err != nil { + log.WithContext(ctx).Debugf("cannot tell whether account %s needs certificate challenges: %v", accountID, err) + return + } + if !wanted { + 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 new file mode 100644 index 000000000..f5e8efe44 --- /dev/null +++ b/management/server/certificate_challenge_test.go @@ -0,0 +1,206 @@ +package server + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/certposture" +) + +// fakeClock lets a test drive the refresher's schedule without waiting for it. +type fakeClock struct { + mu sync.Mutex + at time.Time +} + +func (c *fakeClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.at +} + +func (c *fakeClock) Advance(d time.Duration) { + c.mu.Lock() + defer c.mu.Unlock() + c.at = c.at.Add(d) +} + +// scheduleRefresher builds a refresher on a clock the test drives and never starts its +// loop, so the schedule can be checked by calling takeDue directly. +func scheduleRefresher(refresh func(accountID string) bool) (*certChallengeRefresher, *fakeClock) { + clock := &fakeClock{at: time.Date(2026, 10, 1, 0, 0, 0, 0, time.UTC)} + r := newCertChallengeRefresher(func(_ context.Context, accountID string) bool { + return refresh(accountID) + }) + r.now = clock.Now + return r, clock +} + +// runningRefresher builds a refresher whose period and tick are short enough to observe +// in a test, and starts its loop. +func runningRefresher(t *testing.T, refresh func(accountID string) bool) *certChallengeRefresher { + t.Helper() + r := newCertChallengeRefresher(func(_ context.Context, accountID string) bool { + return refresh(accountID) + }) + r.period = 10 * time.Millisecond + r.tick = time.Millisecond + r.Start(t.Context()) + return r +} + +func TestCertChallengePeriod_LeavesRoomForAMissedRun(t *testing.T) { + // A nonce issued at the very end of a window is accepted for that window and the + // next one only, so the shortest life a peer can be handed is one window. Renewing + // has to stay clear of that edge even when a run is missed. + assert.Less(t, 2*certChallengePeriod, certposture.Window, + "a missed refresh must still leave the peer's nonce valid, with margin") +} + +func TestCertChallengeRefresher_SpreadsAccountsOverThePeriod(t *testing.T) { + // Every account on an instance shares the same global challenge window, so a + // restart arms them all at once. The offset is what keeps them from fanning out + // to their peers in the same moment. + r, clock := scheduleRefresher(func(string) bool { return true }) + + ids := []string{"account-a", "account-b", "account-c", "account-d", "account-e", "account-f"} + for _, id := range ids { + r.Track(context.Background(), id) + } + + start := clock.Now() + offsets := map[time.Duration]bool{} + for _, id := range ids { + offset := r.due[id].Sub(start) + assert.GreaterOrEqual(t, offset, time.Duration(0), "account %s is due in the past", id) + assert.Less(t, offset, r.period, "account %s is due beyond one period", id) + offsets[offset] = true + } + assert.Greater(t, len(offsets), 1, "all accounts were given the same offset, which defeats the spreading") +} + +func TestCertChallengeOffset_IsStablePerAccount(t *testing.T) { + assert.Equal(t, offsetWithin("account-a", certChallengePeriod), offsetWithin("account-a", certChallengePeriod), + "the same account must keep its slot across restarts") + assert.NotEqual(t, offsetWithin("account-a", certChallengePeriod), offsetWithin("account-b", certChallengePeriod), + "two accounts must not share a slot") +} + +func TestCertChallengeRefresher_RenewsOncePerPeriod(t *testing.T) { + r, clock := scheduleRefresher(func(string) bool { return true }) + r.Track(context.Background(), "account-a") + + assert.Empty(t, r.takeDue(), "an account just tracked is not due before its offset elapses") + + clock.Advance(r.period) + assert.Equal(t, []string{"account-a"}, r.takeDue(), "the account is due once its offset has elapsed") + + clock.Advance(r.period / 2) + assert.Empty(t, r.takeDue(), "the next run is booked a full period out") + + clock.Advance(r.period / 2) + assert.Equal(t, []string{"account-a"}, r.takeDue(), "the account is due again one period later") +} + +func TestCertChallengeRefresher_BooksTheNextRunBeforeRefreshing(t *testing.T) { + // takeDue reserves the next run while it holds the lock, so a refresh that outlives + // a tick cannot have the same account handed out twice. + r, clock := scheduleRefresher(func(string) bool { return true }) + r.Track(context.Background(), "account-a") + clock.Advance(r.period) + + require.Len(t, r.takeDue(), 1, "the account is due") + assert.Empty(t, r.takeDue(), "a second pass at the same instant must not hand out the account again") +} + +func TestCertChallengeRefresher_TrackIsIdempotent(t *testing.T) { + r, clock := scheduleRefresher(func(string) bool { return true }) + + r.Track(context.Background(), "account-a") + first := r.due["account-a"] + + clock.Advance(certChallengePeriod) + r.Track(context.Background(), "account-a") + + assert.Equal(t, first, r.due["account-a"], "re-tracking must not push the next run further out") +} + +func TestCertChallengeRefresher_RefreshesATrackedAccount(t *testing.T) { + refreshed := make(chan string, 4) + r := runningRefresher(t, func(accountID string) bool { + select { + case refreshed <- accountID: + default: + } + return true + }) + r.Track(context.Background(), "account-a") + + select { + case got := <-refreshed: + assert.Equal(t, "account-a", got) + case <-time.After(3 * time.Second): + t.Fatal("a tracked account was never refreshed") + } + assert.True(t, r.tracked("account-a"), "an account that still wants challenges stays tracked") +} + +func TestCertChallengeRefresher_DropsAnAccountThatNoLongerWantsChallenges(t *testing.T) { + var mu sync.Mutex + var calls int + + r := runningRefresher(t, func(string) bool { + mu.Lock() + defer mu.Unlock() + calls++ + return false + }) + r.Track(context.Background(), "account-a") + + require.Eventually(t, func() bool { return !r.tracked("account-a") }, 3*time.Second, time.Millisecond, + "an account whose refresh reports it no longer wants challenges must be dropped") + + time.Sleep(20 * r.tick) + mu.Lock() + defer mu.Unlock() + assert.Equal(t, 1, calls, "the account must not be refreshed again after being dropped") +} + +func TestCertChallengeRefresher_RefreshesWithoutHoldingTheLock(t *testing.T) { + // The refresh fans out to every peer of the account, so holding the lock across it + // would stall every peer that connects meanwhile. Calling back into the refresher + // from inside the refresh deadlocks if the loop still holds it. + reentered := make(chan bool, 1) + var once sync.Once + + r := newCertChallengeRefresher(nil) + r.period = 10 * time.Millisecond + r.tick = time.Millisecond + r.refresh = func(_ context.Context, accountID string) bool { + once.Do(func() { + answered := make(chan bool, 1) + go func() { answered <- r.tracked(accountID) }() + select { + case got := <-answered: + reentered <- got + case <-time.After(2 * time.Second): + reentered <- false + } + }) + return true + } + r.Start(t.Context()) + r.Track(context.Background(), "account-a") + + select { + case ok := <-reentered: + assert.True(t, ok, "the refresher was not reachable while a refresh was in flight") + case <-time.After(5 * time.Second): + t.Fatal("the refresh never ran") + } +} diff --git a/management/server/peer.go b/management/server/peer.go index 9f5572252..5f55de9cc 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -117,6 +117,10 @@ func (am *DefaultAccountManager) MarkPeerConnected(ctx context.Context, peerPubK return err } + // Not gated on SSO, unlike the expirations above: a certificate challenge goes to + // every peer the check applies to, however it was enrolled. + am.trackCertificateChallenges(ctx, accountID) + // A login-expired peer reconnecting, or an embedded proxy peer flipping to // connected (which triggers SynthesizePrivateServiceZones), must refresh the // peers reachable from it. The embedded-proxy fan-out tolerates a dispatch error. diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index 081226866..b51e38a82 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -87,6 +87,12 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI am.ExpandAndUpdateAffected(ctx, accountID, snap, change) + // The save itself reaches the peers, but an account that stays quiet afterwards + // would otherwise wait for a reconnect before its challenges start being renewed. + if postureChecks.Checks.CertificateCheck != nil { + am.trackCertificateChallenges(ctx, accountID) + } + return postureChecks, nil } diff --git a/management/server/types/update_reason.go b/management/server/types/update_reason.go index 9d752da9a..e829a1b2a 100644 --- a/management/server/types/update_reason.go +++ b/management/server/types/update_reason.go @@ -34,4 +34,7 @@ const ( UpdateOperationCreate UpdateOperation = "create" UpdateOperationUpdate UpdateOperation = "update" UpdateOperationDelete UpdateOperation = "delete" + // UpdateOperationRefresh is a periodic push that carries no change of its own, so + // it stays out of the counters that track what an administrator actually edited. + UpdateOperationRefresh UpdateOperation = "refresh" )