diff --git a/client/internal/pqkem/bench_test.go b/client/internal/pqkem/bench_test.go index eabcc3a0a..7e500de77 100644 --- a/client/internal/pqkem/bench_test.go +++ b/client/internal/pqkem/bench_test.go @@ -19,8 +19,14 @@ func BenchmarkX25519Keygen(b *testing.B) { func BenchmarkX25519ECDH(b *testing.B) { c := ecdh.X25519() - a, _ := c.GenerateKey(rand.Reader) - p, _ := c.GenerateKey(rand.Reader) + a, err := c.GenerateKey(rand.Reader) + if err != nil { + b.Fatal(err) + } + p, err := c.GenerateKey(rand.Reader) + if err != nil { + b.Fatal(err) + } pub := p.PublicKey() b.ResetTimer() for i := 0; i < b.N; i++ { @@ -39,7 +45,10 @@ func BenchmarkMLKEMKeygen(b *testing.B) { } func BenchmarkMLKEMEncaps(b *testing.B) { - dk, _ := mlkem.GenerateKey768() + dk, err := mlkem.GenerateKey768() + if err != nil { + b.Fatal(err) + } ek := dk.EncapsulationKey() b.ResetTimer() for i := 0; i < b.N; i++ { @@ -48,7 +57,10 @@ func BenchmarkMLKEMEncaps(b *testing.B) { } func BenchmarkMLKEMDecaps(b *testing.B) { - dk, _ := mlkem.GenerateKey768() + dk, err := mlkem.GenerateKey768() + if err != nil { + b.Fatal(err) + } _, ct := dk.EncapsulationKey().Encapsulate() b.ResetTimer() for i := 0; i < b.N; i++ { diff --git a/client/internal/pqkem/concurrency_test.go b/client/internal/pqkem/concurrency_test.go index 0303750a9..bcf389404 100644 --- a/client/internal/pqkem/concurrency_test.go +++ b/client/internal/pqkem/concurrency_test.go @@ -32,7 +32,7 @@ func TestConcurrency_RecoversViaResignalAfterDataPathBreak(t *testing.T) { // Data path breaks: the rotation can no longer converge -> OnRekeyFailed. lbB.drop.Store(true) - _, err := dB.startExchange("aaaa", false, ExchangeID{}) + _, err := dB.startExchangeTest("aaaa", false, ExchangeID{}) require.NoError(t, err) require.Eventually(t, func() bool { return failedCount(wgB) >= 1 }, time.Second, 5*time.Millisecond) diff --git a/client/internal/pqkem/convergence.go b/client/internal/pqkem/convergence.go index e504cbfbd..f818590b1 100644 --- a/client/internal/pqkem/convergence.go +++ b/client/internal/pqkem/convergence.go @@ -28,7 +28,7 @@ func pskFingerprint(psk PSK) string { // same peer cannot both create an exchange. It also refuses to start (and to Add to the // wait group) once the manager is stopping, so it never races Manager.Stop's Wait. func (m *Manager) startExchangeLocked(remoteID RemoteID, viaSignal bool, ackID ExchangeID) ([]byte, error) { - if m.stopping { + if m.rootCtx.Err() != nil { return nil, fmt.Errorf("manager stopping") } init, err := NewInitiator() diff --git a/client/internal/pqkem/convergence_test.go b/client/internal/pqkem/convergence_test.go index f4ce251c4..fed15d0ea 100644 --- a/client/internal/pqkem/convergence_test.go +++ b/client/internal/pqkem/convergence_test.go @@ -61,14 +61,14 @@ func TestManager_RekeyToleratesKFailures(t *testing.T) { // K-1 data-path rekeys must NOT raise OnRekeyFailed. for i := 0; i < DefaultMaxRekeyFailures-1; i++ { - _, err := dB.startExchange("aaaa", false, ExchangeID{}) + _, err := dB.startExchangeTest("aaaa", false, ExchangeID{}) require.NoError(t, err) time.Sleep(50 * time.Millisecond) } require.Equal(t, 0, failedCount(wgB), "no failure before K attempts") // The K-th failure raises it once. - _, err := dB.startExchange("aaaa", false, ExchangeID{}) + _, err := dB.startExchangeTest("aaaa", false, ExchangeID{}) require.NoError(t, err) require.Eventually(t, func() bool { return failedCount(wgB) == 1 }, time.Second, 5*time.Millisecond) } diff --git a/client/internal/pqkem/kem_test.go b/client/internal/pqkem/kem_test.go index 499c82470..9a2c9c8f9 100644 --- a/client/internal/pqkem/kem_test.go +++ b/client/internal/pqkem/kem_test.go @@ -33,19 +33,36 @@ func TestExchange_PSKBoundToPeerIdentities(t *testing.T) { init, err := NewInitiator() require.NoError(t, err) - // responder computes with the honest pair... - _, pskHonest, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA}) + answer, _, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA}) + require.NoError(t, err) + + // Finish twice over the SAME KEM material (same offer/answer/secrets), changing only + // the peer identity binding: the differing PSK is attributable to the binding alone. + pskHonest, err := init.Finish(answer, Binding{LocalID: wgA, RemoteID: wgB}) require.NoError(t, err) - // ...a second responder run with a different peer identity yields a different PSK, - // even though the KEM material would otherwise combine identically. wgC := []byte("peer-C-wireguard-pubkey-32bytes!") - _, pskWrong, err := Respond(init.Offer(), Binding{LocalID: wgC, RemoteID: wgA}) + pskWrong, err := init.Finish(answer, Binding{LocalID: wgA, RemoteID: wgC}) require.NoError(t, err) require.NotEqual(t, pskHonest, pskWrong, "PSK must be bound to the peer pair") } +func TestExchange_RejectsEmptyBinding(t *testing.T) { + init, err := NewInitiator() + require.NoError(t, err) + answer, _, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA}) + require.NoError(t, err) + + // A PSK not bound to both identities could be transplanted to another peer pair. + _, err = init.Finish(answer, Binding{}) + require.Error(t, err, "empty binding must be rejected") + _, err = init.Finish(answer, Binding{LocalID: wgA}) + require.Error(t, err, "missing RemoteID must be rejected") + _, _, err = Respond(init.Offer(), Binding{RemoteID: wgA}) + require.Error(t, err, "missing LocalID must be rejected") +} + func TestExchange_RejectsMalformedMessages(t *testing.T) { init, err := NewInitiator() require.NoError(t, err) diff --git a/client/internal/pqkem/manager.go b/client/internal/pqkem/manager.go index 53fd564c4..4e6ce3bbe 100644 --- a/client/internal/pqkem/manager.go +++ b/client/internal/pqkem/manager.go @@ -250,7 +250,12 @@ func (m *Manager) RemovePeer(remoteID RemoteID) { // Stop cancels all in-flight exchanges, closes the transport, and waits for the // exchange goroutines to exit. func (m *Manager) Stop() { + // Cancel the root context under the lock, before Wait: startExchangeLocked checks + // rootCtx.Err() under the same lock before it Adds to the wait group, so once Stop + // has cancelled here no new Add can race Wait. + m.mu.Lock() m.rootCancel() + m.mu.Unlock() m.wait.Wait() m.mu.Lock() t := m.transport @@ -296,9 +301,11 @@ func (m *Manager) SignalOffer(remoteID RemoteID) ([]byte, error) { m.mu.Unlock() return last, nil } + // Hold the lock across the check above and the install so a concurrent SignalOffer + // for the same peer can't also start an exchange. bootstrap offer acks nothing. + raw, err := m.startExchangeLocked(remoteID, true, ExchangeID{}) m.mu.Unlock() - // bootstrap offer acknowledges nothing (zero AckID). - return m.startExchange(remoteID, true, ExchangeID{}) + return raw, err } // ShouldSendBootstrapOffer reports whether we should emit a fresh KEM offer to kick a @@ -402,17 +409,17 @@ func (m *Manager) OnDataPathRekeyed(remoteID RemoteID, sinceActivity time.Durati m.mu.Lock() ex := m.exchanges[remoteID] chain := ex != nil && ex.state == stateAwaitingRekey - var ackID ExchangeID - if chain { - ackID = ex.id - } - m.mu.Unlock() - - m.trace("pqkem: data-path rekey signal", "peer", remoteID, "chaining", chain) if !chain { + m.mu.Unlock() + m.trace("pqkem: data-path rekey signal", "peer", remoteID, "chaining", false) return } - offer, err := m.startExchange(remoteID, false, ackID) + // Hold the lock across the awaitingRekey check and the install so two rekey clocks + // can't each start a chained exchange for the same peer. + offer, err := m.startExchangeLocked(remoteID, false, ex.id) + m.mu.Unlock() + + m.trace("pqkem: data-path rekey signal", "peer", remoteID, "chaining", true) if err != nil { m.logger.Error("pqkem: chain offer failed to start", "peer", remoteID, "err", err) return diff --git a/client/internal/pqkem/manager_test.go b/client/internal/pqkem/manager_test.go index eab264478..b7a13166a 100644 --- a/client/internal/pqkem/manager_test.go +++ b/client/internal/pqkem/manager_test.go @@ -65,6 +65,14 @@ type fakeWG struct { func newFakeWG() *fakeWG { return &fakeWG{psks: map[RemoteID]PSK{}} } +// startExchangeTest drives startExchangeLocked with the lock held, for tests that kick an +// exchange directly (production callers hold m.mu across their idempotency check). +func (m *Manager) startExchangeTest(remoteID RemoteID, viaSignal bool, ackID ExchangeID) ([]byte, error) { + m.mu.Lock() + defer m.mu.Unlock() + return m.startExchangeLocked(remoteID, viaSignal, ackID) +} + func (f *fakeWG) OnNewPSKReady(remoteID RemoteID, psk PSK) error { f.mu.Lock() defer f.mu.Unlock()