diff --git a/client/internal/certproof/collector.go b/client/internal/certproof/collector.go new file mode 100644 index 000000000..de35bfa04 --- /dev/null +++ b/client/internal/certproof/collector.go @@ -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 +} diff --git a/client/internal/certproof/collector_test.go b/client/internal/certproof/collector_test.go new file mode 100644 index 000000000..73503e9c6 --- /dev/null +++ b/client/internal/certproof/collector_test.go @@ -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") + } +} diff --git a/client/internal/engine.go b/client/internal/engine.go index 0a62e7326..fcadbb939 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -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 {