Epurate wg refs

This commit is contained in:
riccardom
2026-08-27 12:18:07 +02:00
parent 4582c5b4fc
commit da6a56a7ef
8 changed files with 124 additions and 124 deletions
+12 -13
View File
@@ -1,19 +1,18 @@
package pqkem package pqkem
// WGCallbackHandler is implemented by the host (the NetBird wiring layer) and // CallbackHandler is implemented by the host and invoked by the library. The
// invoked by the library. The library only reports events; the host owns the // library only reports events; the host owns the reaction. Keeping this an
// WireGuard reaction. Keeping this an interface — rather than calling wgctrl // interface — rather than touching the transport or keying directly — is what lets
// directly — is what lets the KEM code be extracted as a standalone library. // the KEM code be extracted as a standalone library.
type WGCallbackHandler interface { type CallbackHandler interface {
// OnNewPSKReady fires when a fresh post-quantum PSK has been derived for a // OnNewPSKReady fires when a fresh post-quantum PSK has been derived for a peer
// peer and must be programmed on the WireGuard interface. It is invoked at the // and must be programmed into the consumer's secure channel. It is invoked at
// commit point of each side: the initiator on receiving the answer, the // the commit point of each side: the initiator on receiving the answer, the
// responder on receiving the confirm. // responder on receiving the confirm.
OnNewPSKReady(remoteWgKey string, psk PSK) error OnNewPSKReady(remoteID string, psk PSK) error
// OnRekeyFailed fires when an exchange fails to converge within the allotted // OnRekeyFailed fires when an exchange fails to converge within the allotted
// time. The host should tear the peer connection down so ICE re-establishes, // time. The host should tear the peer connection down so it re-establishes, and
// and log a WARN. The library reports the event; it does not dictate the // log a WARN. The library reports the event; it does not dictate the reaction.
// reaction. OnRekeyFailed(remoteID string) error
OnRekeyFailed(remoteWgKey string) error
} }
+41 -41
View File
@@ -8,9 +8,9 @@ import (
// handleOffer (responder) derives and sends the answer for a new exchange, or // handleOffer (responder) derives and sends the answer for a new exchange, or
// resends the cached answer for a duplicate offer (same exchangeID) without // resends the cached answer for a duplicate offer (same exchangeID) without
// re-deriving. The responder is purely reactive: no retransmit loop, no deadline. // re-deriving. The responder is purely reactive: no retransmit loop, no deadline.
func (m *Manager) handleOffer(remoteWgKey string, o *OfferMsg) error { func (m *Manager) handleOffer(remoteID string, o *OfferMsg) error {
m.mu.Lock() m.mu.Lock()
if ex := m.exchanges[remoteWgKey]; ex != nil && ex.id == o.ExchangeID { if ex := m.exchanges[remoteID]; ex != nil && ex.id == o.ExchangeID {
state, last := ex.state, ex.lastSent state, last := ex.state, ex.lastSent
m.mu.Unlock() m.mu.Unlock()
if state == stateReserved { if state == stateReserved {
@@ -18,13 +18,13 @@ func (m *Manager) handleOffer(remoteWgKey string, o *OfferMsg) error {
// dropping avoids a second (randomized -> divergent) derivation. // dropping avoids a second (randomized -> divergent) derivation.
return nil return nil
} }
return m.transport.Send(remoteWgKey, last) return m.transport.Send(remoteID, last)
} }
// Reserve the slot so a concurrent duplicate offer bails. // Reserve the slot so a concurrent duplicate offer bails.
m.exchanges[remoteWgKey] = &exchangeCtl{id: o.ExchangeID, state: stateReserved, startedAt: time.Now()} m.exchanges[remoteID] = &exchangeCtl{id: o.ExchangeID, state: stateReserved, startedAt: time.Now()}
m.mu.Unlock() m.mu.Unlock()
answerBytes, psk, err := Respond(o.KEMOffer, m.binding(remoteWgKey)) answerBytes, psk, err := Respond(o.KEMOffer, m.binding(remoteID))
if err != nil { if err != nil {
return err return err
} }
@@ -34,22 +34,22 @@ func (m *Manager) handleOffer(remoteWgKey string, o *OfferMsg) error {
} }
m.mu.Lock() m.mu.Lock()
if ex := m.exchanges[remoteWgKey]; ex != nil && ex.id == o.ExchangeID { if ex := m.exchanges[remoteID]; ex != nil && ex.id == o.ExchangeID {
ex.state = stateAwaitingConfirm ex.state = stateAwaitingConfirm
ex.lastSent = raw ex.lastSent = raw
ex.pendingPSK = psk ex.pendingPSK = psk
} }
m.mu.Unlock() m.mu.Unlock()
return m.transport.Send(remoteWgKey, raw) return m.transport.Send(remoteID, raw)
} }
// handleAnswer (initiator) derives the PSK, surfaces it, and sends the confirm, then // handleAnswer (initiator) derives the PSK, surfaces it, and sends the confirm, then
// switches the retransmit payload to the confirm. Only valid in stateAwaitingAnswer; // switches the retransmit payload to the confirm. Only valid in stateAwaitingAnswer;
// advancing the state under the lock makes a concurrent/duplicate answer bail. // advancing the state under the lock makes a concurrent/duplicate answer bail.
func (m *Manager) handleAnswer(remoteWgKey string, a *AnswerMsg) error { func (m *Manager) handleAnswer(remoteID string, a *AnswerMsg) error {
m.mu.Lock() m.mu.Lock()
ex := m.exchanges[remoteWgKey] ex := m.exchanges[remoteID]
if ex == nil || ex.id != a.ExchangeID || ex.state != stateAwaitingAnswer { if ex == nil || ex.id != a.ExchangeID || ex.state != stateAwaitingAnswer {
m.mu.Unlock() m.mu.Unlock()
return nil return nil
@@ -59,11 +59,11 @@ func (m *Manager) handleAnswer(remoteWgKey string, a *AnswerMsg) error {
ex.initiator = nil ex.initiator = nil
m.mu.Unlock() m.mu.Unlock()
psk, err := init.Finish(a.KEMAnswer, m.binding(remoteWgKey)) psk, err := init.Finish(a.KEMAnswer, m.binding(remoteID))
if err != nil { if err != nil {
return err return err
} }
if err := m.wg.OnNewPSKReady(remoteWgKey, psk); err != nil { if err := m.cbHandler.OnNewPSKReady(remoteID, psk); err != nil {
return err return err
} }
raw, err := (&ConfirmMsg{ExchangeID: a.ExchangeID}).Encode() raw, err := (&ConfirmMsg{ExchangeID: a.ExchangeID}).Encode()
@@ -72,42 +72,42 @@ func (m *Manager) handleAnswer(remoteWgKey string, a *AnswerMsg) error {
} }
m.mu.Lock() m.mu.Lock()
if ex := m.exchanges[remoteWgKey]; ex != nil && ex.id == a.ExchangeID { if ex := m.exchanges[remoteID]; ex != nil && ex.id == a.ExchangeID {
ex.lastSent = raw // loop now retransmits the confirm ex.lastSent = raw // loop now retransmits the confirm
m.established[remoteWgKey] = true m.established[remoteID] = true
m.failures[remoteWgKey] = 0 m.failures[remoteID] = 0
_ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step) _ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step)
} }
m.mu.Unlock() m.mu.Unlock()
return m.transport.Send(remoteWgKey, raw) return m.transport.Send(remoteID, raw)
} }
// handleConfirm (responder) commits the pending PSK. Only valid in // handleConfirm (responder) commits the pending PSK. Only valid in
// stateAwaitingConfirm; duplicate confirms (the initiator best-effort resends it) // stateAwaitingConfirm; duplicate confirms (the initiator best-effort resends it)
// find the exchange gone and are ignored. // find the exchange gone and are ignored.
func (m *Manager) handleConfirm(remoteWgKey string, c *ConfirmMsg) error { func (m *Manager) handleConfirm(remoteID string, c *ConfirmMsg) error {
m.mu.Lock() m.mu.Lock()
ex := m.exchanges[remoteWgKey] ex := m.exchanges[remoteID]
if ex == nil || ex.id != c.ExchangeID || ex.state != stateAwaitingConfirm { if ex == nil || ex.id != c.ExchangeID || ex.state != stateAwaitingConfirm {
m.mu.Unlock() m.mu.Unlock()
return nil return nil
} }
psk := ex.pendingPSK psk := ex.pendingPSK
delete(m.exchanges, remoteWgKey) delete(m.exchanges, remoteID)
m.established[remoteWgKey] = true m.established[remoteID] = true
m.failures[remoteWgKey] = 0 m.failures[remoteID] = 0
_ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step) _ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step)
m.mu.Unlock() m.mu.Unlock()
return m.wg.OnNewPSKReady(remoteWgKey, psk) return m.cbHandler.OnNewPSKReady(remoteID, psk)
} }
// initiatorLoop retransmits the initiator's outstanding message, keyed off the // initiatorLoop retransmits the initiator's outstanding message, keyed off the
// exchange state: the offer while awaiting the answer (bounded by maxRetries -> // exchange state: the offer while awaiting the answer (bounded by maxRetries ->
// failure), then the confirm a few best-effort times before stopping. The // failure), then the confirm a few best-effort times before stopping. The
// convergence deadline is thus derived from maxRetries * retryInterval. // convergence deadline is thus derived from maxRetries * retryInterval.
func (m *Manager) initiatorLoop(ctx context.Context, remoteWgKey string, id ExchangeID) { func (m *Manager) initiatorLoop(ctx context.Context, remoteID string, id ExchangeID) {
defer m.wait.Done() defer m.wait.Done()
t := time.NewTicker(m.retryInterval) t := time.NewTicker(m.retryInterval)
defer t.Stop() defer t.Stop()
@@ -119,7 +119,7 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteWgKey string, id Exch
return return
case <-t.C: case <-t.C:
m.mu.Lock() m.mu.Lock()
ex := m.exchanges[remoteWgKey] ex := m.exchanges[remoteID]
if ex == nil || ex.id != id { if ex == nil || ex.id != id {
m.mu.Unlock() m.mu.Unlock()
return return
@@ -128,27 +128,27 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteWgKey string, id Exch
switch ex.state { switch ex.state {
case stateAwaitingAnswer: case stateAwaitingAnswer:
if offerAttempts >= m.maxRetries { if offerAttempts >= m.maxRetries {
delete(m.exchanges, remoteWgKey) delete(m.exchanges, remoteID)
fail := m.registerFailureLocked(remoteWgKey) fail := m.registerFailureLocked(remoteID)
m.mu.Unlock() m.mu.Unlock()
m.raiseFailure(remoteWgKey, fail) m.raiseFailure(remoteID, fail)
return return
} }
msg := ex.lastSent msg := ex.lastSent
offerAttempts++ offerAttempts++
m.mu.Unlock() m.mu.Unlock()
m.retransmit(remoteWgKey, msg) m.retransmit(remoteID, msg)
case stateConfirming: case stateConfirming:
if confirmsSent >= confirmRetransmits { if confirmsSent >= confirmRetransmits {
delete(m.exchanges, remoteWgKey) delete(m.exchanges, remoteID)
m.mu.Unlock() m.mu.Unlock()
return return
} }
msg := ex.lastSent msg := ex.lastSent
confirmsSent++ confirmsSent++
m.mu.Unlock() m.mu.Unlock()
m.retransmit(remoteWgKey, msg) m.retransmit(remoteID, msg)
default: default:
m.mu.Unlock() m.mu.Unlock()
@@ -162,30 +162,30 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteWgKey string, id Exch
// an initial exchange (peer never established) fails immediately; a rekey tolerates // an initial exchange (peer never established) fails immediately; a rekey tolerates
// up to maxRekeyFailures consecutive misses (we stay on the still-valid previous // up to maxRekeyFailures consecutive misses (we stay on the still-valid previous
// PSK) before failing. Assumes m.mu is held. // PSK) before failing. Assumes m.mu is held.
func (m *Manager) registerFailureLocked(remoteWgKey string) bool { func (m *Manager) registerFailureLocked(remoteID string) bool {
if !m.established[remoteWgKey] { if !m.established[remoteID] {
return true return true
} }
m.failures[remoteWgKey]++ m.failures[remoteID]++
if m.failures[remoteWgKey] >= m.maxRekeyFailures { if m.failures[remoteID] >= m.maxRekeyFailures {
m.failures[remoteWgKey] = 0 m.failures[remoteID] = 0
return true return true
} }
return false return false
} }
func (m *Manager) raiseFailure(remoteWgKey string, fail bool) { func (m *Manager) raiseFailure(remoteID string, fail bool) {
if !fail { if !fail {
m.logger.Warn("pqkem rekey attempt timed out, will retry next cycle", "peer", remoteWgKey) m.logger.Warn("pqkem rekey attempt timed out, will retry next cycle", "peer", remoteID)
return return
} }
if err := m.wg.OnRekeyFailed(remoteWgKey); err != nil { if err := m.cbHandler.OnRekeyFailed(remoteID); err != nil {
m.logger.Error("pqkem OnRekeyFailed handler error", "peer", remoteWgKey, "err", err) m.logger.Error("pqkem OnRekeyFailed handler error", "peer", remoteID, "err", err)
} }
} }
func (m *Manager) retransmit(remoteWgKey string, msg []byte) { func (m *Manager) retransmit(remoteID string, msg []byte) {
if err := m.transport.Send(remoteWgKey, msg); err != nil { if err := m.transport.Send(remoteID, msg); err != nil {
m.logger.Warn("pqkem retransmit failed", "peer", remoteWgKey, "err", err) m.logger.Warn("pqkem retransmit failed", "peer", remoteID, "err", err)
} }
} }
+7 -7
View File
@@ -1,6 +1,6 @@
// Package pqkem is a spike (NET-1406) for a post-quantum pre-shared-key exchange // Package pqkem is a spike (NET-1406) for a post-quantum pre-shared-key exchange
// that could replace Rosenpass. It performs an X25519MLKEM768 hybrid key // that could replace Rosenpass. It performs an X25519MLKEM768 hybrid key
// encapsulation and derives a 32-byte WireGuard PSK. // encapsulation and derives a 32-byte pre-shared key (PSK).
// //
// The exchange is a single round trip designed to ride the (already // The exchange is a single round trip designed to ride the (already
// authenticated) Signal offer/answer channel: // authenticated) Signal offer/answer channel:
@@ -9,7 +9,7 @@
// initiator <--Answer(1120B)-- responder // initiator <--Answer(1120B)-- responder
// //
// Both sides then hold the same PSK, which is bound to the two peers' identities // Both sides then hold the same PSK, which is bound to the two peers' identities
// (their WireGuard static public keys) so the derived key cannot be transplanted // (their peer identity keys) so the derived key cannot be transplanted
// to a different peer pair even if the transport authentication were bypassed. // to a different peer pair even if the transport authentication were bypassed.
// //
// Combiner note: this follows the IETF hybrid layout (X25519 ‖ ML-KEM on the // Combiner note: this follows the IETF hybrid layout (X25519 ‖ ML-KEM on the
@@ -35,14 +35,14 @@ const (
pskLabel = "netbird-pq-psk-v1" pskLabel = "netbird-pq-psk-v1"
) )
// PSK is the 32-byte pre-shared key handed to WireGuard. // PSK is the 32-byte derived pre-shared key handed to the consumer to key its channel.
type PSK [32]byte type PSK [32]byte
// Binding identifies the peer pair the PSK is derived for. Callers set both // Binding identifies the peer pair the PSK is derived for. Callers set both
// WireGuard static public keys; the order does not matter (it is canonicalised). // peer identity keys; the order does not matter (it is canonicalised).
type Binding struct { type Binding struct {
LocalWgPub []byte LocalID []byte
RemoteWgPub []byte RemoteID []byte
} }
// Initiator holds the ephemeral secrets between Offer and Finish. // Initiator holds the ephemeral secrets between Offer and Finish.
@@ -141,7 +141,7 @@ func Respond(offer []byte, b Binding) (answer []byte, psk PSK, err error) {
// TODO(NET-1406): replace the SHA-256 concat with the RFC HKDF combiner // TODO(NET-1406): replace the SHA-256 concat with the RFC HKDF combiner
// (crypto/hkdf, Go 1.24+) and proper labels before this leaves spike status. // (crypto/hkdf, Go 1.24+) and proper labels before this leaves spike status.
func derivePSK(ssMLKEM, ssX, offer, answer []byte, b Binding) PSK { func derivePSK(ssMLKEM, ssX, offer, answer []byte, b Binding) PSK {
lo, hi := canonicalPair(b.LocalWgPub, b.RemoteWgPub) lo, hi := canonicalPair(b.LocalID, b.RemoteID)
h := sha256.New() h := sha256.New()
h.Write([]byte(pskLabel)) h.Write([]byte(pskLabel))
+6 -6
View File
@@ -18,11 +18,11 @@ func TestExchange_DerivesMatchingPSK(t *testing.T) {
require.Len(t, init.Offer(), OfferSize) require.Len(t, init.Offer(), OfferSize)
answer, pskB, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA}) answer, pskB, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA})
require.NoError(t, err) require.NoError(t, err)
require.Len(t, answer, AnswerSize) require.Len(t, answer, AnswerSize)
pskA, err := init.Finish(answer, Binding{LocalWgPub: wgA, RemoteWgPub: wgB}) pskA, err := init.Finish(answer, Binding{LocalID: wgA, RemoteID: wgB})
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, pskB, pskA, "both sides must derive the same PSK") require.Equal(t, pskB, pskA, "both sides must derive the same PSK")
@@ -34,13 +34,13 @@ func TestExchange_PSKBoundToPeerIdentities(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// responder computes with the honest pair... // responder computes with the honest pair...
_, pskHonest, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA}) _, pskHonest, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA})
require.NoError(t, err) require.NoError(t, err)
// ...a second responder run with a different peer identity yields a different PSK, // ...a second responder run with a different peer identity yields a different PSK,
// even though the KEM material would otherwise combine identically. // even though the KEM material would otherwise combine identically.
wgC := []byte("peer-C-wireguard-pubkey-32bytes!") wgC := []byte("peer-C-wireguard-pubkey-32bytes!")
_, pskWrong, err := Respond(init.Offer(), Binding{LocalWgPub: wgC, RemoteWgPub: wgA}) _, pskWrong, err := Respond(init.Offer(), Binding{LocalID: wgC, RemoteID: wgA})
require.NoError(t, err) require.NoError(t, err)
require.NotEqual(t, pskHonest, pskWrong, "PSK must be bound to the peer pair") require.NotEqual(t, pskHonest, pskWrong, "PSK must be bound to the peer pair")
@@ -70,12 +70,12 @@ func TestExchange_ReportSizesAndTiming(t *testing.T) {
tInit += time.Since(s0) tInit += time.Since(s0)
s1 := time.Now() s1 := time.Now()
answer, _, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA}) answer, _, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA})
require.NoError(t, err) require.NoError(t, err)
tResp += time.Since(s1) tResp += time.Since(s1)
s2 := time.Now() s2 := time.Now()
_, err = init.Finish(answer, Binding{LocalWgPub: wgA, RemoteWgPub: wgB}) _, err = init.Finish(answer, Binding{LocalID: wgA, RemoteID: wgB})
require.NoError(t, err) require.NoError(t, err)
tFinish += time.Since(s2) tFinish += time.Since(s2)
} }
+48 -47
View File
@@ -10,8 +10,9 @@ import (
) )
const ( const (
// DefaultRekeyInterval matches WireGuard's own REKEY_AFTER_TIME so the freshly // DefaultRekeyInterval is the default PSK rotation cadence (~2 min), chosen so a
// rotated PSK is naturally adopted by WG's next handshake without forcing one. // rotated PSK is adopted by the consumer's next transport handshake without
// forcing one.
DefaultRekeyInterval = 2 * time.Minute DefaultRekeyInterval = 2 * time.Minute
// DefaultRetryInterval is how often the initiator retransmits its outstanding // DefaultRetryInterval is how often the initiator retransmits its outstanding
// message (offer, then confirm) while an exchange is in flight. // message (offer, then confirm) while an exchange is in flight.
@@ -29,11 +30,11 @@ const (
) )
// Transport hands an already-encoded exchange message to the peer. The host routes // Transport hands an already-encoded exchange message to the peer. The host routes
// it over the appropriate channel — Signal before the tunnel is up, the WireGuard // it over the appropriate channel — a signalling channel before the tunnel is up,
// tunnel for rekeys — so the Manager never needs to know which is in use. It is the // the data tunnel for rekeys — so the Manager never needs to know which is in use.
// analogue of go-rosenpass's Conn seam. // It is the analogue of go-rosenpass's Conn seam.
type Transport interface { type Transport interface {
Send(remoteWgKey string, msg []byte) error Send(remoteID string, msg []byte) error
} }
// exchangeState is the single source of truth for an exchange's role and phase. // exchangeState is the single source of truth for an exchange's role and phase.
@@ -66,14 +67,14 @@ type exchangeCtl struct {
// Manager is the stateful orchestrator — the analogue of go-rosenpass's Server. It // Manager is the stateful orchestrator — the analogue of go-rosenpass's Server. It
// runs the per-peer rekey timer, drives the X25519MLKEM768 exchange over a pluggable // runs the per-peer rekey timer, drives the X25519MLKEM768 exchange over a pluggable
// Transport, and surfaces the derived PSK to the host via WGCallbackHandler. The // Transport, and surfaces the derived PSK to the host via CallbackHandler. The
// cryptography is the pure kem.go primitives; all per-exchange and per-peer state // cryptography is the pure kem.go primitives; all per-exchange and per-peer state
// lives here under one lock. // lives here under one lock.
type Manager struct { type Manager struct {
localWgKey string localID string
transport Transport transport Transport
wg WGCallbackHandler cbHandler CallbackHandler
logger *slog.Logger logger *slog.Logger
rekeyInterval time.Duration rekeyInterval time.Duration
retryInterval time.Duration retryInterval time.Duration
@@ -91,11 +92,11 @@ type Manager struct {
wait sync.WaitGroup wait sync.WaitGroup
} }
// NewManager builds a manager for the local peer identified by its WireGuard public // 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 zero // key (used for the deterministic initiator role and the identity binding). A zero
// interval falls back to DefaultRekeyInterval; a nil logger to slog.Default(). // interval falls back to DefaultRekeyInterval; a nil logger to slog.Default().
// Retry/retries/K use their defaults and can be overridden before use. // Retry/retries/K use their defaults and can be overridden before use.
func NewManager(localWgKey string, t Transport, h WGCallbackHandler, interval time.Duration, logger *slog.Logger) *Manager { func NewManager(localID string, t Transport, h CallbackHandler, interval time.Duration, logger *slog.Logger) *Manager {
if interval <= 0 { if interval <= 0 {
interval = DefaultRekeyInterval interval = DefaultRekeyInterval
} }
@@ -104,9 +105,9 @@ func NewManager(localWgKey string, t Transport, h WGCallbackHandler, interval ti
} }
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
return &Manager{ return &Manager{
localWgKey: localWgKey, localID: localID,
transport: t, transport: t,
wg: h, cbHandler: h,
logger: logger, logger: logger,
rekeyInterval: interval, rekeyInterval: interval,
retryInterval: DefaultRetryInterval, retryInterval: DefaultRetryInterval,
@@ -122,40 +123,40 @@ func NewManager(localWgKey string, t Transport, h WGCallbackHandler, interval ti
} }
// IsInitiator reports whether the local peer drives the exchange for this remote // IsInitiator reports whether the local peer drives the exchange for this remote
// peer. Roles are deterministic (lexicographic WG-key compare) so exactly one side // peer. Roles are deterministic (lexicographic identity-key compare) so exactly one
// initiates, mirroring how Rosenpass picks its handshake initiator. // side initiates, mirroring how Rosenpass picks its handshake initiator.
func (m *Manager) IsInitiator(remoteWgKey string) bool { func (m *Manager) IsInitiator(remoteID string) bool {
return m.localWgKey > remoteWgKey return m.localID > remoteID
} }
// AddPeer registers a remote peer and starts its rekey timer. Re-adding is a no-op. // AddPeer registers a remote peer and starts its rekey timer. Re-adding is a no-op.
func (m *Manager) AddPeer(remoteWgKey string) { func (m *Manager) AddPeer(remoteID string) {
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock() defer m.mu.Unlock()
if _, ok := m.peers[remoteWgKey]; ok { if _, ok := m.peers[remoteID]; ok {
return return
} }
ctx, cancel := context.WithCancel(m.rootCtx) ctx, cancel := context.WithCancel(m.rootCtx)
m.peers[remoteWgKey] = cancel m.peers[remoteID] = cancel
m.wait.Add(1) m.wait.Add(1)
go m.rekeyLoop(ctx, remoteWgKey) go m.rekeyLoop(ctx, remoteID)
} }
// RemovePeer stops a peer's rekey timer and any in-flight exchange, and drops state. // RemovePeer stops a peer's rekey timer and any in-flight exchange, and drops state.
func (m *Manager) RemovePeer(remoteWgKey string) { func (m *Manager) RemovePeer(remoteID string) {
m.mu.Lock() m.mu.Lock()
if cancel, ok := m.peers[remoteWgKey]; ok { if cancel, ok := m.peers[remoteID]; ok {
cancel() cancel()
delete(m.peers, remoteWgKey) delete(m.peers, remoteID)
} }
if ex, ok := m.exchanges[remoteWgKey]; ok { if ex, ok := m.exchanges[remoteID]; ok {
if ex.cancel != nil { if ex.cancel != nil {
ex.cancel() ex.cancel()
} }
delete(m.exchanges, remoteWgKey) delete(m.exchanges, remoteID)
} }
delete(m.established, remoteWgKey) delete(m.established, remoteID)
delete(m.failures, remoteWgKey) delete(m.failures, remoteID)
m.mu.Unlock() m.mu.Unlock()
} }
@@ -171,28 +172,28 @@ func (m *Manager) Stop() {
// HandleInbound decodes an incoming message and drives the exchange, sending any // HandleInbound decodes an incoming message and drives the exchange, sending any
// response via the transport and surfacing derived PSKs / convergence to the host. // response via the transport and surfacing derived PSKs / convergence to the host.
func (m *Manager) HandleInbound(remoteWgKey string, raw []byte) error { func (m *Manager) HandleInbound(remoteID string, raw []byte) error {
typ, msg, err := Decode(raw) typ, msg, err := Decode(raw)
if err != nil { if err != nil {
return fmt.Errorf("decode from %s: %w", remoteWgKey, err) return fmt.Errorf("decode from %s: %w", remoteID, err)
} }
switch typ { switch typ {
case MsgOffer: case MsgOffer:
return m.handleOffer(remoteWgKey, msg.(*OfferMsg)) return m.handleOffer(remoteID, msg.(*OfferMsg))
case MsgAnswer: case MsgAnswer:
return m.handleAnswer(remoteWgKey, msg.(*AnswerMsg)) return m.handleAnswer(remoteID, msg.(*AnswerMsg))
case MsgConfirm: case MsgConfirm:
return m.handleConfirm(remoteWgKey, msg.(*ConfirmMsg)) return m.handleConfirm(remoteID, msg.(*ConfirmMsg))
default: default:
return fmt.Errorf("unhandled message type %d from %s", typ, remoteWgKey) return fmt.Errorf("unhandled message type %d from %s", typ, remoteID)
} }
} }
// initiateRekey starts a fresh exchange when the local peer is the initiator for // initiateRekey starts a fresh exchange when the local peer is the initiator for
// this remote peer; the responder waits for the offer instead. Exposed (unexported // this remote peer; the responder waits for the offer instead. Exposed (unexported
// but directly callable) so tests can drive a rekey without waiting on the ticker. // but directly callable) so tests can drive a rekey without waiting on the ticker.
func (m *Manager) initiateRekey(remoteWgKey string) error { func (m *Manager) initiateRekey(remoteID string) error {
if !m.IsInitiator(remoteWgKey) { if !m.IsInitiator(remoteID) {
return nil return nil
} }
init, err := NewInitiator() init, err := NewInitiator()
@@ -210,10 +211,10 @@ func (m *Manager) initiateRekey(remoteWgKey string) error {
ctx, cancel := context.WithCancel(m.rootCtx) ctx, cancel := context.WithCancel(m.rootCtx)
m.mu.Lock() m.mu.Lock()
if old := m.exchanges[remoteWgKey]; old != nil && old.cancel != nil { if old := m.exchanges[remoteID]; old != nil && old.cancel != nil {
old.cancel() old.cancel()
} }
m.exchanges[remoteWgKey] = &exchangeCtl{ m.exchanges[remoteID] = &exchangeCtl{
id: id, id: id,
state: stateAwaitingAnswer, state: stateAwaitingAnswer,
startedAt: time.Now(), startedAt: time.Now(),
@@ -224,12 +225,12 @@ func (m *Manager) initiateRekey(remoteWgKey string) error {
m.mu.Unlock() m.mu.Unlock()
m.wait.Add(1) m.wait.Add(1)
go m.initiatorLoop(ctx, remoteWgKey, id) go m.initiatorLoop(ctx, remoteID, id)
return m.transport.Send(remoteWgKey, raw) return m.transport.Send(remoteID, raw)
} }
func (m *Manager) rekeyLoop(ctx context.Context, remoteWgKey string) { func (m *Manager) rekeyLoop(ctx context.Context, remoteID string) {
defer m.wait.Done() defer m.wait.Done()
t := time.NewTicker(m.rekeyInterval) t := time.NewTicker(m.rekeyInterval)
defer t.Stop() defer t.Stop()
@@ -238,15 +239,15 @@ func (m *Manager) rekeyLoop(ctx context.Context, remoteWgKey string) {
case <-ctx.Done(): case <-ctx.Done():
return return
case <-t.C: case <-t.C:
if err := m.initiateRekey(remoteWgKey); err != nil { if err := m.initiateRekey(remoteID); err != nil {
m.logger.Error("pqkem rekey failed to start", "peer", remoteWgKey, "err", err) m.logger.Error("pqkem rekey failed to start", "peer", remoteID, "err", err)
} }
} }
} }
} }
func (m *Manager) binding(remoteWgKey string) Binding { func (m *Manager) binding(remoteID string) Binding {
return Binding{LocalWgPub: []byte(m.localWgKey), RemoteWgPub: []byte(remoteWgKey)} return Binding{LocalID: []byte(m.localID), RemoteID: []byte(remoteID)}
} }
func newExchangeID() (ExchangeID, error) { func newExchangeID() (ExchangeID, error) {
+5 -5
View File
@@ -15,7 +15,7 @@ type loopback struct {
peer *Manager peer *Manager
} }
func (l *loopback) Send(remoteWgKey string, msg []byte) error { func (l *loopback) Send(remoteID string, msg []byte) error {
cp := append([]byte(nil), msg...) cp := append([]byte(nil), msg...)
return l.peer.HandleInbound(l.localKey, cp) return l.peer.HandleInbound(l.localKey, cp)
} }
@@ -28,17 +28,17 @@ type fakeWG struct {
func newFakeWG() *fakeWG { return &fakeWG{psks: map[string]PSK{}} } func newFakeWG() *fakeWG { return &fakeWG{psks: map[string]PSK{}} }
func (f *fakeWG) OnNewPSKReady(remoteWgKey string, psk PSK) error { func (f *fakeWG) OnNewPSKReady(remoteID string, psk PSK) error {
f.mu.Lock() f.mu.Lock()
defer f.mu.Unlock() defer f.mu.Unlock()
f.psks[remoteWgKey] = psk f.psks[remoteID] = psk
return nil return nil
} }
func (f *fakeWG) OnRekeyFailed(remoteWgKey string) error { func (f *fakeWG) OnRekeyFailed(remoteID string) error {
f.mu.Lock() f.mu.Lock()
defer f.mu.Unlock() defer f.mu.Unlock()
f.failed = append(f.failed, remoteWgKey) f.failed = append(f.failed, remoteID)
return nil return nil
} }
+4 -4
View File
@@ -7,7 +7,7 @@ import (
// Wire framing for the PQ-KEM exchange. Messages are self-contained, versioned, // Wire framing for the PQ-KEM exchange. Messages are self-contained, versioned,
// transport-agnostic byte blobs: the same bytes ride the Signal offer/answer // transport-agnostic byte blobs: the same bytes ride the Signal offer/answer
// (initial, pre-tunnel) or a WireGuard-tunnel packet (rekey). They are NOT a gRPC // (initial, pre-tunnel) or a data-tunnel packet (rekey). They are NOT a gRPC
// service — the network layer only sees opaque []byte. // service — the network layer only sees opaque []byte.
// //
// Layout (all messages): [type:1][version:1][exchangeID:16][payload...] // Layout (all messages): [type:1][version:1][exchangeID:16][payload...]
@@ -59,9 +59,9 @@ type AnswerMsg struct {
// derive it to encapsulate), so the initiator needs no confirmation. Only the // derive it to encapsulate), so the initiator needs no confirmation. Only the
// responder is left unsure whether the initiator received the answer and committed // responder is left unsure whether the initiator received the answer and committed
// the key — this message resolves that. As a bonus, being sent under the new PSK it // the key — this message resolves that. As a bonus, being sent under the new PSK it
// triggers the WireGuard handshake with the new key: the responder converges on // exercises the consumer's channel handshake with the new key: the responder
// receiving it, and the initiator converges by observing that handshake succeed // converges on receiving it, and the initiator converges by observing that handshake
// (which fails on a PSK mismatch, so success proves the responder also committed). // succeed (which fails on a PSK mismatch, so success proves the responder committed too).
type ConfirmMsg struct { type ConfirmMsg struct {
ExchangeID ExchangeID ExchangeID ExchangeID
} }
+1 -1
View File
@@ -9,7 +9,7 @@ import (
func TestMessageRoundTrip(t *testing.T) { func TestMessageRoundTrip(t *testing.T) {
init, err := NewInitiator() init, err := NewInitiator()
require.NoError(t, err) require.NoError(t, err)
answer, _, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA}) answer, _, err := Respond(init.Offer(), Binding{LocalID: wgB, RemoteID: wgA})
require.NoError(t, err) require.NoError(t, err)
id := ExchangeID{1, 2, 3, 4} id := ExchangeID{1, 2, 3, 4}