[client] pqkem: carry a generation on PSK callbacks to drop stale applies

OnNewPSKReady could be applied out of order: two exchanges for a peer can derive
concurrently (one over signal, one over the data path), and the callbacks run
outside the manager lock, so an older exchange's apply could land after a newer
one and restore a stale WireGuard PSK, splitting the tunnel.

Give each exchange a per-peer monotonic generation, assigned under the lock at
creation so a later exchange always carries a higher one, and pass it to
OnNewPSKReady. The host adapter records the newest generation applied per peer
and drops any callback that is not newer, with the check-and-record atomic so the
slow SetPresharedKey call stays off that lock.

Found in cubic review on #7098 (client/internal/pqkem/callbacks.go:12).
This commit is contained in:
riccardom
2026-10-07 13:30:52 +02:00
parent d0f7e2ac7a
commit 1182239faa
7 changed files with 90 additions and 40 deletions
+2 -1
View File
@@ -589,7 +589,8 @@ func (e *Engine) startPQKEMManager(publicKey wgtypes.Key) error {
return nil
}
cbHandler := pqCallbackHandler{
wg: e.wgInterface,
wg: e.wgInterface,
applied: newAppliedGenerations(),
// On a persistent rekey failure, re-bootstrap the KEM over Signal: a fresh
// signalling offer starts a new exchange that overwrites the stalled PSK on both
// sides, recovering from a data-path desync.
+6 -1
View File
@@ -11,7 +11,12 @@ type CallbackHandler interface {
// fires it right after deriving the PSK from the offer and before sending that
// answer. A fired callback therefore means the key is derived locally, not that the
// peer has confirmed it — the next offer is the later acknowledgement.
OnNewPSKReady(remoteID RemoteID, psk PSK) error
//
// gen is a per-peer monotonic generation: a later exchange always carries a higher
// gen. Callbacks can be applied out of order (two exchanges deriving concurrently),
// so the host must ignore a call whose gen is not newer than the one it last applied
// for that peer, or it may restore an older PSK over a newer one and split the tunnel.
OnNewPSKReady(remoteID RemoteID, gen uint64, psk PSK) error
// OnRekeyFailed fires when an exchange fails to converge within the allotted time.
// Recovery is host-defined: the library reports the event and does not dictate the
+6 -3
View File
@@ -67,6 +67,7 @@ func (m *Manager) startExchangeLocked(remoteID RemoteID, viaSignal bool, ackID E
m.exchanges[remoteID] = &exchangeCtl{
id: id,
state: stateAwaitingAnswer,
gen: m.nextGenLocked(remoteID),
cancel: cancel,
lastSent: raw,
initiator: init,
@@ -117,7 +118,7 @@ func (m *Manager) processOffer(remoteID RemoteID, o *OfferMsg, via string) ([]by
return last, nil
}
// Reserve the slot so a concurrent duplicate offer bails.
m.exchanges[remoteID] = &exchangeCtl{id: o.ExchangeID, state: stateReserved}
m.exchanges[remoteID] = &exchangeCtl{id: o.ExchangeID, state: stateReserved, gen: m.nextGenLocked(remoteID)}
m.mu.Unlock()
answerBytes, psk, err := Respond(o.KEMOffer, m.binding(remoteID))
@@ -147,6 +148,7 @@ func (m *Manager) processOffer(remoteID RemoteID, o *OfferMsg, via string) ([]by
ex.state = stateAwaitingAck
ex.lastSent = raw
ex.pendingPSK = psk
gen := ex.gen
m.psks[remoteID] = psk
m.capable[remoteID] = true // a real KEM offer proves the peer runs the exchange
m.mu.Unlock()
@@ -157,7 +159,7 @@ func (m *Manager) processOffer(remoteID RemoteID, o *OfferMsg, via string) ([]by
// part of the commit: if the host can't program it, drop the exchange and do NOT send
// the answer, so the initiator times out and re-bootstraps instead of converging on a
// key we could not apply. Route it through the failure path.
if err := m.cbHandler.OnNewPSKReady(remoteID, psk); err != nil {
if err := m.cbHandler.OnNewPSKReady(remoteID, gen, psk); err != nil {
m.mu.Lock()
if c := m.exchanges[remoteID]; c != nil && c.id == o.ExchangeID {
delete(m.exchanges, remoteID)
@@ -238,6 +240,7 @@ func (m *Manager) processAnswer(remoteID RemoteID, a *AnswerMsg, via string) err
return nil
}
wasEstablished := m.established[remoteID]
gen := cur.gen
cur.state = stateAwaitingRekey
m.established[remoteID] = true
m.failures[remoteID] = 0
@@ -247,7 +250,7 @@ func (m *Manager) processAnswer(remoteID RemoteID, a *AnswerMsg, via string) err
m.debug("pqkem: PSK derived", "peer", remoteID, "exchange", idHex(a.ExchangeID), "role", "initiator", "via", via, "kind", kind, "psk_fp", pskFingerprint(psk))
if err := m.cbHandler.OnNewPSKReady(remoteID, psk); err != nil {
if err := m.cbHandler.OnNewPSKReady(remoteID, gen, psk); err != nil {
// Applying the PSK is part of the commit: if the host fails to program it, the
// exchange is not really converged. Drop it and route the failure through recovery
// so a re-bootstrap re-derives and re-applies, instead of leaving the peer parked
+12
View File
@@ -82,6 +82,7 @@ const (
type exchangeCtl struct {
id ExchangeID
state exchangeState
gen uint64 // local, per-peer monotonic generation; lets the host reject a stale PSK apply
cancel context.CancelFunc
lastSent []byte
initiator *Initiator
@@ -116,9 +117,19 @@ type Manager struct {
capable map[RemoteID]bool // peer runs the KEM (advertised a PQ port); false = known non-capable
peerAddrs map[RemoteID]netip.AddrPort // remoteID -> data-path endpoint (send routing)
peersByAddr map[netip.AddrPort]RemoteID // reverse: source endpoint -> remoteID (inbound)
genCounter map[RemoteID]uint64 // per-peer monotonic exchange generation
wait sync.WaitGroup
}
// nextGenLocked returns the next monotonic generation for a peer's exchange. Assumes
// m.mu is held. Because exchanges for a given peer are created under the lock in order,
// a later exchange always carries a higher generation, so the host can drop a PSK apply
// that arrives out of order (generation <= the one it already applied).
func (m *Manager) nextGenLocked(remoteID RemoteID) uint64 {
m.genCounter[remoteID]++
return m.genCounter[remoteID]
}
// NewManager builds a manager for the local peer identified by its peer identity key
// (used for the deterministic initiator role and the identity binding). A nil logger
// falls back to slog.Default(). Install the data-path transport with Start.
@@ -143,6 +154,7 @@ func NewManager(localID LocalID, h CallbackHandler, logger *slog.Logger) *Manage
capable: make(map[RemoteID]bool),
peerAddrs: make(map[RemoteID]netip.AddrPort),
peersByAddr: make(map[netip.AddrPort]RemoteID),
genCounter: make(map[RemoteID]uint64),
}
}
+1 -1
View File
@@ -74,7 +74,7 @@ func (m *Manager) startExchangeTest(remoteID RemoteID, viaSignal bool, ackID Exc
return m.startExchangeLocked(remoteID, viaSignal, ackID)
}
func (f *fakeWG) OnNewPSKReady(remoteID RemoteID, psk PSK) error {
func (f *fakeWG) OnNewPSKReady(remoteID RemoteID, _ uint64, psk PSK) error {
f.mu.Lock()
defer f.mu.Unlock()
f.psks[remoteID] = psk
+38 -3
View File
@@ -2,6 +2,7 @@ package internal
import (
"net/netip"
"sync"
"time"
log "github.com/sirupsen/logrus"
@@ -10,6 +11,31 @@ import (
"github.com/netbirdio/netbird/client/internal/pqkem"
)
// appliedGenerations records the newest PSK generation applied per peer, so an
// out-of-order OnNewPSKReady (two exchanges deriving concurrently) can't restore an
// older PSK over a newer one.
type appliedGenerations struct {
mu sync.Mutex
last map[string]uint64
}
func newAppliedGenerations() *appliedGenerations {
return &appliedGenerations{last: make(map[string]uint64)}
}
// claim reports whether gen is newer than the last applied for the peer; when it is, it
// records gen as the newest and returns true. The check and record are atomic, so the
// slow SetPresharedKey call runs outside this lock.
func (a *appliedGenerations) claim(peer string, gen uint64) bool {
a.mu.Lock()
defer a.mu.Unlock()
if gen <= a.last[peer] {
return false
}
a.last[peer] = gen
return true
}
// pqPresharedKeySetter is the subset of the WireGuard interface the ML-KEM callback
// needs: programming a peer's preshared key. *iface.WGIface satisfies it.
type pqPresharedKeySetter interface {
@@ -20,18 +46,27 @@ type pqPresharedKeySetter interface {
// engine-side implementation of pqkem.CallbackHandler.
type pqCallbackHandler struct {
wg pqPresharedKeySetter
// applied drops an out-of-order PSK apply so a stale callback can't overwrite a newer
// key. Nil disables the guard.
applied *appliedGenerations
// reoffer re-bootstraps the KEM over Signal for a peer (a fresh signalling offer)
// to recover from a persistent data-path rekey failure. Nil disables recovery.
reoffer func(remoteKey string)
}
// OnNewPSKReady programs the freshly derived PSK for the peer (updateOnly: a no-op
// if the peer is not present, mirroring Rosenpass).
func (h pqCallbackHandler) OnNewPSKReady(remoteID pqkem.RemoteID, psk pqkem.PSK) error {
// if the peer is not present, mirroring Rosenpass). A callback whose generation is not
// newer than the last applied for the peer is dropped, so a reordered apply can't
// restore an older PSK over a newer one.
func (h pqCallbackHandler) OnNewPSKReady(remoteID pqkem.RemoteID, gen uint64, psk pqkem.PSK) error {
if h.applied != nil && !h.applied.claim(string(remoteID), gen) {
log.Tracef("pqkem: dropping stale PSK apply for peer %s (gen %d)", remoteID, gen)
return nil
}
// updateOnly: applies to an already-configured peer (rotation). At bootstrap the
// peer is not configured yet, so this is a no-op there and the PSK is instead
// pulled at peer-config time (pqHandshaker.PSK / conn.presharedKey).
log.Tracef("pqkem: programming PSK for peer %s", remoteID)
log.Tracef("pqkem: programming PSK for peer %s (gen %d)", remoteID, gen)
return h.wg.SetPresharedKey(string(remoteID), wgtypes.Key(psk), true)
}
+25 -31
View File
@@ -1,37 +1,31 @@
package internal
import (
"testing"
import "testing"
"github.com/stretchr/testify/require"
// TestAppliedGenerations_DropsStale verifies the out-of-order guard: a generation is
// accepted only when it is strictly newer than the last one applied for that peer, so a
// reordered PSK callback cannot restore an older key over a newer one.
func TestAppliedGenerations_DropsStale(t *testing.T) {
a := newAppliedGenerations()
"github.com/netbirdio/netbird/client/internal/pqkem"
)
if !a.claim("peerA", 1) {
t.Fatal("first generation must be accepted")
}
if !a.claim("peerA", 2) {
t.Fatal("a newer generation must be accepted")
}
if a.claim("peerA", 2) {
t.Fatal("re-applying the same generation must be dropped")
}
if a.claim("peerA", 1) {
t.Fatal("an older generation arriving late must be dropped")
}
if !a.claim("peerA", 3) {
t.Fatal("a newer generation after a dropped stale one must still be accepted")
}
type pqNoopHandler struct{}
func (pqNoopHandler) OnNewPSKReady(pqkem.RemoteID, pqkem.PSK) error { return nil }
func (pqNoopHandler) OnRekeyFailed(pqkem.RemoteID) error { return nil }
// TestPQAdapter_CapabilityRoleAware locks the role-aware capability signal: the KEM
// payload only flows initiator-offer -> responder-answer, so an empty message in the
// other direction comes from a perfectly capable peer and must NOT flag it. Only the
// message that should carry material (the answer we receive as initiator) marks a peer
// non-capable when empty.
func TestPQAdapter_CapabilityRoleAware(t *testing.T) {
// localID "zzzz" > "aaaa" => this manager is the KEM initiator for peer "aaaa".
mgr := pqkem.NewManager("zzzz", pqNoopHandler{}, nil)
defer mgr.Stop()
h := pqHandshaker{mgr: mgr}
// An empty OFFER from our peer is normal here: as the initiator's responder it puts
// its material in the answer, not the offer. It must not disable our offering.
h.AnswerPayload("aaaa", nil)
payload, _ := h.OfferPayload("aaaa")
require.NotNil(t, payload, "an empty offer from a responder-role peer must not mark it non-capable")
// An empty ANSWER to our offer means the peer does not run the KEM -> stop offering.
h.OnAnswer("aaaa", nil)
payload2, _ := h.OfferPayload("aaaa")
require.Nil(t, payload2, "an empty answer to our offer marks the peer non-capable, so we stop offering")
// Generations are tracked independently per peer.
if !a.claim("peerB", 1) {
t.Fatal("a different peer's first generation must be accepted")
}
}