diff --git a/client/internal/pqkem/manager.go b/client/internal/pqkem/manager.go index 429d38716..cb92a9b2e 100644 --- a/client/internal/pqkem/manager.go +++ b/client/internal/pqkem/manager.go @@ -57,7 +57,8 @@ type Transport interface { // 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 stops delivery, drains any in-flight onInbound callback, and releases the + // socket, so no callback is still running when Close returns. Close() error } @@ -281,19 +282,25 @@ func (m *Manager) Stop() { // has cancelled here no new Add can race Wait. m.mu.Lock() m.rootCancel() - m.mu.Unlock() - m.wait.Wait() - m.mu.Lock() t := m.transport m.transport = nil - m.exchanges = make(map[RemoteID]*exchangeCtl) - m.psks = make(map[RemoteID]PSK) m.mu.Unlock() + + // Close and drain the data-path receive loop BEFORE clearing peer state: Close waits + // for any in-flight onInbound to return, so no inbound message can derive a PSK into a + // cleared map, and the engine only tears the WireGuard interface down after this Stop + // returns, so no late SetPresharedKey hits a torn-down device. if t != nil { if err := t.Close(); err != nil { m.logger.Warn("pqkem: closing data-path transport", "err", err) } } + m.wait.Wait() + + m.mu.Lock() + m.exchanges = make(map[RemoteID]*exchangeCtl) + m.psks = make(map[RemoteID]PSK) + m.mu.Unlock() } // ---- Signalling channel (host-driven; rides the host's negotiation) ---- diff --git a/client/internal/pqkem_transport.go b/client/internal/pqkem_transport.go index 920e655c4..7096c9e9d 100644 --- a/client/internal/pqkem_transport.go +++ b/client/internal/pqkem_transport.go @@ -4,6 +4,7 @@ import ( "fmt" "net" "net/netip" + "sync" log "github.com/sirupsen/logrus" ) @@ -20,6 +21,7 @@ const DefaultPort = 51833 type pqTransport struct { conn *net.UDPConn port int + wg sync.WaitGroup } // newPQTransport binds a UDP socket on the WG overlay IP, preferring DefaultPort and @@ -58,7 +60,9 @@ func (t *pqTransport) LocalPort() int { return t.port } // Run implements pqkem.Transport: the receive loop, delivering each datagram as // (source endpoint, msg). Exits when the socket is closed. func (t *pqTransport) Run(onInbound func(src netip.AddrPort, msg []byte)) { + t.wg.Add(1) go func() { + defer t.wg.Done() buf := make([]byte, 2048) for { n, src, err := t.conn.ReadFromUDPAddrPort(buf) @@ -72,5 +76,10 @@ func (t *pqTransport) Run(onInbound func(src netip.AddrPort, msg []byte)) { }() } -// Close implements pqkem.Transport. -func (t *pqTransport) Close() error { return t.conn.Close() } +// Close implements pqkem.Transport. It closes the socket (unblocking the read) and drains +// the receive loop, so no onInbound callback is still running when Close returns. +func (t *pqTransport) Close() error { + err := t.conn.Close() + t.wg.Wait() + return err +}