From 8aa45b7c8b5b519c6d43f9cefd2cdd43d6785456 Mon Sep 17 00:00:00 2001 From: breken-ai <312387581+breken-ai@users.noreply.github.com> Date: Fri, 25 Sep 2026 12:12:10 -0700 Subject: [PATCH] fix(relay): key IPv6 proxy mappings the way packets look them up --- relay/relay.go | 6 ++--- relay/relay_test.go | 66 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 69 insertions(+), 3 deletions(-) create mode 100644 relay/relay_test.go diff --git a/relay/relay.go b/relay/relay.go index df48136..ae0a54b 100644 --- a/relay/relay.go +++ b/relay/relay.go @@ -670,7 +670,7 @@ const addrCacheTTL = 5 * time.Minute // getCachedAddr returns a cached UDP address or resolves and caches it. // This avoids per-packet DNS lookups which are a major throughput bottleneck. func (s *UDPProxyServer) getCachedAddr(ip string, port int) (*net.UDPAddr, error) { - key := fmt.Sprintf("%s:%d", ip, port) + key := net.JoinHostPort(ip, strconv.Itoa(port)) // Check cache first if cached, ok := s.addrCache.Load(key); ok { @@ -1132,7 +1132,7 @@ func (s *UDPProxyServer) notifyServer(endpoint ClientEndpoint) { logger.Debug("Received proxy mapping from server: %v", mapping) // Store the mapping with current timestamp - key := fmt.Sprintf("%s:%d", endpoint.IP, endpoint.Port) + key := net.JoinHostPort(endpoint.IP, strconv.Itoa(endpoint.Port)) logger.Debug("About to store proxy mapping with key: %s (from endpoint IP=%s, Port=%d)", key, endpoint.IP, endpoint.Port) mapping.LastUsed = time.Now() if _, existed := s.proxyMappings.Load(key); existed { @@ -1147,7 +1147,7 @@ func (s *UDPProxyServer) notifyServer(endpoint ClientEndpoint) { // Updated to support multiple destinations func (s *UDPProxyServer) UpdateProxyMapping(sourceIP string, sourcePort int, destinations []PeerDestination) { - key := fmt.Sprintf("%s:%d", sourceIP, sourcePort) + key := net.JoinHostPort(sourceIP, strconv.Itoa(sourcePort)) mapping := ProxyMapping{ Destinations: destinations, LastUsed: time.Now(), diff --git a/relay/relay_test.go b/relay/relay_test.go new file mode 100644 index 0000000..16d60cf --- /dev/null +++ b/relay/relay_test.go @@ -0,0 +1,66 @@ +package relay + +import ( + "context" + "net" + "testing" + "time" + + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// TestRelayForwardsIPv6Peer checks that a mapping registered for an IPv6 +// client (as /update-destinations and the hole-punch path do) is found when +// that client's WireGuard packets arrive, and that the packet reaches an IPv6 +// destination. +func TestRelayForwardsIPv6Peer(t *testing.T) { + dest, err := net.ListenUDP("udp6", &net.UDPAddr{IP: net.IPv6loopback}) + if err != nil { + t.Skipf("IPv6 loopback not available: %v", err) + } + defer dest.Close() + + client, err := net.ListenUDP("udp6", &net.UDPAddr{IP: net.IPv6loopback}) + if err != nil { + t.Fatalf("listen client: %v", err) + } + defer client.Close() + + key, err := wgtypes.GeneratePrivateKey() + if err != nil { + t.Fatalf("generate key: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + s := NewUDPProxyServer(ctx, "[::1]:0", "http://127.0.0.1:0", key, "") + if err := s.Start(); err != nil { + t.Fatalf("start relay: %v", err) + } + defer s.Stop() + + clientAddr := client.LocalAddr().(*net.UDPAddr) + destAddr := dest.LocalAddr().(*net.UDPAddr) + s.UpdateProxyMapping(clientAddr.IP.String(), clientAddr.Port, []PeerDestination{ + {DestinationIP: destAddr.IP.String(), DestinationPort: destAddr.Port}, + }) + + // Minimal WireGuard handshake initiation: type 1, sender index 7. + initiation := make([]byte, 148) + initiation[0] = WireGuardMessageTypeHandshakeInitiation + initiation[4] = 7 + if _, err := client.WriteToUDP(initiation, s.conn.LocalAddr().(*net.UDPAddr)); err != nil { + t.Fatalf("send initiation: %v", err) + } + + if err := dest.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatalf("set deadline: %v", err) + } + buf := make([]byte, 1500) + n, _, err := dest.ReadFromUDP(buf) + if err != nil { + t.Fatalf("destination did not receive the forwarded initiation: %v", err) + } + if n != len(initiation) || buf[0] != WireGuardMessageTypeHandshakeInitiation { + t.Fatalf("destination got %d bytes (type %d), want %d bytes of type 1", n, buf[0], len(initiation)) + } +}