From bca463d4c845ff96a83937b3b97d272079f84f6a Mon Sep 17 00:00:00 2001 From: riccardom Date: Tue, 21 Jul 2026 17:29:28 +0200 Subject: [PATCH] Models reattempts --- client/internal/pqkem/driver.go | 186 ++++++++++++++++++++++---------- 1 file changed, 128 insertions(+), 58 deletions(-) diff --git a/client/internal/pqkem/driver.go b/client/internal/pqkem/driver.go index 10e5e2a10..c26d4ec94 100644 --- a/client/internal/pqkem/driver.go +++ b/client/internal/pqkem/driver.go @@ -8,9 +8,21 @@ import ( "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 +const ( + // DefaultRekeyInterval matches WireGuard's own REKEY_AFTER_TIME so the freshly + // rotated PSK is naturally adopted by WG's next handshake without forcing one. + DefaultRekeyInterval = 2 * time.Minute + // DefaultRetryInterval is how often an in-flight exchange retransmits. + DefaultRetryInterval = 2 * time.Second + // DefaultConvergenceTimeout bounds a single exchange before it is declared failed. + DefaultConvergenceTimeout = 20 * time.Second + // DefaultMaxRekeyFailures is how many consecutive rekey (non-initial) failures are + // tolerated before OnRekeyFailed is raised. The initial exchange fails immediately. + DefaultMaxRekeyFailures = 3 + // confirmRetransmits is how many extra times the initiator best-effort resends the + // confirm (the exchange's unacked last message) to reduce a stuck responder. + confirmRetransmits = 3 +) // 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 @@ -19,29 +31,48 @@ 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. +// exchangeCtl tracks one in-flight exchange for a peer: its id, the cancel for its +// retransmit/deadline goroutine, when it started (for convergence latency) and +// whether it is the peer's first (initial) exchange (which fails hard on timeout). +type exchangeCtl struct { + id ExchangeID + cancel context.CancelFunc + startedAt time.Time + isInitial bool +} + +// Driver ties the pure Manager to the outside world: per-peer rekey timer, inbound +// dispatch, retransmit/convergence tracking, and surfacing events to the host via +// WGCallbackHandler. 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 + rekeyInterval time.Duration + retryInterval time.Duration + convergenceTimeout time.Duration + maxRekeyFailures int + + rootCtx context.Context + rootCancel context.CancelFunc + + mu sync.Mutex + peers map[string]context.CancelFunc // per-peer rekey loop + exchanges map[string]*exchangeCtl // in-flight exchange per peer + established map[string]bool // peer has completed at least one exchange + failures map[string]int // consecutive rekey failures per peer + 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(). +// DefaultRekeyInterval; a nil logger falls back to slog.Default(). Retry/timeout/K +// use their defaults and can be overridden on the returned struct before use. func NewDriver(localWgKey string, t Transport, h WGCallbackHandler, interval time.Duration, logger *slog.Logger) *Driver { if interval <= 0 { interval = DefaultRekeyInterval @@ -49,54 +80,66 @@ func NewDriver(localWgKey string, t Transport, h WGCallbackHandler, interval tim if logger == nil { logger = slog.Default() } + ctx, cancel := context.WithCancel(context.Background()) return &Driver{ - mgr: NewManager(localWgKey), - transport: t, - wg: h, - interval: interval, - logger: logger, - peers: make(map[string]context.CancelFunc), + mgr: NewManager(localWgKey), + transport: t, + wg: h, + logger: logger, + rekeyInterval: interval, + retryInterval: DefaultRetryInterval, + convergenceTimeout: DefaultConvergenceTimeout, + maxRekeyFailures: DefaultMaxRekeyFailures, + rootCtx: ctx, + rootCancel: cancel, + peers: make(map[string]context.CancelFunc), + exchanges: make(map[string]*exchangeCtl), + established: make(map[string]bool), + failures: make(map[string]int), } } -// AddPeer registers a remote peer and starts its rekey timer. Re-adding an existing -// peer is a no-op. +// AddPeer registers a remote peer and starts its rekey timer. Re-adding 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()) + ctx, cancel := context.WithCancel(d.rootCtx) d.peers[remoteWgKey] = cancel d.wait.Add(1) go d.rekeyLoop(ctx, remoteWgKey) } -// RemovePeer stops a peer's rekey timer and drops its state. +// RemovePeer stops a peer's rekey timer and any in-flight exchange, and drops state. func (d *Driver) RemovePeer(remoteWgKey string) { d.mu.Lock() - cancel, ok := d.peers[remoteWgKey] - delete(d.peers, remoteWgKey) - d.mu.Unlock() - if ok { + if cancel, ok := d.peers[remoteWgKey]; ok { cancel() + delete(d.peers, remoteWgKey) } + if ex, ok := d.exchanges[remoteWgKey]; ok { + ex.cancel() + delete(d.exchanges, remoteWgKey) + } + delete(d.established, remoteWgKey) + delete(d.failures, remoteWgKey) + d.mu.Unlock() } -// Stop cancels all peer timers and waits for their goroutines to exit. +// Stop cancels all timers and in-flight exchanges and waits for 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.rootCancel() d.wait.Wait() + d.mu.Lock() + d.peers = make(map[string]context.CancelFunc) + d.exchanges = make(map[string]*exchangeCtl) + d.mu.Unlock() } // HandleInbound decodes an incoming message and drives the exchange, sending any -// response via the transport and surfacing a derived PSK through the callback. +// response via the transport and surfacing derived PSKs / convergence to the host. func (d *Driver) HandleInbound(remoteWgKey string, raw []byte) error { typ, msg, err := Decode(raw) if err != nil { @@ -105,34 +148,56 @@ func (d *Driver) HandleInbound(remoteWgKey string, raw []byte) error { switch typ { case MsgOffer: - answer, err := d.mgr.HandleOffer(remoteWgKey, msg.(*OfferMsg)) - if err != nil { - return err - } - return d.send(remoteWgKey, answer) - + return d.handleOffer(remoteWgKey, msg.(*OfferMsg)) 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) - + return d.handleAnswer(remoteWgKey, msg.(*AnswerMsg)) case MsgConfirm: - psk, err := d.mgr.HandleConfirm(remoteWgKey, msg.(*ConfirmMsg)) - if err != nil { - return err - } - return d.wg.OnNewPSKReady(remoteWgKey, psk) - + return d.handleConfirm(remoteWgKey, msg.(*ConfirmMsg)) default: return fmt.Errorf("unhandled message type %d from %s", typ, remoteWgKey) } } +func (d *Driver) handleOffer(remoteWgKey string, o *OfferMsg) error { + answer, err := d.mgr.HandleOffer(remoteWgKey, o) + if err != nil { + return err + } + // Arm a responder exchange (deadline waiting for the confirm) unless one for this + // id is already armed — a retransmitted offer just re-sends the same answer. + d.armExchange(remoteWgKey, o.ExchangeID, nil) + return d.send(remoteWgKey, answer) +} + +func (d *Driver) handleAnswer(remoteWgKey string, a *AnswerMsg) error { + psk, confirm, err := d.mgr.HandleAnswer(remoteWgKey, a) + if err != nil { + return err + } + if err := d.wg.OnNewPSKReady(remoteWgKey, psk); err != nil { + return err + } + // Initiator has the PSK now -> this side has converged. + d.onConverged(remoteWgKey, a.ExchangeID) + if err := d.send(remoteWgKey, confirm); err != nil { + return err + } + d.retransmitConfirm(remoteWgKey, confirm) + return nil +} + +func (d *Driver) handleConfirm(remoteWgKey string, c *ConfirmMsg) error { + psk, err := d.mgr.HandleConfirm(remoteWgKey, c) + if err != nil { + return err + } + if err := d.wg.OnNewPSKReady(remoteWgKey, psk); err != nil { + return err + } + d.onConverged(remoteWgKey, c.ExchangeID) + return nil +} + // 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. @@ -144,12 +209,17 @@ func (d *Driver) initiateRekey(remoteWgKey string) error { if err != nil { return err } - return d.send(remoteWgKey, offer) + raw, err := offer.Encode() + if err != nil { + return err + } + d.armExchange(remoteWgKey, offer.ExchangeID, raw) + return d.transport.Send(remoteWgKey, raw) } func (d *Driver) rekeyLoop(ctx context.Context, remoteWgKey string) { defer d.wait.Done() - t := time.NewTicker(d.interval) + t := time.NewTicker(d.rekeyInterval) defer t.Stop() for { select {