Assume two transports: initial "signal" (control plane) one (no data path established yet) + data path one

Define OnDataPathRekeyed event to transition from control plane path to data plane path over the WG tunnel.

Keep confirm ALWAYS on NEW established WG tunnel (posthandshake with rekeying). We keep an active method
irrelevant of the WG handshake (we might decide that the indirect wg handshake is sufficient in the future).

Optimistic commit on responder(when sending answer), while on initiator we set it on getting the answer
This commit is contained in:
riccardom
2026-09-11 14:48:54 +02:00
parent f8bb816dea
commit 2ed0d3bf03
4 changed files with 175 additions and 103 deletions
+57 -51
View File
@@ -5,9 +5,10 @@ import (
"time" "time"
) )
// handleOffer (responder) derives and sends the answer for a new exchange, or // handleOffer (responder) derives the PSK, commits it optimistically, and sends the
// resends the cached answer for a duplicate offer (same exchangeID) without // answer over the current channel. A duplicate offer (same exchangeID) resends the
// re-deriving. The responder is purely reactive: no retransmit loop, no deadline. // cached answer without re-deriving. The responder is otherwise reactive: no
// retransmit loop.
func (m *Manager) handleOffer(remoteID string, o *OfferMsg) error { func (m *Manager) handleOffer(remoteID string, o *OfferMsg) error {
m.mu.Lock() m.mu.Lock()
if ex := m.exchanges[remoteID]; ex != nil && ex.id == o.ExchangeID { if ex := m.exchanges[remoteID]; ex != nil && ex.id == o.ExchangeID {
@@ -18,7 +19,7 @@ func (m *Manager) handleOffer(remoteID 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(remoteID, last) return m.send(remoteID, last)
} }
// Reserve the slot so a concurrent duplicate offer bails. // Reserve the slot so a concurrent duplicate offer bails.
m.exchanges[remoteID] = &exchangeCtl{id: o.ExchangeID, state: stateReserved, startedAt: time.Now()} m.exchanges[remoteID] = &exchangeCtl{id: o.ExchangeID, state: stateReserved, startedAt: time.Now()}
@@ -34,18 +35,27 @@ func (m *Manager) handleOffer(remoteID string, o *OfferMsg) error {
} }
m.mu.Lock() m.mu.Lock()
if ex := m.exchanges[remoteID]; ex != nil && ex.id == o.ExchangeID { ex := m.exchanges[remoteID]
ex.state = stateAwaitingConfirm if ex == nil || ex.id != o.ExchangeID {
ex.lastSent = raw m.mu.Unlock()
ex.pendingPSK = psk return nil
} }
ex.state = stateAwaitingConfirm
ex.lastSent = raw
ex.pendingPSK = psk
m.mu.Unlock() m.mu.Unlock()
return m.transport.Send(remoteID, raw) // Commit optimistically so our data path can rekey to the new PSK; a lost answer
// simply means the data path won't come up (initial) or the previous PSK keeps
// working until it does (rekey grace) — both self-heal via retry.
if err := m.cbHandler.OnNewPSKReady(remoteID, psk); err != nil {
return err
}
return m.send(remoteID, raw)
} }
// handleAnswer (initiator) derives the PSK, surfaces it, and sends the confirm, then // handleAnswer (initiator) derives and commits the PSK, then waits for
// switches the retransmit payload to the confirm. Only valid in stateAwaitingAnswer; // OnDataPathRekeyed to send 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(remoteID string, a *AnswerMsg) error { func (m *Manager) handleAnswer(remoteID string, a *AnswerMsg) error {
m.mu.Lock() m.mu.Lock()
@@ -54,7 +64,7 @@ func (m *Manager) handleAnswer(remoteID string, a *AnswerMsg) error {
m.mu.Unlock() m.mu.Unlock()
return nil return nil
} }
ex.state = stateConfirming ex.state = stateAwaitingRekey
init := ex.initiator init := ex.initiator
ex.initiator = nil ex.initiator = nil
m.mu.Unlock() m.mu.Unlock()
@@ -63,29 +73,15 @@ func (m *Manager) handleAnswer(remoteID string, a *AnswerMsg) error {
if err != nil { if err != nil {
return err return err
} }
if err := m.cbHandler.OnNewPSKReady(remoteID, psk); err != nil { // Commit now; the confirm is deferred until OnDataPathRekeyed so it rides the
return err // data path under the new key.
} return m.cbHandler.OnNewPSKReady(remoteID, psk)
raw, err := (&ConfirmMsg{ExchangeID: a.ExchangeID}).Encode()
if err != nil {
return err
}
m.mu.Lock()
if ex := m.exchanges[remoteID]; ex != nil && ex.id == a.ExchangeID {
ex.lastSent = raw // loop now retransmits the confirm
m.established[remoteID] = true
m.failures[remoteID] = 0
_ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step)
}
m.mu.Unlock()
return m.transport.Send(remoteID, raw)
} }
// handleConfirm (responder) commits the pending PSK. Only valid in // handleConfirm (responder) records convergence. The PSK was already committed in
// stateAwaitingConfirm; duplicate confirms (the initiator best-effort resends it) // handleOffer; a confirm arriving over the data path with a matching exchangeID
// find the exchange gone and are ignored. // proves we operate on the new key from this exchange. Stale/duplicate confirms find
// the exchange gone and are ignored.
func (m *Manager) handleConfirm(remoteID string, c *ConfirmMsg) error { func (m *Manager) handleConfirm(remoteID string, c *ConfirmMsg) error {
m.mu.Lock() m.mu.Lock()
ex := m.exchanges[remoteID] ex := m.exchanges[remoteID]
@@ -93,26 +89,25 @@ func (m *Manager) handleConfirm(remoteID string, c *ConfirmMsg) error {
m.mu.Unlock() m.mu.Unlock()
return nil return nil
} }
psk := ex.pendingPSK
delete(m.exchanges, remoteID) delete(m.exchanges, remoteID)
m.established[remoteID] = true m.established[remoteID] = true
m.failures[remoteID] = 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 nil
return m.cbHandler.OnNewPSKReady(remoteID, psk)
} }
// initiatorLoop retransmits the initiator's outstanding message, keyed off the // initiatorLoop retransmits the initiator's outstanding message and enforces the
// exchange state: the offer while awaiting the answer (bounded by maxRetries -> // convergence deadline, keyed off the exchange state: resend the offer while
// failure), then the confirm a few best-effort times before stopping. The // awaiting the answer; wait (counting toward the deadline) while awaiting the data
// convergence deadline is thus derived from maxRetries * retryInterval. // path rekey; resend the confirm a few best-effort times once sent. Exhausting the
// deadline before reaching the confirm phase is a failure.
func (m *Manager) initiatorLoop(ctx context.Context, remoteID 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()
offerAttempts, confirmsSent := 0, 0 attempts, confirmsSent := 0, 0
for { for {
select { select {
case <-ctx.Done(): case <-ctx.Done():
@@ -127,7 +122,7 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteID string, id Exchang
switch ex.state { switch ex.state {
case stateAwaitingAnswer: case stateAwaitingAnswer:
if offerAttempts >= m.maxRetries { if attempts >= m.maxRetries {
delete(m.exchanges, remoteID) delete(m.exchanges, remoteID)
fail := m.registerFailureLocked(remoteID) fail := m.registerFailureLocked(remoteID)
m.mu.Unlock() m.mu.Unlock()
@@ -135,9 +130,24 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteID string, id Exchang
return return
} }
msg := ex.lastSent msg := ex.lastSent
offerAttempts++ attempts++
m.mu.Unlock()
if err := m.send(remoteID, msg); err != nil {
m.logger.Warn("pqkem offer retransmit failed", "peer", remoteID, "err", err)
}
case stateAwaitingRekey:
// Waiting for OnDataPathRekeyed; nothing to send, but the deadline
// still applies (the data path may never come up with the new key).
if attempts >= m.maxRetries {
delete(m.exchanges, remoteID)
fail := m.registerFailureLocked(remoteID)
m.mu.Unlock()
m.raiseFailure(remoteID, fail)
return
}
attempts++
m.mu.Unlock() m.mu.Unlock()
m.retransmit(remoteID, msg)
case stateConfirming: case stateConfirming:
if confirmsSent >= confirmRetransmits { if confirmsSent >= confirmRetransmits {
@@ -148,7 +158,9 @@ func (m *Manager) initiatorLoop(ctx context.Context, remoteID string, id Exchang
msg := ex.lastSent msg := ex.lastSent
confirmsSent++ confirmsSent++
m.mu.Unlock() m.mu.Unlock()
m.retransmit(remoteID, msg) if err := m.transport.SendDataPath(remoteID, msg); err != nil {
m.logger.Warn("pqkem confirm retransmit failed", "peer", remoteID, "err", err)
}
default: default:
m.mu.Unlock() m.mu.Unlock()
@@ -183,9 +195,3 @@ func (m *Manager) raiseFailure(remoteID string, fail bool) {
m.logger.Error("pqkem OnRekeyFailed handler error", "peer", remoteID, "err", err) m.logger.Error("pqkem OnRekeyFailed handler error", "peer", remoteID, "err", err)
} }
} }
func (m *Manager) retransmit(remoteID string, msg []byte) {
if err := m.transport.Send(remoteID, msg); err != nil {
m.logger.Warn("pqkem retransmit failed", "peer", remoteID, "err", err)
}
}
+17 -10
View File
@@ -10,21 +10,25 @@ import (
type dropTransport struct{} type dropTransport struct{}
func (dropTransport) Send(string, []byte) error { return nil } func (dropTransport) SendDataPath(string, []byte) error { return nil }
func (dropTransport) SendSignal(string, []byte) error { return nil }
// gate is a loopback transport with a switchable drop flag. // gate is a loopback transport with a switchable drop flag.
type gate struct { type gate struct {
local string localID string
peer *Manager peer *Manager
drop atomic.Bool drop atomic.Bool
} }
func (g *gate) Send(remote string, msg []byte) error { func (g *gate) SendDataPath(remoteID string, msg []byte) error { return g.deliver(msg) }
func (g *gate) SendSignal(remoteID string, msg []byte) error { return g.deliver(msg) }
func (g *gate) deliver(msg []byte) error {
if g.drop.Load() { if g.drop.Load() {
return nil return nil
} }
cp := append([]byte(nil), msg...) cp := append([]byte(nil), msg...)
return g.peer.HandleInbound(g.local, cp) return g.peer.HandleInbound(g.localID, cp)
} }
func TestManager_InitialTimeoutFailsImmediately(t *testing.T) { func TestManager_InitialTimeoutFailsImmediately(t *testing.T) {
@@ -46,8 +50,8 @@ func TestManager_InitialTimeoutFailsImmediately(t *testing.T) {
} }
func TestManager_RekeyToleratesKFailures(t *testing.T) { func TestManager_RekeyToleratesKFailures(t *testing.T) {
gA := &gate{local: "aaaa"} gA := &gate{localID: "aaaa"}
gB := &gate{local: "bbbb"} gB := &gate{localID: "bbbb"}
wgA := newFakeWG() wgA := newFakeWG()
wgB := newFakeWG() wgB := newFakeWG()
@@ -62,11 +66,14 @@ func TestManager_RekeyToleratesKFailures(t *testing.T) {
defer dA.Stop() defer dA.Stop()
defer dB.Stop() defer dB.Stop()
// first exchange succeeds -> peer becomes established (subsequent failures are rekeys). // First exchange succeeds -> peer becomes established (subsequent failures are
// rekeys). Drive the data-path-rekeyed event so the confirm converges A.
require.NoError(t, dB.initiateRekey("aaaa")) require.NoError(t, dB.initiateRekey("aaaa"))
dA.OnDataPathRekeyed("bbbb")
dB.OnDataPathRekeyed("aaaa")
require.NotEqual(t, PSK{}, wgB.psk("aaaa")) require.NotEqual(t, PSK{}, wgB.psk("aaaa"))
// now drop B's outbound: rekeys can no longer converge. // Now drop B's outbound: rekeys can no longer converge.
gB.drop.Store(true) gB.drop.Store(true)
// K-1 failures must NOT raise OnRekeyFailed. // K-1 failures must NOT raise OnRekeyFailed.
+78 -28
View File
@@ -17,43 +17,49 @@ const (
// 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.
DefaultRetryInterval = 2 * time.Second DefaultRetryInterval = 2 * time.Second
// DefaultMaxRetries bounds how many times the offer is retransmitted before the // DefaultMaxRetries bounds how many ticks an exchange may run before the offer
// exchange is declared failed. The convergence deadline is thus derived as // or the wait for the data-path rekey is declared failed. The convergence
// MaxRetries * RetryInterval — there is no separate deadline timer. // deadline is thus derived as MaxRetries * RetryInterval — no separate timer.
DefaultMaxRetries = 10 DefaultMaxRetries = 10
// DefaultMaxRekeyFailures is how many consecutive rekey (non-initial) failures // DefaultMaxRekeyFailures is how many consecutive rekey (non-initial) failures
// are tolerated before OnRekeyFailed. The initial exchange fails immediately. // are tolerated before OnRekeyFailed. The initial exchange fails immediately.
DefaultMaxRekeyFailures = 3 DefaultMaxRekeyFailures = 3
// confirmRetransmits is how many times the initiator best-effort resends the // confirmRetransmits is how many times the initiator best-effort resends the
// confirm after converging, to cover its loss without a dedicated goroutine. // confirm over the data path, to cover its loss without a dedicated goroutine.
confirmRetransmits = 3 confirmRetransmits = 3
) )
// Transport hands an already-encoded exchange message to the peer. The host routes // Transport carries exchange messages over the two available channels; the library
// it over the appropriate channel — a signalling channel before the tunnel is up, // picks which. It is the analogue of go-rosenpass's Conn seam.
// the data tunnel for rekeys — so the Manager never needs to know which is in use. //
// It is the analogue of go-rosenpass's Conn seam. // Model: a message rides the channel keyed with the CURRENTLY valid key. Before a
// data path exists (initial bootstrap) that is the out-of-band signalling channel;
// once the data path is up (after OnDataPathRekeyed) offers/answers ride it, and the
// confirm ALWAYS rides the data path (so its arrival proves the new key works).
type Transport interface { type Transport interface {
Send(remoteID string, msg []byte) error // SendDataPath sends over the peer's established data path (e.g. a WireGuard
// tunnel). The consumer routes it accordingly.
SendDataPath(remoteID string, msg []byte) error
// SendSignal sends over the out-of-band signalling channel, handled outside the
// library by the consumer.
SendSignal(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.
// Every handler and the retransmit loop key off it, so no role/phase is re-derived
// from other fields.
type exchangeState uint8 type exchangeState uint8
const ( const (
stateReserved exchangeState = iota // responder: deriving the answer stateReserved exchangeState = iota // responder: deriving the answer
stateAwaitingAnswer // initiator: offer sent, awaiting the answer stateAwaitingAnswer // initiator: offer sent, awaiting the answer
stateAwaitingConfirm // responder: answer sent, awaiting the confirm stateAwaitingRekey // initiator: PSK derived+set, awaiting OnDataPathRekeyed to send the confirm
stateConfirming // initiator: answer in, PSK committed, flushing the confirm stateConfirming // initiator: confirm sent over the data path, best-effort retransmit
stateAwaitingConfirm // responder: answer sent, awaiting the confirm over the data path
) )
// exchangeCtl holds all state for one in-flight exchange with a peer, under the // exchangeCtl holds all state for one in-flight exchange with a peer, under the
// Manager's single lock. state drives every decision. lastSent is the message // Manager's single lock. state drives every decision. lastSent is the current
// currently being (re)transmitted. The crypto payloads live here too: initiator is // retransmit payload (offer, then confirm). initiator is the ephemeral handle used
// the ephemeral handle used at Finish (initiator side); pendingPSK is the derived // at Finish; pendingPSK is the responder's derived key. Only the initiator runs a
// key held until the confirm commits it (responder side). Only the initiator runs a
// retransmit loop, so only it sets cancel. // retransmit loop, so only it sets cancel.
type exchangeCtl struct { type exchangeCtl struct {
id ExchangeID id ExchangeID
@@ -66,10 +72,10 @@ 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 the
// Transport, and surfaces the derived PSK to the host via CallbackHandler. The // Transport, and surfaces the derived PSK and convergence to the host via
// cryptography is the pure kem.go primitives; all per-exchange and per-peer state // CallbackHandler. The cryptography is the pure kem.go primitives; all state lives
// lives here under one lock. // here under one lock.
type Manager struct { type Manager struct {
localID string localID string
transport Transport transport Transport
@@ -89,13 +95,13 @@ type Manager struct {
exchanges map[string]*exchangeCtl // in-flight exchange per peer exchanges map[string]*exchangeCtl // in-flight exchange per peer
established map[string]bool // peer has completed at least one exchange established map[string]bool // peer has completed at least one exchange
failures map[string]int // consecutive rekey failures per peer failures map[string]int // consecutive rekey failures per peer
dataPathUp map[string]bool // peer's data path is up (offers/answers may ride it)
wait sync.WaitGroup wait sync.WaitGroup
} }
// NewManager builds a manager for the local peer identified by its peer identity // 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.
func NewManager(localID string, t Transport, h CallbackHandler, 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
@@ -119,6 +125,7 @@ func NewManager(localID string, t Transport, h CallbackHandler, interval time.Du
exchanges: make(map[string]*exchangeCtl), exchanges: make(map[string]*exchangeCtl),
established: make(map[string]bool), established: make(map[string]bool),
failures: make(map[string]int), failures: make(map[string]int),
dataPathUp: make(map[string]bool),
} }
} }
@@ -157,6 +164,7 @@ func (m *Manager) RemovePeer(remoteID string) {
} }
delete(m.established, remoteID) delete(m.established, remoteID)
delete(m.failures, remoteID) delete(m.failures, remoteID)
delete(m.dataPathUp, remoteID)
m.mu.Unlock() m.mu.Unlock()
} }
@@ -170,8 +178,38 @@ func (m *Manager) Stop() {
m.mu.Unlock() m.mu.Unlock()
} }
// HandleInbound decodes an incoming message and drives the exchange, sending any // OnDataPathRekeyed notifies the library that the peer's data path is up and freshly
// response via the transport and surfacing derived PSKs / convergence to the host. // keyed with the latest PSK (fired on first establishment AND every rekey — the same
// event). It marks the data path usable and, if we are the initiator waiting to
// confirm, sends the confirm over the data path (its arrival proves to the responder
// that we operate on the new key, correlated by exchangeID).
func (m *Manager) OnDataPathRekeyed(remoteID string) {
m.mu.Lock()
m.dataPathUp[remoteID] = true
ex := m.exchanges[remoteID]
if ex == nil || ex.state != stateAwaitingRekey {
m.mu.Unlock()
return
}
confirm, err := (&ConfirmMsg{ExchangeID: ex.id}).Encode()
if err != nil {
m.mu.Unlock()
m.logger.Error("pqkem encode confirm", "peer", remoteID, "err", err)
return
}
ex.state = stateConfirming
ex.lastSent = confirm
m.established[remoteID] = true
m.failures[remoteID] = 0
_ = time.Since(ex.startedAt) // convergence latency (metrics hook, later step)
m.mu.Unlock()
if err := m.transport.SendDataPath(remoteID, confirm); err != nil {
m.logger.Warn("pqkem send confirm failed", "peer", remoteID, "err", err)
}
}
// HandleInbound decodes an incoming message and drives the exchange.
func (m *Manager) HandleInbound(remoteID 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 {
@@ -189,9 +227,9 @@ func (m *Manager) HandleInbound(remoteID string, raw []byte) error {
} }
} }
// 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. The
// this remote peer; the responder waits for the offer instead. Exposed (unexported // offer rides the data path if it is up (rekey) or the signalling channel otherwise
// but directly callable) so tests can drive a rekey without waiting on the ticker. // (initial bootstrap).
func (m *Manager) initiateRekey(remoteID string) error { func (m *Manager) initiateRekey(remoteID string) error {
if !m.IsInitiator(remoteID) { if !m.IsInitiator(remoteID) {
return nil return nil
@@ -227,7 +265,7 @@ func (m *Manager) initiateRekey(remoteID string) error {
m.wait.Add(1) m.wait.Add(1)
go m.initiatorLoop(ctx, remoteID, id) go m.initiatorLoop(ctx, remoteID, id)
return m.transport.Send(remoteID, raw) return m.send(remoteID, raw)
} }
func (m *Manager) rekeyLoop(ctx context.Context, remoteID string) { func (m *Manager) rekeyLoop(ctx context.Context, remoteID string) {
@@ -246,6 +284,18 @@ func (m *Manager) rekeyLoop(ctx context.Context, remoteID string) {
} }
} }
// send routes offer/answer over the data path when it is up, else the signalling
// channel. The confirm never goes through here — it always uses SendDataPath.
func (m *Manager) send(remoteID string, msg []byte) error {
m.mu.Lock()
viaDataPath := m.dataPathUp[remoteID]
m.mu.Unlock()
if viaDataPath {
return m.transport.SendDataPath(remoteID, msg)
}
return m.transport.SendSignal(remoteID, msg)
}
func (m *Manager) binding(remoteID string) Binding { func (m *Manager) binding(remoteID string) Binding {
return Binding{LocalID: []byte(m.localID), RemoteID: []byte(remoteID)} return Binding{LocalID: []byte(m.localID), RemoteID: []byte(remoteID)}
} }
+23 -14
View File
@@ -8,16 +8,20 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// loopback delivers a sent message synchronously to the peer driver's HandleInbound, // loopback delivers a sent message synchronously to the peer manager's HandleInbound,
// attributing it to localKey (the sender). // attributing it to localID (the sender). Both channels deliver the same way — the
// test does not care which physical channel is used.
type loopback struct { type loopback struct {
localKey string localID string
peer *Manager peer *Manager
} }
func (l *loopback) Send(remoteID string, msg []byte) error { func (l *loopback) SendDataPath(remoteID string, msg []byte) error { return l.deliver(msg) }
func (l *loopback) SendSignal(remoteID string, msg []byte) error { return l.deliver(msg) }
func (l *loopback) deliver(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.localID, cp)
} }
type fakeWG struct { type fakeWG struct {
@@ -49,8 +53,8 @@ func (f *fakeWG) psk(peer string) PSK {
} }
func TestManager_ExchangeConverges(t *testing.T) { func TestManager_ExchangeConverges(t *testing.T) {
lbA := &loopback{localKey: "aaaa"} lbA := &loopback{localID: "aaaa"}
lbB := &loopback{localKey: "bbbb"} lbB := &loopback{localID: "bbbb"}
wgA := newFakeWG() wgA := newFakeWG()
wgB := newFakeWG() wgB := newFakeWG()
@@ -65,21 +69,26 @@ func TestManager_ExchangeConverges(t *testing.T) {
defer dA.Stop() defer dA.Stop()
defer dB.Stop() defer dB.Stop()
// B is the initiator ("bbbb" > "aaaa"). // B is the initiator ("bbbb" > "aaaa"). Offer/answer flow synchronously over the
// loopback; both commit their PSK, and B parks in stateAwaitingRekey.
require.NoError(t, dB.initiateRekey("aaaa")) require.NoError(t, dB.initiateRekey("aaaa"))
pskB := wgB.psk("aaaa") // B committed on the answer // The consumer reports the data path is (re)keyed on both sides; this makes B
pskA := wgA.psk("bbbb") // A committed on the confirm // send the confirm, which converges A.
dA.OnDataPathRekeyed("bbbb")
dB.OnDataPathRekeyed("aaaa")
pskB := wgB.psk("aaaa")
pskA := wgA.psk("bbbb")
require.NotEqual(t, PSK{}, pskA, "responder A must have a PSK") require.NotEqual(t, PSK{}, pskA, "responder A must have a PSK")
require.NotEqual(t, PSK{}, pskB, "initiator B must have a PSK") require.NotEqual(t, PSK{}, pskB, "initiator B must have a PSK")
require.Equal(t, pskB, pskA, "both sides converge on the same PSK") require.Equal(t, pskB, pskA, "both sides converge on the same PSK")
} }
func TestManager_NonInitiatorDoesNothing(t *testing.T) { func TestManager_NonInitiatorDoesNothing(t *testing.T) {
lbA := &loopback{localKey: "aaaa"} lbA := &loopback{localID: "aaaa"}
wgA := newFakeWG() wgA := newFakeWG()
dA := NewManager("aaaa", lbA, wgA, time.Hour, nil) dA := NewManager("aaaa", lbA, wgA, time.Hour, nil)
// no peer driver wired; if A wrongly initiated, Send would nil-panic.
dA.AddPeer("bbbb") dA.AddPeer("bbbb")
defer dA.Stop() defer dA.Stop()
@@ -89,7 +98,7 @@ func TestManager_NonInitiatorDoesNothing(t *testing.T) {
} }
func TestManager_StopIsIdempotent(t *testing.T) { func TestManager_StopIsIdempotent(t *testing.T) {
dA := NewManager("aaaa", &loopback{localKey: "aaaa"}, newFakeWG(), time.Hour, nil) dA := NewManager("aaaa", &loopback{localID: "aaaa"}, newFakeWG(), time.Hour, nil)
dA.AddPeer("bbbb") dA.AddPeer("bbbb")
dA.Stop() dA.Stop()
dA.Stop() // must not panic or hang dA.Stop() // must not panic or hang