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

155 lines
4.2 KiB
Go

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")
}
}