diff --git a/peer/forwarder.go b/peer/forwarder.go index 0555d94..8382f08 100644 --- a/peer/forwarder.go +++ b/peer/forwarder.go @@ -1,6 +1,7 @@ package peer import ( + "fmt" "log" "net" "net/netip" @@ -10,75 +11,85 @@ import ( const ForwarderBasePort = 5000 // Forwarder is a relay-side UDP forwarding service. It listens on one UDP port -// per peer on the relay's VPN interface and forwards received bytes to each -// peer's physical WireGuard port. Port assignment: ForwarderBasePort + +// per possible peer on the relay's VPN interface and forwards received bytes to +// each peer's physical WireGuard port. Port assignment: ForwarderBasePort + // peer_vpn_ip_byte (peer .11 → port 5011). // -// Incoming packets are authenticated for free by the relay's outer WireGuard -// session before they reach the forwarder. +// All 256 listeners are opened at construction time and kept open for the +// lifetime of the Forwarder. SetEndpoint and RemoveEndpoint only update the +// dialed send connection for a given peer byte. type Forwarder struct { - vpnIP netip.Addr - peers [256]atomic.Pointer[fwdPeer] + vpnIP netip.Addr + sendConns [256]atomic.Pointer[net.UDPConn] // updated atomically; read by run goroutines + listenConns [256]*net.UDPConn // written once in NewForwarder; read by Close } -// fwdPeer is immutable once created. Endpoint changes produce a new fwdPeer. -type fwdPeer struct { - listenConn *net.UDPConn - endpoint netip.AddrPort -} - -// NewForwarder creates a Forwarder that will bind listeners to vpnIP. -func NewForwarder(vpnIP netip.Addr) *Forwarder { - return &Forwarder{vpnIP: vpnIP} -} - -// SetEndpoint registers or updates the physical WireGuard endpoint for the -// peer identified by peerIPByte (last octet of its VPN IP). The old fwdPeer, -// if any, is closed before the new one is opened on the same port. -func (f *Forwarder) SetEndpoint(peerIPByte byte, ep netip.AddrPort) { - if old := f.peers[peerIPByte].Swap(nil); old != nil { - old.listenConn.Close() +// NewForwarder opens all 256 forwarding listeners on vpnIP and starts their +// goroutines. It returns an error if any listener fails to bind. +func NewForwarder(vpnIP netip.Addr) (*Forwarder, error) { + f := &Forwarder{vpnIP: vpnIP} + for i := range 256 { + listenAddr := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(i)) + listenConn, err := net.ListenUDP("udp4", net.UDPAddrFromAddrPort(listenAddr)) + if err != nil { + for j := range i { + f.listenConns[j].Close() + } + return nil, fmt.Errorf("forwarder listen %v: %w", listenAddr, err) + } + f.listenConns[i] = listenConn + go f.run(byte(i), listenConn) } - - listenAddr := netip.AddrPortFrom(f.vpnIP, ForwarderBasePort+uint16(peerIPByte)) - listenConn, err := net.ListenUDP("udp4", net.UDPAddrFromAddrPort(listenAddr)) - if err != nil { - log.Printf("[Forwarder] listen %v: %v", listenAddr, err) - return - } - - p := &fwdPeer{listenConn: listenConn, endpoint: ep} - f.peers[peerIPByte].Store(p) - go p.run() + return f, nil } -// RemoveEndpoint stops forwarding for the given peer and closes its listener. -func (f *Forwarder) RemoveEndpoint(peerIPByte byte) { - if old := f.peers[peerIPByte].Swap(nil); old != nil { - old.listenConn.Close() +// Close shuts down all listeners and open send connections. +func (f *Forwarder) Close() { + for i := range 256 { + if old := f.sendConns[i].Swap(nil); old != nil { + old.Close() + } + f.listenConns[i].Close() } } -func (p *fwdPeer) run() { +// SetEndpoint registers or replaces the physical WireGuard endpoint for the +// peer identified by b (last octet of its VPN IP). +func (f *Forwarder) SetEndpoint(b byte, ep netip.AddrPort) { network := "udp4" - if p.endpoint.Addr().Is6() { + if ep.Addr().Is6() { network = "udp6" } - sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(p.endpoint)) + sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(ep)) if err != nil { - log.Printf("[Forwarder] dial %v: %v", p.endpoint, err) + log.Printf("[Forwarder] dial %v: %v", ep, err) return } - defer sendConn.Close() + if old := f.sendConns[b].Swap(sendConn); old != nil { + old.Close() + } +} +// RemoveEndpoint stops forwarding for the given peer. +func (f *Forwarder) RemoveEndpoint(b byte) { + if old := f.sendConns[b].Swap(nil); old != nil { + old.Close() + } +} + +func (f *Forwarder) run(b byte, listenConn *net.UDPConn) { buf := make([]byte, 1<<16) for { - n, err := p.listenConn.Read(buf) + n, err := listenConn.Read(buf) if err != nil { return } - if _, err := sendConn.Write(buf[:n]); err != nil { - log.Printf("[Forwarder] write to %v: %v", p.endpoint, err) + conn := f.sendConns[b].Load() + if conn == nil { + continue + } + if _, err := conn.Write(buf[:n]); err != nil { + log.Printf("[Forwarder] write: %v", err) } } } diff --git a/peer/forwarder_test.go b/peer/forwarder_test.go index 1418304..a788d4a 100644 --- a/peer/forwarder_test.go +++ b/peer/forwarder_test.go @@ -11,12 +11,11 @@ import ( // no real WireGuard interface is needed. func newTestForwarder(t *testing.T) *Forwarder { t.Helper() - f := NewForwarder(netip.MustParseAddr("127.0.0.1")) - t.Cleanup(func() { - for b := range 256 { - f.RemoveEndpoint(byte(b)) - } - }) + f, err := NewForwarder(netip.MustParseAddr("127.0.0.1")) + if err != nil { + t.Fatalf("NewForwarder: %v", err) + } + t.Cleanup(f.Close) return f }