Models reattempts

This commit is contained in:
riccardom
2026-08-03 12:48:46 +02:00
parent 951614b803
commit 60a96c2aff
+128 -58
View File
@@ -8,9 +8,21 @@ import (
"time" "time"
) )
// DefaultRekeyInterval matches WireGuard's own REKEY_AFTER_TIME so the freshly const (
// rotated PSK is naturally adopted by WG's next handshake without forcing one. // DefaultRekeyInterval matches WireGuard's own REKEY_AFTER_TIME so the freshly
const DefaultRekeyInterval = 2 * time.Minute // 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 // 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 — Signal before the tunnel is up, the WireGuard
@@ -19,29 +31,48 @@ type Transport interface {
Send(remoteWgKey string, msg []byte) error 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 { type encodableMsg interface {
Encode() ([]byte, error) Encode() ([]byte, error)
} }
// Driver ties the pure Manager to the outside world: it runs the per-peer rekey // exchangeCtl tracks one in-flight exchange for a peer: its id, the cancel for its
// timer, dispatches inbound messages, and surfaces the derived PSK to the host via // retransmit/deadline goroutine, when it started (for convergence latency) and
// WGCallbackHandler. Convergence/retry/OnRekeyFailed are layered on in a later step. // 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 { type Driver struct {
mgr *Manager mgr *Manager
transport Transport transport Transport
wg WGCallbackHandler wg WGCallbackHandler
interval time.Duration
logger *slog.Logger logger *slog.Logger
mu sync.Mutex rekeyInterval time.Duration
peers map[string]context.CancelFunc retryInterval time.Duration
wait sync.WaitGroup 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 // 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 { func NewDriver(localWgKey string, t Transport, h WGCallbackHandler, interval time.Duration, logger *slog.Logger) *Driver {
if interval <= 0 { if interval <= 0 {
interval = DefaultRekeyInterval interval = DefaultRekeyInterval
@@ -49,54 +80,66 @@ func NewDriver(localWgKey string, t Transport, h WGCallbackHandler, interval tim
if logger == nil { if logger == nil {
logger = slog.Default() logger = slog.Default()
} }
ctx, cancel := context.WithCancel(context.Background())
return &Driver{ return &Driver{
mgr: NewManager(localWgKey), mgr: NewManager(localWgKey),
transport: t, transport: t,
wg: h, wg: h,
interval: interval, logger: logger,
logger: logger, rekeyInterval: interval,
peers: make(map[string]context.CancelFunc), 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 // AddPeer registers a remote peer and starts its rekey timer. Re-adding is a no-op.
// peer is a no-op.
func (d *Driver) AddPeer(remoteWgKey string) { func (d *Driver) AddPeer(remoteWgKey string) {
d.mu.Lock() d.mu.Lock()
defer d.mu.Unlock() defer d.mu.Unlock()
if _, ok := d.peers[remoteWgKey]; ok { if _, ok := d.peers[remoteWgKey]; ok {
return return
} }
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(d.rootCtx)
d.peers[remoteWgKey] = cancel d.peers[remoteWgKey] = cancel
d.wait.Add(1) d.wait.Add(1)
go d.rekeyLoop(ctx, remoteWgKey) 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) { func (d *Driver) RemovePeer(remoteWgKey string) {
d.mu.Lock() d.mu.Lock()
cancel, ok := d.peers[remoteWgKey] if cancel, ok := d.peers[remoteWgKey]; ok {
delete(d.peers, remoteWgKey)
d.mu.Unlock()
if ok {
cancel() 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() { func (d *Driver) Stop() {
d.mu.Lock() d.rootCancel()
for _, cancel := range d.peers {
cancel()
}
d.peers = make(map[string]context.CancelFunc)
d.mu.Unlock()
d.wait.Wait() 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 // 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 { func (d *Driver) HandleInbound(remoteWgKey string, raw []byte) error {
typ, msg, err := Decode(raw) typ, msg, err := Decode(raw)
if err != nil { if err != nil {
@@ -105,34 +148,56 @@ func (d *Driver) HandleInbound(remoteWgKey string, raw []byte) error {
switch typ { switch typ {
case MsgOffer: case MsgOffer:
answer, err := d.mgr.HandleOffer(remoteWgKey, msg.(*OfferMsg)) return d.handleOffer(remoteWgKey, msg.(*OfferMsg))
if err != nil {
return err
}
return d.send(remoteWgKey, answer)
case MsgAnswer: case MsgAnswer:
psk, confirm, err := d.mgr.HandleAnswer(remoteWgKey, msg.(*AnswerMsg)) return d.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: case MsgConfirm:
psk, err := d.mgr.HandleConfirm(remoteWgKey, msg.(*ConfirmMsg)) return d.handleConfirm(remoteWgKey, msg.(*ConfirmMsg))
if err != nil {
return err
}
return d.wg.OnNewPSKReady(remoteWgKey, psk)
default: default:
return fmt.Errorf("unhandled message type %d from %s", typ, remoteWgKey) 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 // 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.
@@ -144,12 +209,17 @@ func (d *Driver) initiateRekey(remoteWgKey string) error {
if err != nil { if err != nil {
return err 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) { func (d *Driver) rekeyLoop(ctx context.Context, remoteWgKey string) {
defer d.wait.Done() defer d.wait.Done()
t := time.NewTicker(d.interval) t := time.NewTicker(d.rekeyInterval)
defer t.Stop() defer t.Stop()
for { for {
select { select {