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, err := NewForwarder(netip.MustParseAddr("127.0.0.1")) if err != nil { t.Fatalf("NewForwarder: %v", err) } t.Cleanup(f.Close) 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") } }