Makes Transport just a UDP socket.

Manager owns maps for remoteID <-> remote UDP addr
Engine talks to manager only
This commit is contained in:
riccardom
2026-08-07 10:24:01 +02:00
parent 5d8d87c050
commit 11d0ad13bb
5 changed files with 244 additions and 218 deletions
+4 -10
View File
@@ -199,9 +199,8 @@ type Engine struct {
rpManager *rosenpass.Manager rpManager *rosenpass.Manager
// pqkemManager runs the ML-KEM post-quantum PSK exchange (gated by NB_ENABLE_PQ_MLKEM). // pqkemManager runs the ML-KEM post-quantum PSK exchange (gated by NB_ENABLE_PQ_MLKEM).
// It owns the data-path transport and peer endpoint routing.
pqkemManager *pqkem.Manager pqkemManager *pqkem.Manager
// pqTransport is the ML-KEM data-path UDP transport, bound on the WG overlay IP.
pqTransport *pqTransport
// syncMsgMux is used to guarantee sequential Management Service message processing // syncMsgMux is used to guarantee sequential Management Service message processing
syncMsgMux *sync.Mutex syncMsgMux *sync.Mutex
@@ -657,11 +656,9 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
if pqErr != nil { if pqErr != nil {
log.Errorf("ML-KEM PQ transport bind failed, PQ exchange disabled: %v", pqErr) log.Errorf("ML-KEM PQ transport bind failed, PQ exchange disabled: %v", pqErr)
} else { } else {
e.pqTransport = tr e.pqkemManager = pqkem.NewManager(publicKey.String(), pqCallbackHandler{wg: e.wgInterface}, nil)
e.pqkemManager = pqkem.NewManager(publicKey.String(), tr, pqCallbackHandler{wg: e.wgInterface}, nil) e.pqkemManager.SetTransport(tr)
tr.setManager(e.pqkemManager) log.Infof("ML-KEM post-quantum exchange enabled (udp port %d on overlay %s)", e.pqkemManager.LocalPort(), e.config.WgAddr.IP)
go tr.run()
log.Infof("ML-KEM post-quantum exchange enabled (udp port %d on overlay %s)", tr.Port(), e.config.WgAddr.IP)
} }
} }
@@ -2097,9 +2094,6 @@ func (e *Engine) close() {
_ = e.rpManager.Close() _ = e.rpManager.Close()
} }
if e.pqTransport != nil {
_ = e.pqTransport.Close()
}
if e.pqkemManager != nil { if e.pqkemManager != nil {
e.pqkemManager.Stop() e.pqkemManager.Stop()
} }
+24 -50
View File
@@ -1,36 +1,31 @@
package pqkem package pqkem
import ( import (
"sync/atomic" "net/netip"
"testing" "testing"
"time" "time"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// dropTransport is a pqkem.Transport that silently discards everything.
type dropTransport struct{} type dropTransport struct{}
func (dropTransport) SendDataPath(string, []byte) error { return nil } func (dropTransport) Send(netip.AddrPort, []byte) error { return nil }
func (dropTransport) LocalPort() int { return 0 }
func (dropTransport) Run(func(netip.AddrPort, []byte)) {}
func (dropTransport) Close() error { return nil }
// gate is a data-path loopback with a switchable drop flag. When dropping it reports func failedCount(f *fakeWG) int {
// success but does not deliver (mimics a lossy/broken tunnel). f.mu.Lock()
type gate struct { defer f.mu.Unlock()
localID string return len(f.failed)
peer *Manager
drop atomic.Bool
}
func (g *gate) SendDataPath(remoteID string, msg []byte) error {
if g.drop.Load() {
return nil
}
cp := append([]byte(nil), msg...)
return g.peer.OnDataPathMessage(g.localID, cp)
} }
func TestManager_InitialTimeoutFailsImmediately(t *testing.T) { func TestManager_InitialTimeoutFailsImmediately(t *testing.T) {
wg := newFakeWG() wg := newFakeWG()
d := NewManager("bbbb", dropTransport{}, wg, nil) // bbbb > aaaa -> initiator d := NewManager("bbbb", wg, nil) // bbbb > aaaa -> initiator
d.SetTransport(dropTransport{})
d.retryInterval = 5 * time.Millisecond d.retryInterval = 5 * time.Millisecond
d.maxRetries = 3 d.maxRetries = 3
defer d.Stop() defer d.Stop()
@@ -41,40 +36,25 @@ func TestManager_InitialTimeoutFailsImmediately(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
require.NotNil(t, offer) require.NotNil(t, offer)
require.Eventually(t, func() bool { require.Eventually(t, func() bool { return failedCount(wg) == 1 }, time.Second, 5*time.Millisecond)
wg.mu.Lock()
defer wg.mu.Unlock()
return len(wg.failed) == 1
}, time.Second, 5*time.Millisecond)
} }
func TestManager_RekeyToleratesKFailures(t *testing.T) { func TestManager_RekeyToleratesKFailures(t *testing.T) {
gA := &gate{localID: "aaaa"} dA, dB, _, wgB, lbB := pair(t)
gB := &gate{localID: "bbbb"}
wgA := newFakeWG()
wgB := newFakeWG()
dA := NewManager("aaaa", gA, wgA, nil)
dB := NewManager("bbbb", gB, wgB, nil)
gA.peer = dB
gB.peer = dA
dB.retryInterval = 5 * time.Millisecond
dB.maxRetries = 2
defer dA.Stop() defer dA.Stop()
defer dB.Stop() defer dB.Stop()
// Bootstrap over signalling -> B becomes established. // Establish: bootstrap + data-path-rekeyed so B becomes established and its data
offer, err := dB.SignalOffer("aaaa") // path is usable.
require.NoError(t, err) bootstrap(t, dA, dB)
answer, err := dA.SignalOnOffer("bbbb", offer)
require.NoError(t, err)
require.NoError(t, dB.SignalOnAnswer("aaaa", answer))
require.NotEqual(t, PSK{}, wgB.psk("aaaa"))
// Bring the data path up on both, then drop B's delivery so rekeys can't converge.
dA.OnDataPathRekeyed("bbbb") dA.OnDataPathRekeyed("bbbb")
dB.OnDataPathRekeyed("aaaa") dB.OnDataPathRekeyed("aaaa")
gB.drop.Store(true) require.NotEqual(t, PSK{}, wgB.psk("aaaa"))
// Tighten timings and drop B's outbound so rekeys can no longer converge.
dB.retryInterval = 5 * time.Millisecond
dB.maxRetries = 2
lbB.drop.Store(true)
// K-1 data-path rekeys must NOT raise OnRekeyFailed. // K-1 data-path rekeys must NOT raise OnRekeyFailed.
for i := 0; i < DefaultMaxRekeyFailures-1; i++ { for i := 0; i < DefaultMaxRekeyFailures-1; i++ {
@@ -85,13 +65,7 @@ func TestManager_RekeyToleratesKFailures(t *testing.T) {
require.Equal(t, 0, failedCount(wgB), "no failure before K attempts") require.Equal(t, 0, failedCount(wgB), "no failure before K attempts")
// The K-th failure raises it once. // The K-th failure raises it once.
_, err = dB.startExchange("aaaa", false, ExchangeID{}) _, err := dB.startExchange("aaaa", false, ExchangeID{})
require.NoError(t, err) require.NoError(t, err)
require.Eventually(t, func() bool { return failedCount(wgB) == 1 }, time.Second, 5*time.Millisecond) require.Eventually(t, func() bool { return failedCount(wgB) == 1 }, time.Second, 5*time.Millisecond)
} }
func failedCount(f *fakeWG) int {
f.mu.Lock()
defer f.mu.Unlock()
return len(f.failed)
}
+117 -46
View File
@@ -5,6 +5,7 @@ import (
"crypto/rand" "crypto/rand"
"fmt" "fmt"
"log/slog" "log/slog"
"net/netip"
"sync" "sync"
"time" "time"
) )
@@ -21,12 +22,22 @@ const (
DefaultMaxRekeyFailures = 3 DefaultMaxRekeyFailures = 3
) )
// Transport pushes a message over the peer's data path (e.g. a WireGuard tunnel). // Transport is the data-path socket the Manager drives (the analogue of
// It is the library's only outbound send: the control-plane (signalling) channel is // go-rosenpass's Conn). It is a dumb mover of bytes to/from endpoints: the Manager
// host-driven — the library hands the host offer/answer payloads to piggyback on the // owns the remoteID<->endpoint routing and hands the transport a resolved endpoint
// host's own negotiation, it never pushes there itself. // to Send, and reverse-resolves the source of each inbound datagram. Its lifecycle
// belongs to the Manager (Run at SetTransport, Close at Stop).
type Transport interface { type Transport interface {
SendDataPath(remoteID string, msg []byte) error // Send delivers msg to the given data-path endpoint.
Send(endpoint netip.AddrPort, msg []byte) error
// LocalPort is the bound local UDP port, announced to peers so they know where
// to send data-path messages.
LocalPort() int
// Run starts delivering inbound datagrams as (source endpoint, msg) to onInbound
// and returns immediately; it runs until Close.
Run(onInbound func(src netip.AddrPort, msg []byte))
// Close stops delivery and releases the socket.
Close() 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.
@@ -58,14 +69,13 @@ 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
// drives the X25519MLKEM768 exchange and surfaces the derived PSK and convergence to // drives the X25519MLKEM768 exchange, owns the peer endpoint routing and the data-path
// the host via CallbackHandler. It is fully event-driven: the bootstrap is triggered // transport, and surfaces the derived PSK and convergence to the host via
// by the host (SignalOffer) and each rotation is clocked by OnDataPathRekeyed (the // CallbackHandler. It is event-driven: the bootstrap is triggered by the host
// consumer's transport rekey). The cryptography is the pure kem.go primitives; all // (SignalOffer) and each rotation is clocked by OnDataPathRekeyed. The cryptography is
// state lives here under one lock. // the pure kem.go primitives; all state lives here under one lock.
type Manager struct { type Manager struct {
localID string localID string
transport Transport
cbHandler CallbackHandler cbHandler CallbackHandler
logger *slog.Logger logger *slog.Logger
@@ -77,27 +87,25 @@ type Manager struct {
rootCancel context.CancelFunc rootCancel context.CancelFunc
mu sync.Mutex mu sync.Mutex
exchanges map[string]*exchangeCtl // in-flight exchange per peer transport Transport
established map[string]bool // peer has completed at least one exchange exchanges map[string]*exchangeCtl // in-flight exchange per peer
failures map[string]int // consecutive rekey failures per peer established map[string]bool // peer has completed at least one exchange
// dataSend holds the peer's data-path sender when the data path is up; nil failures map[string]int // consecutive rekey failures per peer
// (absent) means it is down. Toggled by OnDataPathRekeyed / OnDataPathDown. peers map[string]netip.AddrPort // remoteID -> data-path endpoint (send routing)
dataSend map[string]func(string, []byte) error peersByAddr map[netip.AddrPort]string // reverse: source endpoint -> remoteID (inbound)
wait sync.WaitGroup wait sync.WaitGroup
} }
// NewManager builds a manager for the local peer identified by its peer identity key // 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 // (used for the deterministic initiator role and the identity binding). A nil logger
// falls back to slog.Default(). Retry/retries/K use their defaults and can be // falls back to slog.Default(). Set the data-path transport with SetTransport.
// overridden before use. func NewManager(localID string, h CallbackHandler, logger *slog.Logger) *Manager {
func NewManager(localID string, t Transport, h CallbackHandler, logger *slog.Logger) *Manager {
if logger == nil { if logger == nil {
logger = slog.Default() logger = slog.Default()
} }
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
return &Manager{ return &Manager{
localID: localID, localID: localID,
transport: t,
cbHandler: h, cbHandler: h,
logger: logger, logger: logger,
retryInterval: DefaultRetryInterval, retryInterval: DefaultRetryInterval,
@@ -108,10 +116,34 @@ func NewManager(localID string, t Transport, h CallbackHandler, logger *slog.Log
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),
dataSend: make(map[string]func(string, []byte) error), peers: make(map[string]netip.AddrPort),
peersByAddr: make(map[netip.AddrPort]string),
} }
} }
// SetTransport installs the data-path transport and starts its inbound delivery. The
// Manager owns it from here: Stop closes it.
func (m *Manager) SetTransport(t Transport) {
m.mu.Lock()
m.transport = t
m.mu.Unlock()
if t != nil {
t.Run(m.onDataPathInbound)
}
}
// LocalPort is the data-path transport's bound UDP port (0 if no transport), to be
// announced to peers.
func (m *Manager) LocalPort() int {
m.mu.Lock()
t := m.transport
m.mu.Unlock()
if t == nil {
return 0
}
return t.LocalPort()
}
// 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 identity-key compare) so exactly one // peer. Roles are deterministic (lexicographic identity-key compare) so exactly one
// side initiates, mirroring how Rosenpass picks its handshake initiator. // side initiates, mirroring how Rosenpass picks its handshake initiator.
@@ -119,7 +151,22 @@ func (m *Manager) IsInitiator(remoteID string) bool {
return m.localID > remoteID return m.localID > remoteID
} }
// RemovePeer stops any in-flight exchange for a peer and drops its state. // AddPeer registers where a peer's data-path messages are sent and received: its
// overlay endpoint (IP:port). Re-adding updates the endpoint.
func (m *Manager) AddPeer(remoteID string, endpoint netip.AddrPort) {
if !endpoint.IsValid() {
return
}
m.mu.Lock()
if old, ok := m.peers[remoteID]; ok {
delete(m.peersByAddr, old)
}
m.peers[remoteID] = endpoint
m.peersByAddr[endpoint] = remoteID
m.mu.Unlock()
}
// RemovePeer stops any in-flight exchange for a peer and drops its state and routing.
func (m *Manager) RemovePeer(remoteID string) { func (m *Manager) RemovePeer(remoteID string) {
m.mu.Lock() m.mu.Lock()
if ex, ok := m.exchanges[remoteID]; ok { if ex, ok := m.exchanges[remoteID]; ok {
@@ -130,17 +177,26 @@ func (m *Manager) RemovePeer(remoteID string) {
} }
delete(m.established, remoteID) delete(m.established, remoteID)
delete(m.failures, remoteID) delete(m.failures, remoteID)
delete(m.dataSend, remoteID) if ep, ok := m.peers[remoteID]; ok {
delete(m.peersByAddr, ep)
delete(m.peers, remoteID)
}
m.mu.Unlock() m.mu.Unlock()
} }
// Stop cancels all in-flight exchanges and waits for their goroutines to exit. // Stop cancels all in-flight exchanges, closes the transport, and waits for the
// exchange goroutines to exit.
func (m *Manager) Stop() { func (m *Manager) Stop() {
m.rootCancel() m.rootCancel()
m.wait.Wait() m.wait.Wait()
m.mu.Lock() m.mu.Lock()
t := m.transport
m.transport = nil
m.exchanges = make(map[string]*exchangeCtl) m.exchanges = make(map[string]*exchangeCtl)
m.mu.Unlock() m.mu.Unlock()
if t != nil {
_ = t.Close()
}
} }
// ---- Signalling channel (host-driven; rides the host's negotiation) ---- // ---- Signalling channel (host-driven; rides the host's negotiation) ----
@@ -190,10 +246,24 @@ func (m *Manager) SignalOnAnswer(remoteID string, answer []byte) error {
return m.processAnswer(remoteID, msg.(*AnswerMsg)) return m.processAnswer(remoteID, msg.(*AnswerMsg))
} }
// ---- Data path (library-driven push) ---- // ---- Data path ----
// OnDataPathMessage feeds a KEM message received over the data path (tunnel) and // onDataPathInbound is the transport's inbound handler: it reverse-resolves the
// pushes any reply back over the data path. // source endpoint to a peer and dispatches. Unknown sources are dropped.
func (m *Manager) onDataPathInbound(src netip.AddrPort, msg []byte) {
m.mu.Lock()
remoteID := m.peersByAddr[src]
m.mu.Unlock()
if remoteID == "" {
return
}
if err := m.OnDataPathMessage(remoteID, msg); err != nil {
m.logger.Debug("pqkem inbound", "peer", remoteID, "err", err)
}
}
// OnDataPathMessage handles a KEM message received over the data path from remoteID
// and pushes any reply back over the data path.
func (m *Manager) OnDataPathMessage(remoteID string, raw []byte) error { func (m *Manager) OnDataPathMessage(remoteID string, raw []byte) error {
typ, msg, err := Decode(raw) typ, msg, err := Decode(raw)
if err != nil { if err != nil {
@@ -217,13 +287,12 @@ func (m *Manager) OnDataPathMessage(remoteID string, raw []byte) error {
} }
// OnDataPathRekeyed notifies that the peer's data path is up and freshly keyed with // OnDataPathRekeyed notifies that the peer's data path is up and freshly keyed with
// the latest PSK (fired on first establishment AND every rekey). It marks the data // the latest PSK (fired on first establishment AND every rekey). If we are the
// path usable and, if we are the initiator that just derived a PSK, chains the next // initiator that just derived a PSK, it chains the next exchange: a fresh offer over
// exchange: a fresh offer over the data path that acknowledges the just-completed one // the data path that acknowledges the just-completed one (its arrival under the new
// (its arrival under the new key proves to the responder that the key works). // key proves to the responder that the key works).
func (m *Manager) OnDataPathRekeyed(remoteID string) { func (m *Manager) OnDataPathRekeyed(remoteID string) {
m.mu.Lock() m.mu.Lock()
m.dataSend[remoteID] = m.transport.SendDataPath
ex := m.exchanges[remoteID] ex := m.exchanges[remoteID]
chain := ex != nil && ex.state == stateAwaitingRekey chain := ex != nil && ex.state == stateAwaitingRekey
var ackID ExchangeID var ackID ExchangeID
@@ -245,25 +314,27 @@ func (m *Manager) OnDataPathRekeyed(remoteID string) {
} }
} }
// OnDataPathDown notifies that the peer's data path went down; rotations pause until // OnDataPathDown notifies that the peer's data path went down. Rotations resume once
// it is up again (the host re-bootstraps over signalling on reconnect). // the host re-bootstraps over signalling on reconnect; in-flight data-path sends will
func (m *Manager) OnDataPathDown(remoteID string) { // simply fail until then. Reserved as an explicit hook.
m.mu.Lock() func (m *Manager) OnDataPathDown(remoteID string) {}
delete(m.dataSend, remoteID)
m.mu.Unlock()
}
// ---- internals ---- // ---- internals ----
// pushDataPath sends over the peer's data path, erroring if it is down. // pushDataPath resolves the peer's endpoint and sends over the data-path transport,
// erroring if the peer is unknown or no transport is set.
func (m *Manager) pushDataPath(remoteID string, msg []byte) error { func (m *Manager) pushDataPath(remoteID string, msg []byte) error {
m.mu.Lock() m.mu.Lock()
send := m.dataSend[remoteID] ep, ok := m.peers[remoteID]
t := m.transport
m.mu.Unlock() m.mu.Unlock()
if send == nil { if !ok {
return fmt.Errorf("no data path for peer %s", remoteID) return fmt.Errorf("no data-path endpoint for peer %s", remoteID)
} }
return send(remoteID, msg) if t == nil {
return fmt.Errorf("no data-path transport")
}
return t.Send(ep, msg)
} }
func (m *Manager) binding(remoteID string) Binding { func (m *Manager) binding(remoteID string) Binding {
+71 -24
View File
@@ -1,25 +1,62 @@
package pqkem package pqkem
import ( import (
"fmt"
"net/netip"
"sync" "sync"
"sync/atomic"
"testing" "testing"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// loopback is a data-path transport: SendDataPath delivers synchronously to the peer // netSwitch is an in-memory UDP fabric: transports register their endpoint and get
// manager's OnDataPathMessage, attributing it to localID (the sender). The signalling // datagrams delivered to their inbound handler.
// channel is driven by the test directly via the SignalX methods. type netSwitch struct {
type loopback struct { mu sync.Mutex
localID string h map[netip.AddrPort]func(netip.AddrPort, []byte)
peer *Manager
} }
func (l *loopback) SendDataPath(remoteID string, msg []byte) error { func newSwitch() *netSwitch {
cp := append([]byte(nil), msg...) return &netSwitch{h: map[netip.AddrPort]func(netip.AddrPort, []byte){}}
return l.peer.OnDataPathMessage(l.localID, cp)
} }
func (s *netSwitch) register(ep netip.AddrPort, fn func(netip.AddrPort, []byte)) {
s.mu.Lock()
s.h[ep] = fn
s.mu.Unlock()
}
func (s *netSwitch) deliver(dst, src netip.AddrPort, msg []byte) error {
s.mu.Lock()
fn := s.h[dst]
s.mu.Unlock()
if fn == nil {
return fmt.Errorf("no route to %s", dst)
}
fn(src, msg)
return nil
}
// loopback is an endpoint-based pqkem.Transport over a netSwitch, with a switchable
// drop flag.
type loopback struct {
ep netip.AddrPort
sw *netSwitch
drop atomic.Bool
}
func (l *loopback) Send(dst netip.AddrPort, msg []byte) error {
if l.drop.Load() {
return nil
}
return l.sw.deliver(dst, l.ep, append([]byte(nil), msg...))
}
func (l *loopback) LocalPort() int { return int(l.ep.Port()) }
func (l *loopback) Run(onInbound func(netip.AddrPort, []byte)) { l.sw.register(l.ep, onInbound) }
func (l *loopback) Close() error { return nil }
type fakeWG struct { type fakeWG struct {
mu sync.Mutex mu sync.Mutex
psks map[string]PSK psks map[string]PSK
@@ -48,18 +85,27 @@ func (f *fakeWG) psk(peer string) PSK {
return f.psks[peer] return f.psks[peer]
} }
// pair builds two wired managers (B is the initiator, "bbbb" > "aaaa"). var (
func pair(t *testing.T) (dA, dB *Manager, wgA, wgB *fakeWG) { epA = netip.MustParseAddrPort("100.64.0.1:51833")
epB = netip.MustParseAddrPort("100.64.0.2:51833")
)
// pair builds two wired managers (B is the initiator, "bbbb" > "aaaa") sharing a
// netSwitch, with each peer's data-path endpoint registered. lbB is B's loopback
// (for toggling drop).
func pair(t *testing.T) (dA, dB *Manager, wgA, wgB *fakeWG, lbB *loopback) {
t.Helper() t.Helper()
lbA := &loopback{localID: "aaaa"} sw := newSwitch()
lbB := &loopback{localID: "bbbb"}
wgA = newFakeWG() wgA = newFakeWG()
wgB = newFakeWG() wgB = newFakeWG()
dA = NewManager("aaaa", lbA, wgA, nil) dA = NewManager("aaaa", wgA, nil)
dB = NewManager("bbbb", lbB, wgB, nil) dB = NewManager("bbbb", wgB, nil)
lbA.peer = dB dA.SetTransport(&loopback{ep: epA, sw: sw})
lbB.peer = dA lbB = &loopback{ep: epB, sw: sw}
return dA, dB, wgA, wgB dB.SetTransport(lbB)
dA.AddPeer("bbbb", epB)
dB.AddPeer("aaaa", epA)
return dA, dB, wgA, wgB, lbB
} }
// bootstrap runs the signalling offer/answer (the test plays the host carrying bytes). // bootstrap runs the signalling offer/answer (the test plays the host carrying bytes).
@@ -75,7 +121,7 @@ func bootstrap(t *testing.T, dA, dB *Manager) {
} }
func TestManager_BootstrapDerivesSamePSK(t *testing.T) { func TestManager_BootstrapDerivesSamePSK(t *testing.T) {
dA, dB, wgA, wgB := pair(t) dA, dB, wgA, wgB, _ := pair(t)
defer dA.Stop() defer dA.Stop()
defer dB.Stop() defer dB.Stop()
@@ -88,15 +134,15 @@ func TestManager_BootstrapDerivesSamePSK(t *testing.T) {
} }
func TestManager_ChainRotatesAndAcks(t *testing.T) { func TestManager_ChainRotatesAndAcks(t *testing.T) {
dA, dB, wgA, wgB := pair(t) dA, dB, wgA, wgB, _ := pair(t)
defer dA.Stop() defer dA.Stop()
defer dB.Stop() defer dB.Stop()
bootstrap(t, dA, dB) bootstrap(t, dA, dB)
psk1 := wgB.psk("aaaa") psk1 := wgB.psk("aaaa")
// Data path up on both sides; B chains the next offer (acking exchange 1) over the // Data path up: B (initiator) chains the next offer over the data path, which
// data path, which rotates both to a fresh PSK and acknowledges A. // rotates both to a fresh PSK and acknowledges A.
dA.OnDataPathRekeyed("bbbb") dA.OnDataPathRekeyed("bbbb")
dB.OnDataPathRekeyed("aaaa") dB.OnDataPathRekeyed("aaaa")
@@ -107,7 +153,7 @@ func TestManager_ChainRotatesAndAcks(t *testing.T) {
} }
func TestManager_NonInitiatorReturnsNoOffer(t *testing.T) { func TestManager_NonInitiatorReturnsNoOffer(t *testing.T) {
dA := NewManager("aaaa", &loopback{localID: "aaaa"}, newFakeWG(), nil) dA := NewManager("aaaa", newFakeWG(), nil)
defer dA.Stop() defer dA.Stop()
offer, err := dA.SignalOffer("bbbb") // not the initiator vs "bbbb" offer, err := dA.SignalOffer("bbbb") // not the initiator vs "bbbb"
@@ -116,7 +162,8 @@ func TestManager_NonInitiatorReturnsNoOffer(t *testing.T) {
} }
func TestManager_StopIsIdempotent(t *testing.T) { func TestManager_StopIsIdempotent(t *testing.T) {
dA := NewManager("aaaa", &loopback{localID: "aaaa"}, newFakeWG(), nil) dA := NewManager("aaaa", newFakeWG(), nil)
dA.SetTransport(&loopback{ep: epA, sw: newSwitch()})
dA.Stop() dA.Stop()
dA.Stop() // must not panic or hang dA.Stop() // must not panic or hang
} }
+28 -88
View File
@@ -4,11 +4,8 @@ import (
"fmt" "fmt"
"net" "net"
"net/netip" "net/netip"
"sync"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/pqkem"
) )
// DefaultPort is the preferred UDP port for the ML-KEM data-path service, bound on // DefaultPort is the preferred UDP port for the ML-KEM data-path service, bound on
@@ -17,25 +14,18 @@ import (
// announces Body.mlkemPort when a collision forced it onto a different port. // announces Body.mlkemPort when a collision forced it onto a different port.
const DefaultPort = 51833 const DefaultPort = 51833
// pqTransport is the ML-KEM data-path transport: a dedicated UDP socket bound on the // pqTransport is the ML-KEM data-path transport: a dumb UDP socket bound on the WG
// WireGuard overlay IP. Rekey messages travel through the tunnel to each peer's // overlay IP. It implements pqkem.Transport — the manager owns the remoteID<->endpoint
// overlay IP and announced pqkem port. It implements pqkem.Transport and feeds // routing and drives this socket's lifecycle (Run / Close).
// inbound datagrams to the manager (set via setManager after construction, since the
// manager is built with this transport).
type pqTransport struct { type pqTransport struct {
conn *net.UDPConn conn *net.UDPConn
port int port int
mu sync.RWMutex
mgr *pqkem.Manager
peers map[string]*net.UDPAddr // remoteID (WG pubkey) -> overlay UDP addr to send to
byAddr map[string]string // source addr string -> remoteID (inbound dispatch)
} }
// newPQTransport binds a UDP socket on the WG overlay IP, preferring DefaultPort and // newPQTransport binds a UDP socket on the WG overlay IP, preferring DefaultPort and
// falling back to an OS-assigned ephemeral port if it is in use. It must be called // falling back to an OS-assigned ephemeral port if it is in use. Call it after the WG
// after the WG interface is up so the overlay IP is assigned; when the bound port is // interface is up so the overlay IP is assigned; when the bound port is not
// not DefaultPort it is announced to peers via Body.mlkemPort. // DefaultPort it must be announced to peers via Body.mlkemPort.
func newPQTransport(overlayIP netip.Addr) (*pqTransport, error) { func newPQTransport(overlayIP netip.Addr) (*pqTransport, error) {
if !overlayIP.IsValid() { if !overlayIP.IsValid() {
return nil, fmt.Errorf("invalid overlay IP for pqkem transport") return nil, fmt.Errorf("invalid overlay IP for pqkem transport")
@@ -43,90 +33,40 @@ func newPQTransport(overlayIP netip.Addr) (*pqTransport, error) {
ip := net.IP(overlayIP.AsSlice()) ip := net.IP(overlayIP.AsSlice())
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: ip, Port: DefaultPort}) conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: ip, Port: DefaultPort})
if err != nil { if err != nil {
// Default port unavailable (rare on a dedicated overlay IP): fall back to an
// ephemeral port, which will be announced to peers.
log.Debugf("pqkem: default port %d unavailable on %s (%v), using an ephemeral port", DefaultPort, overlayIP, err) log.Debugf("pqkem: default port %d unavailable on %s (%v), using an ephemeral port", DefaultPort, overlayIP, err)
conn, err = net.ListenUDP("udp4", &net.UDPAddr{IP: ip, Port: 0}) conn, err = net.ListenUDP("udp4", &net.UDPAddr{IP: ip, Port: 0})
if err != nil { if err != nil {
return nil, fmt.Errorf("bind pqkem udp on overlay %s: %w", overlayIP, err) return nil, fmt.Errorf("bind pqkem udp on overlay %s: %w", overlayIP, err)
} }
} }
return &pqTransport{ return &pqTransport{conn: conn, port: conn.LocalAddr().(*net.UDPAddr).Port}, nil
conn: conn,
port: conn.LocalAddr().(*net.UDPAddr).Port,
peers: make(map[string]*net.UDPAddr),
byAddr: make(map[string]string),
}, nil
} }
// Port is the bound UDP port, announced to peers as Body.mlkemPort. // Send implements pqkem.Transport.
func (t *pqTransport) Port() int { return t.port } func (t *pqTransport) Send(endpoint netip.AddrPort, msg []byte) error {
_, err := t.conn.WriteToUDPAddrPort(msg, endpoint)
func (t *pqTransport) setManager(m *pqkem.Manager) {
t.mu.Lock()
t.mgr = m
t.mu.Unlock()
}
// AddPeer records where a peer's rekey messages are sent: its overlay IP and the
// pqkem port it announced. A zero port or invalid IP is ignored.
func (t *pqTransport) AddPeer(remoteID string, overlayIP netip.Addr, port int) {
if !overlayIP.IsValid() || port <= 0 {
return
}
addr := &net.UDPAddr{IP: net.IP(overlayIP.AsSlice()), Port: port}
t.mu.Lock()
if old, ok := t.peers[remoteID]; ok {
delete(t.byAddr, old.String())
}
t.peers[remoteID] = addr
t.byAddr[addr.String()] = remoteID
t.mu.Unlock()
}
func (t *pqTransport) RemovePeer(remoteID string) {
t.mu.Lock()
if a, ok := t.peers[remoteID]; ok {
delete(t.byAddr, a.String())
delete(t.peers, remoteID)
}
t.mu.Unlock()
}
// SendDataPath implements pqkem.Transport.
func (t *pqTransport) SendDataPath(remoteID string, msg []byte) error {
t.mu.RLock()
addr := t.peers[remoteID]
t.mu.RUnlock()
if addr == nil {
return fmt.Errorf("no data-path address for peer %s", remoteID)
}
_, err := t.conn.WriteToUDP(msg, addr)
return err return err
} }
// run is the receive loop: it maps each datagram's source overlay address to a peer // LocalPort implements pqkem.Transport.
// and feeds it to the manager. Exits when the socket is closed. func (t *pqTransport) LocalPort() int { return t.port }
func (t *pqTransport) run() {
buf := make([]byte, 2048) // Run implements pqkem.Transport: the receive loop, delivering each datagram as
for { // (source endpoint, msg). Exits when the socket is closed.
n, src, err := t.conn.ReadFromUDP(buf) func (t *pqTransport) Run(onInbound func(src netip.AddrPort, msg []byte)) {
if err != nil { go func() {
return buf := make([]byte, 2048)
for {
n, src, err := t.conn.ReadFromUDPAddrPort(buf)
if err != nil {
return
}
msg := make([]byte, n)
copy(msg, buf[:n])
onInbound(src, msg)
} }
t.mu.RLock() }()
remoteID := t.byAddr[src.String()]
mgr := t.mgr
t.mu.RUnlock()
if remoteID == "" || mgr == nil {
continue
}
msg := make([]byte, n)
copy(msg, buf[:n])
if err := mgr.OnDataPathMessage(remoteID, msg); err != nil {
log.Debugf("pqkem: inbound from %s: %v", remoteID, err)
}
}
} }
// Close implements pqkem.Transport.
func (t *pqTransport) Close() error { return t.conn.Close() } func (t *pqTransport) Close() error { return t.conn.Close() }