package peer import ( "log" "net" "net/netip" "sync/atomic" ) 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 + // 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. type Forwarder struct { vpnIP netip.Addr peers [256]atomic.Pointer[fwdPeer] } // 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() } 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() } // 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() } } func (p *fwdPeer) run() { network := "udp4" if p.endpoint.Addr().Is6() { network = "udp6" } sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(p.endpoint)) if err != nil { log.Printf("[Forwarder] dial %v: %v", p.endpoint, err) return } defer sendConn.Close() buf := make([]byte, 1<<16) for { n, err := p.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) } } }