diff --git a/client/internal/engine.go b/client/internal/engine.go index a52b561ec..d3becfb04 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -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. diff --git a/client/internal/pqkem/callbacks.go b/client/internal/pqkem/callbacks.go index 2b7c3b1c6..f1b681648 100644 --- a/client/internal/pqkem/callbacks.go +++ b/client/internal/pqkem/callbacks.go @@ -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 diff --git a/client/internal/pqkem/convergence.go b/client/internal/pqkem/convergence.go index fe7d309a6..d7fc65c99 100644 --- a/client/internal/pqkem/convergence.go +++ b/client/internal/pqkem/convergence.go @@ -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 diff --git a/client/internal/pqkem/manager.go b/client/internal/pqkem/manager.go index c5cdad09c..76d07a1a7 100644 --- a/client/internal/pqkem/manager.go +++ b/client/internal/pqkem/manager.go @@ -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), } } diff --git a/client/internal/pqkem/manager_test.go b/client/internal/pqkem/manager_test.go index 94ab70490..239037f03 100644 --- a/client/internal/pqkem/manager_test.go +++ b/client/internal/pqkem/manager_test.go @@ -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 diff --git a/client/internal/pqkem_adapter.go b/client/internal/pqkem_adapter.go index 8814dc229..4bf20145c 100644 --- a/client/internal/pqkem_adapter.go +++ b/client/internal/pqkem_adapter.go @@ -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) } diff --git a/client/internal/pqkem_adapter_test.go b/client/internal/pqkem_adapter_test.go index dec3b5943..9674075bb 100644 --- a/client/internal/pqkem_adapter_test.go +++ b/client/internal/pqkem_adapter_test.go @@ -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") + } }