From 3f15ec14c14a437659da90f703b4a3ca8afd02cf Mon Sep 17 00:00:00 2001 From: riccardom Date: Tue, 21 Jul 2026 13:09:21 +0200 Subject: [PATCH] Adds driver to glue together manager and outside world --- client/internal/pqkem/driver.go | 172 +++++++++++++++++++++++++++ client/internal/pqkem/driver_test.go | 96 +++++++++++++++ 2 files changed, 268 insertions(+) create mode 100644 client/internal/pqkem/driver.go create mode 100644 client/internal/pqkem/driver_test.go diff --git a/client/internal/pqkem/driver.go b/client/internal/pqkem/driver.go new file mode 100644 index 000000000..10e5e2a10 --- /dev/null +++ b/client/internal/pqkem/driver.go @@ -0,0 +1,172 @@ +package pqkem + +import ( + "context" + "fmt" + "log/slog" + "sync" + "time" +) + +// DefaultRekeyInterval matches WireGuard's own REKEY_AFTER_TIME so the freshly +// rotated PSK is naturally adopted by WG's next handshake without forcing one. +const DefaultRekeyInterval = 2 * time.Minute + +// 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 +// tunnel for rekeys — so the driver never needs to know which is in use. +type Transport interface { + Send(remoteWgKey string, msg []byte) error +} + +// encodableMsg is the common shape of the three wire messages, letting the driver +// send any of them uniformly. +type encodableMsg interface { + Encode() ([]byte, error) +} + +// Driver ties the pure Manager to the outside world: it runs the per-peer rekey +// timer, dispatches inbound messages, and surfaces the derived PSK to the host via +// WGCallbackHandler. Convergence/retry/OnRekeyFailed are layered on in a later step. +type Driver struct { + mgr *Manager + transport Transport + wg WGCallbackHandler + interval time.Duration + logger *slog.Logger + + mu sync.Mutex + peers map[string]context.CancelFunc + wait sync.WaitGroup +} + +// NewDriver builds a driver for the local peer. A zero interval falls back to +// DefaultRekeyInterval; a nil logger falls back to slog.Default(). +func NewDriver(localWgKey string, t Transport, h WGCallbackHandler, interval time.Duration, logger *slog.Logger) *Driver { + if interval <= 0 { + interval = DefaultRekeyInterval + } + if logger == nil { + logger = slog.Default() + } + return &Driver{ + mgr: NewManager(localWgKey), + transport: t, + wg: h, + interval: interval, + logger: logger, + peers: make(map[string]context.CancelFunc), + } +} + +// AddPeer registers a remote peer and starts its rekey timer. Re-adding an existing +// peer is a no-op. +func (d *Driver) AddPeer(remoteWgKey string) { + d.mu.Lock() + defer d.mu.Unlock() + if _, ok := d.peers[remoteWgKey]; ok { + return + } + ctx, cancel := context.WithCancel(context.Background()) + d.peers[remoteWgKey] = cancel + d.wait.Add(1) + go d.rekeyLoop(ctx, remoteWgKey) +} + +// RemovePeer stops a peer's rekey timer and drops its state. +func (d *Driver) RemovePeer(remoteWgKey string) { + d.mu.Lock() + cancel, ok := d.peers[remoteWgKey] + delete(d.peers, remoteWgKey) + d.mu.Unlock() + if ok { + cancel() + } +} + +// Stop cancels all peer timers and waits for their goroutines to exit. +func (d *Driver) Stop() { + d.mu.Lock() + for _, cancel := range d.peers { + cancel() + } + d.peers = make(map[string]context.CancelFunc) + d.mu.Unlock() + d.wait.Wait() +} + +// HandleInbound decodes an incoming message and drives the exchange, sending any +// response via the transport and surfacing a derived PSK through the callback. +func (d *Driver) HandleInbound(remoteWgKey string, raw []byte) error { + typ, msg, err := Decode(raw) + if err != nil { + return fmt.Errorf("decode from %s: %w", remoteWgKey, err) + } + + switch typ { + case MsgOffer: + answer, err := d.mgr.HandleOffer(remoteWgKey, msg.(*OfferMsg)) + if err != nil { + return err + } + return d.send(remoteWgKey, answer) + + case MsgAnswer: + psk, confirm, err := d.mgr.HandleAnswer(remoteWgKey, msg.(*AnswerMsg)) + if err != nil { + return err + } + if err := d.wg.OnNewPSKReady(remoteWgKey, psk); err != nil { + return err + } + return d.send(remoteWgKey, confirm) + + case MsgConfirm: + psk, err := d.mgr.HandleConfirm(remoteWgKey, msg.(*ConfirmMsg)) + if err != nil { + return err + } + return d.wg.OnNewPSKReady(remoteWgKey, psk) + + default: + return fmt.Errorf("unhandled message type %d from %s", typ, remoteWgKey) + } +} + +// 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 +// but directly callable) so tests can drive a rekey without waiting on the ticker. +func (d *Driver) initiateRekey(remoteWgKey string) error { + if !d.mgr.IsInitiator(remoteWgKey) { + return nil + } + offer, err := d.mgr.StartExchange(remoteWgKey) + if err != nil { + return err + } + return d.send(remoteWgKey, offer) +} + +func (d *Driver) rekeyLoop(ctx context.Context, remoteWgKey string) { + defer d.wait.Done() + t := time.NewTicker(d.interval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + if err := d.initiateRekey(remoteWgKey); err != nil { + d.logger.Error("pqkem rekey failed to start", "peer", remoteWgKey, "err", err) + } + } + } +} + +func (d *Driver) send(remoteWgKey string, m encodableMsg) error { + raw, err := m.Encode() + if err != nil { + return err + } + return d.transport.Send(remoteWgKey, raw) +} diff --git a/client/internal/pqkem/driver_test.go b/client/internal/pqkem/driver_test.go new file mode 100644 index 000000000..59fb381da --- /dev/null +++ b/client/internal/pqkem/driver_test.go @@ -0,0 +1,96 @@ +package pqkem + +import ( + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// loopback delivers a sent message synchronously to the peer driver's HandleInbound, +// attributing it to localKey (the sender). +type loopback struct { + localKey string + peer *Driver +} + +func (l *loopback) Send(remoteWgKey string, msg []byte) error { + cp := append([]byte(nil), msg...) + return l.peer.HandleInbound(l.localKey, cp) +} + +type fakeWG struct { + mu sync.Mutex + psks map[string]PSK + failed []string +} + +func newFakeWG() *fakeWG { return &fakeWG{psks: map[string]PSK{}} } + +func (f *fakeWG) OnNewPSKReady(remoteWgKey string, psk PSK) error { + f.mu.Lock() + defer f.mu.Unlock() + f.psks[remoteWgKey] = psk + return nil +} + +func (f *fakeWG) OnRekeyFailed(remoteWgKey string) error { + f.mu.Lock() + defer f.mu.Unlock() + f.failed = append(f.failed, remoteWgKey) + return nil +} + +func (f *fakeWG) psk(peer string) PSK { + f.mu.Lock() + defer f.mu.Unlock() + return f.psks[peer] +} + +func TestDriver_ExchangeConverges(t *testing.T) { + lbA := &loopback{localKey: "aaaa"} + lbB := &loopback{localKey: "bbbb"} + wgA := newFakeWG() + wgB := newFakeWG() + + // long interval so the ticker never fires during the test; we drive manually. + dA := NewDriver("aaaa", lbA, wgA, time.Hour, nil) + dB := NewDriver("bbbb", lbB, wgB, time.Hour, nil) + lbA.peer = dB // A sends -> B receives + lbB.peer = dA // B sends -> A receives + + dA.AddPeer("bbbb") + dB.AddPeer("aaaa") + defer dA.Stop() + defer dB.Stop() + + // B is the initiator ("bbbb" > "aaaa"). + require.NoError(t, dB.initiateRekey("aaaa")) + + pskB := wgB.psk("aaaa") // B committed on the answer + pskA := wgA.psk("bbbb") // A committed on the confirm + require.NotEqual(t, PSK{}, pskA, "responder A 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") +} + +func TestDriver_NonInitiatorDoesNothing(t *testing.T) { + lbA := &loopback{localKey: "aaaa"} + wgA := newFakeWG() + dA := NewDriver("aaaa", lbA, wgA, time.Hour, nil) + // no peer driver wired; if A wrongly initiated, Send would nil-panic. + dA.AddPeer("bbbb") + defer dA.Stop() + + // A is NOT the initiator vs "bbbb" -> initiateRekey is a no-op, no Send. + require.NoError(t, dA.initiateRekey("bbbb")) + require.Equal(t, PSK{}, wgA.psk("bbbb")) +} + +func TestDriver_StopIsIdempotent(t *testing.T) { + dA := NewDriver("aaaa", &loopback{localKey: "aaaa"}, newFakeWG(), time.Hour, nil) + dA.AddPeer("bbbb") + dA.Stop() + dA.Stop() // must not panic or hang +}