Files
vppn/peer/forwarder.go
2026-06-08 18:36:32 +02:00

96 lines
2.6 KiB
Go

package peer
import (
"fmt"
"log"
"net"
"net/netip"
"sync/atomic"
)
const ForwarderBasePort = 5000
// Forwarder is a relay-side UDP forwarding service. It listens on one UDP port
// 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).
//
// 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
sendConns [256]atomic.Pointer[net.UDPConn] // updated atomically; read by run goroutines
listenConns [256]*net.UDPConn // written once in NewForwarder; read by 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)
}
return f, nil
}
// 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()
}
}
// 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 ep.Addr().Is6() {
network = "udp6"
}
sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(ep))
if err != nil {
log.Printf("[Forwarder] dial %v: %v", ep, err)
return
}
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 := listenConn.Read(buf)
if err != nil {
return
}
conn := f.sendConns[b].Load()
if conn == nil {
continue
}
if _, err := conn.Write(buf[:n]); err != nil {
log.Printf("[Forwarder] write: %v", err)
}
}
}