Bound certificate proof collection so a stuck token or keychain cannot hold the sync loop

This commit is contained in:
Viktor Liu
2026-10-01 08:21:13 +02:00
parent ab2f8972be
commit 83e7bf8ab4
3 changed files with 160 additions and 1 deletions
+67
View File
@@ -0,0 +1,67 @@
package certproof
import (
"context"
"sync/atomic"
"time"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/shared/management/certposture"
"github.com/netbirdio/netbird/shared/management/proto"
)
const collectTimeout = 10 * time.Second
// Collector runs CollectProofs with a deadline and at most one collection at a time.
// Token, TPM and keychain calls cannot be interrupted, so a collection that overruns is
// abandoned rather than awaited, and a new one is refused until it has finished. The
// zero value is ready to use.
type Collector struct {
busy atomic.Bool
// timeout overrides collectTimeout when set.
timeout time.Duration
}
// Collect answers the certificate challenges in checks, returning no proofs when there
// are no challenges, when a previous collection is still running, or when this one does
// not finish in time. Missing proofs fail the certificate check on management.
func (c *Collector) Collect(ctx context.Context, checks []*proto.Checks, peerKey []byte, cfg Config) []certposture.Proof {
return c.collect(ctx, checks, func(ctx context.Context) []certposture.Proof {
return CollectProofs(ctx, checks, peerKey, cfg)
})
}
func (c *Collector) collect(ctx context.Context, checks []*proto.Checks, run func(context.Context) []certposture.Proof) []certposture.Proof {
if len(certificateChallenges(checks)) == 0 {
return nil
}
if !c.busy.CompareAndSwap(false, true) {
log.Warnf("certificate posture: previous proof collection is still running, sending no proofs")
return nil
}
ctx, cancel := context.WithTimeout(ctx, c.deadline())
defer cancel()
done := make(chan []certposture.Proof, 1)
go func() {
defer c.busy.Store(false)
done <- run(ctx)
}()
select {
case proofs := <-done:
return proofs
case <-ctx.Done():
log.Warnf("certificate posture: proof collection did not finish within %s, sending no proofs", c.deadline())
return nil
}
}
func (c *Collector) deadline() time.Duration {
if c.timeout > 0 {
return c.timeout
}
return collectTimeout
}
@@ -0,0 +1,88 @@
package certproof
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/shared/management/certposture"
"github.com/netbirdio/netbird/shared/management/proto"
)
var challengeChecks = []*proto.Checks{{CertificateChallenge: &proto.CertificateChallenge{Nonce: []byte("nonce")}}}
func TestCollector_SkipsChecksWithoutChallenges(t *testing.T) {
var c Collector
called := false
proofs := c.collect(context.Background(), []*proto.Checks{{Files: []string{"/bin/agent"}}}, func(context.Context) []certposture.Proof {
called = true
return nil
})
assert.Nil(t, proofs)
assert.False(t, called, "no store is touched when no check carries a challenge")
}
func TestCollector_ReturnsProofs(t *testing.T) {
var c Collector
want := []certposture.Proof{{Nonce: []byte("nonce")}}
proofs := c.collect(context.Background(), challengeChecks, func(context.Context) []certposture.Proof { return want })
assert.Equal(t, want, proofs, "a collection that finishes in time is returned as is")
}
func TestCollector_AbandonsStuckCollection(t *testing.T) {
c := Collector{timeout: 50 * time.Millisecond}
release := make(chan struct{})
finished := make(chan struct{})
// A token or keychain call that ignores its context and blocks well past the deadline.
stuck := func(context.Context) []certposture.Proof {
defer close(finished)
<-release
return []certposture.Proof{{Nonce: []byte("late")}}
}
start := time.Now()
proofs := c.collect(context.Background(), challengeChecks, stuck)
assert.Nil(t, proofs, "an overrunning collection yields no proofs")
assert.Less(t, time.Since(start), time.Second, "the caller is released at the deadline, not when the call returns")
called := false
proofs = c.collect(context.Background(), challengeChecks, func(context.Context) []certposture.Proof {
called = true
return nil
})
assert.Nil(t, proofs)
assert.False(t, called, "no second collection starts while the first is still running")
close(release)
<-finished
require.Eventually(t, func() bool { return !c.busy.Load() }, time.Second, 5*time.Millisecond, "the collector frees up once the stuck call returns")
want := []certposture.Proof{{Nonce: []byte("nonce")}}
assert.Equal(t, want, c.collect(context.Background(), challengeChecks, func(context.Context) []certposture.Proof { return want }),
"collection works again after the stuck call returned")
}
func TestCollector_CancelsContextAtDeadline(t *testing.T) {
c := Collector{timeout: 20 * time.Millisecond}
cancelled := make(chan struct{})
c.collect(context.Background(), challengeChecks, func(ctx context.Context) []certposture.Proof {
<-ctx.Done()
close(cancelled)
return nil
})
select {
case <-cancelled:
case <-time.After(time.Second):
t.Fatal("a collection that honours its context, like the helper process, must see it cancelled")
}
}
+5 -1
View File
@@ -276,6 +276,9 @@ type Engine struct {
// checks are the client-applied posture checks that need to be evaluated on the client
checks []*mgmProto.Checks
// certProofs answers the certificate challenges in checks within a bounded time.
certProofs certproof.Collector
infoSource system.InfoSource
relayManager *relayClient.Manager
@@ -1296,9 +1299,10 @@ func (e *Engine) applyInfoFlags(info *system.Info) {
// attachCertificateProofs answers the certificate challenges in checks with the
// certificates reachable on this device, signing each challenge nonce for our peer key.
// Collection is bounded in time because callers hold the sync loop while it runs.
func (e *Engine) attachCertificateProofs(info *system.Info, checks []*mgmProto.Checks) {
peerKey := e.config.WgPrivateKey.PublicKey()
info.CertificateProofs = certproof.CollectProofs(e.ctx, checks, peerKey[:], e.config.CertStore)
info.CertificateProofs = e.certProofs.Collect(e.ctx, checks, peerKey[:], e.config.CertStore)
}
func (e *Engine) currentSystemInfo(ctx context.Context) *system.Info {