From 465881708fd6d8462b23410df0f78fd000c92d6f Mon Sep 17 00:00:00 2001 From: jdl Date: Mon, 8 Jun 2026 17:59:14 +0200 Subject: [PATCH] WIP --- peer/forwarder.go | 84 ++++++++++++++++++++++ peer/forwarder_test.go | 155 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 239 insertions(+) create mode 100644 peer/forwarder.go create mode 100644 peer/forwarder_test.go diff --git a/peer/forwarder.go b/peer/forwarder.go new file mode 100644 index 0000000..0555d94 --- /dev/null +++ b/peer/forwarder.go @@ -0,0 +1,84 @@ +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) + } + } +} diff --git a/peer/forwarder_test.go b/peer/forwarder_test.go new file mode 100644 index 0000000..1418304 --- /dev/null +++ b/peer/forwarder_test.go @@ -0,0 +1,155 @@ +package peer + +import ( + "net" + "net/netip" + "testing" + "time" +) + +// newTestForwarder creates a Forwarder using 127.0.0.1 as the VPN IP so that +// 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)) + } + }) + return f +} + +// listenRandom opens a UDP listener on 127.0.0.1 with a kernel-assigned port. +func listenRandom(t *testing.T) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatalf("listenRandom: %v", err) + } + t.Cleanup(func() { conn.Close() }) + return conn +} + +func addrPort(conn *net.UDPConn) netip.AddrPort { + ap, _ := netip.ParseAddrPort(conn.LocalAddr().String()) + return ap +} + +func recvFrom(t *testing.T, conn *net.UDPConn, timeout time.Duration) []byte { + t.Helper() + conn.SetReadDeadline(time.Now().Add(timeout)) + buf := make([]byte, 1<<16) + n, _, err := conn.ReadFromUDP(buf) + if err != nil { + t.Fatalf("recvFrom: %v", err) + } + return buf[:n] +} + +func sendTo(t *testing.T, dst netip.AddrPort, payload []byte) { + t.Helper() + conn, err := net.DialUDP("udp4", nil, net.UDPAddrFromAddrPort(dst)) + if err != nil { + t.Fatalf("sendTo dial: %v", err) + } + defer conn.Close() + if _, err := conn.Write(payload); err != nil { + t.Fatalf("sendTo write: %v", err) + } +} + +func TestForwarderBasic(t *testing.T) { + f := newTestForwarder(t) + target := listenRandom(t) + + const peerByte = byte(11) + f.SetEndpoint(peerByte, addrPort(target)) + + fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte)) + sendTo(t, fwdAddr, []byte("hello wireguard")) + + got := recvFrom(t, target, time.Second) + if string(got) != "hello wireguard" { + t.Fatalf("got %q, want %q", got, "hello wireguard") + } +} + +func TestForwarderEndpointUpdate(t *testing.T) { + f := newTestForwarder(t) + target1 := listenRandom(t) + target2 := listenRandom(t) + + const peerByte = byte(22) + fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte)) + + f.SetEndpoint(peerByte, addrPort(target1)) + sendTo(t, fwdAddr, []byte("first")) + recvFrom(t, target1, time.Second) + + // Update endpoint to target2. + f.SetEndpoint(peerByte, addrPort(target2)) + sendTo(t, fwdAddr, []byte("second")) + + got := recvFrom(t, target2, time.Second) + if string(got) != "second" { + t.Fatalf("got %q, want %q", got, "second") + } + + // Nothing should have arrived at target1 after the update. + target1.SetReadDeadline(time.Now().Add(50 * time.Millisecond)) + buf := make([]byte, 16) + if n, _, err := target1.ReadFromUDP(buf); err == nil { + t.Fatalf("unexpected packet at old target: %q", buf[:n]) + } +} + +func TestForwarderRemoveEndpoint(t *testing.T) { + f := newTestForwarder(t) + target := listenRandom(t) + + const peerByte = byte(33) + fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte)) + + f.SetEndpoint(peerByte, addrPort(target)) + sendTo(t, fwdAddr, []byte("before")) + recvFrom(t, target, time.Second) + + f.RemoveEndpoint(peerByte) + + // After removal the forwarder port is closed. Send a packet and verify + // nothing arrives at target within a short window. + sendTo(t, fwdAddr, []byte("after")) + target.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) + buf := make([]byte, 16) + if n, _, err := target.ReadFromUDP(buf); err == nil { + t.Fatalf("packet arrived after RemoveEndpoint: %q", buf[:n]) + } +} + +func TestForwarderMultiplePeers(t *testing.T) { + f := newTestForwarder(t) + targetA := listenRandom(t) + targetB := listenRandom(t) + + const ( + byteA = byte(10) + byteB = byte(11) + ) + vpnIP := netip.MustParseAddr("127.0.0.1") + fwdA := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(byteA)) + fwdB := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(byteB)) + + f.SetEndpoint(byteA, addrPort(targetA)) + f.SetEndpoint(byteB, addrPort(targetB)) + + sendTo(t, fwdA, []byte("for-a")) + sendTo(t, fwdB, []byte("for-b")) + + if got := string(recvFrom(t, targetA, time.Second)); got != "for-a" { + t.Fatalf("targetA got %q, want %q", got, "for-a") + } + if got := string(recvFrom(t, targetB, time.Second)); got != "for-b" { + t.Fatalf("targetB got %q, want %q", got, "for-b") + } +}