From b2efd984f7accb0cec6ad6fbd1b9593078cceacb Mon Sep 17 00:00:00 2001 From: Owen Date: Thu, 24 Sep 2026 14:32:38 -0400 Subject: [PATCH] Resolve exit node endpoints over whichever address family is reachable Fixes fosrl/android#42 Fixes fosrl/pangolin#3471 Fixes fosrl/olm#108 --- holepunch/holepunch.go | 130 +++++++++++++++++------- util/util.go | 221 +++++++++++++++++++++-------------------- 2 files changed, 205 insertions(+), 146 deletions(-) diff --git a/holepunch/holepunch.go b/holepunch/holepunch.go index d7e1c14..99828d7 100644 --- a/holepunch/holepunch.go +++ b/holepunch/holepunch.go @@ -39,6 +39,13 @@ type Manager struct { updateChan chan struct{} // signals the goroutine to refresh exit nodes publicDNS []string + // disabled, when true, makes Start/StartMultipleExitNodes/TriggerHolePunch + // no-ops so no UDP hole punch packet is ever sent - e.g. a user-configured + // "disable hole punching" setting must fully suppress outbound hole punch + // traffic, not just change what's reported to the server (which is all it + // did before - see https://github.com/fosrl/olm/issues/134). + disabled bool + sendHolepunchInterval time.Duration sendHolepunchIntervalMin time.Duration sendHolepunchIntervalMax time.Duration @@ -66,6 +73,21 @@ func NewManager(sharedBind *bind.SharedBind, ID string, clientType string, publi } } +// SetEnabled controls whether this manager may send UDP hole punch packets. +// When disabled, Start/StartMultipleExitNodes/TriggerHolePunch are no-ops. +// Safe to call before or after Start; disabling an already-running manager +// stops it immediately. +func (m *Manager) SetEnabled(enabled bool) { + m.mu.Lock() + m.disabled = !enabled + running := m.running + m.mu.Unlock() + + if m.disabled && running { + m.Stop() + } +} + // SetToken updates the authentication token used for hole punching func (m *Manager) SetToken(token string) { m.mu.Lock() @@ -269,11 +291,51 @@ func (m *Manager) ResetServerHolepunchInterval() { } } +// resolveExitNodeAddrs resolves exitNode.Endpoint to every candidate UDP +// address (all address families) it currently has, rather than collapsing to +// a single IPv4-preferred address. Hole punch sends are cheap, best-effort +// UDP packets, so trying every candidate costs little and means whichever +// address family the local network path actually has a route for gets used - +// e.g. on an IPv6-only/NAT64 network where an IPv4 candidate exists in DNS +// but has no route at all. See https://github.com/fosrl/olm/issues/108. +func (m *Manager) resolveExitNodeAddrs(exitNode ExitNode) ([]*net.UDPAddr, error) { + var hosts []string + var err error + if len(m.publicDNS) > 0 { + hosts, err = util.ResolveDomainAllUpstream(exitNode.Endpoint, m.publicDNS) + } else { + hosts, err = util.ResolveDomainAll(exitNode.Endpoint) + } + if err != nil { + return nil, err + } + + var addrs []*net.UDPAddr + for _, host := range hosts { + serverAddr := net.JoinHostPort(host, strconv.Itoa(int(exitNode.RelayPort))) + remoteAddr, err := net.ResolveUDPAddr("udp", serverAddr) + if err != nil { + logger.Error("Failed to resolve UDP address %s: %v", serverAddr, err) + continue + } + addrs = append(addrs, remoteAddr) + } + if len(addrs) == 0 { + return nil, fmt.Errorf("no usable addresses resolved for endpoint %s", exitNode.Endpoint) + } + return addrs, nil +} + // TriggerHolePunch sends an immediate hole punch packet to all configured exit nodes // This is useful for triggering hole punching on demand without waiting for the interval func (m *Manager) TriggerHolePunch() error { m.mu.Lock() + if m.disabled { + m.mu.Unlock() + return fmt.Errorf("hole punching is disabled") + } + if len(m.exitNodes) == 0 { m.mu.Unlock() return fmt.Errorf("no exit nodes configured") @@ -291,32 +353,25 @@ func (m *Manager) TriggerHolePunch() error { // Send hole punch to all exit nodes successCount := 0 for _, exitNode := range currentExitNodes { - var host string - var err error - if len(m.publicDNS) > 0 { - host, err = util.ResolveDomainUpstream(exitNode.Endpoint, m.publicDNS) - } else { - host, err = util.ResolveDomain(exitNode.Endpoint) - } + remoteAddrs, err := m.resolveExitNodeAddrs(exitNode) if err != nil { logger.Warn("Failed to resolve endpoint %s: %v", exitNode.Endpoint, err) continue } - serverAddr := net.JoinHostPort(host, strconv.Itoa(int(exitNode.RelayPort))) - remoteAddr, err := net.ResolveUDPAddr("udp", serverAddr) - if err != nil { - logger.Error("Failed to resolve UDP address %s: %v", serverAddr, err) - continue + sentAny := false + for _, remoteAddr := range remoteAddrs { + if err := m.sendHolePunch(remoteAddr, exitNode.PublicKey); err != nil { + logger.Warn("Failed to send on-demand hole punch to %s: %v", remoteAddr, err) + continue + } + sentAny = true } - if err := m.sendHolePunch(remoteAddr, exitNode.PublicKey); err != nil { - logger.Warn("Failed to send on-demand hole punch to %s: %v", exitNode.Endpoint, err) - continue + if sentAny { + logger.Debug("Sent on-demand hole punch to %s", exitNode.Endpoint) + successCount++ } - - logger.Debug("Sent on-demand hole punch to %s", exitNode.Endpoint) - successCount++ } if successCount == 0 { @@ -331,6 +386,12 @@ func (m *Manager) TriggerHolePunch() error { func (m *Manager) StartMultipleExitNodes(exitNodes []ExitNode) error { m.mu.Lock() + if m.disabled { + m.mu.Unlock() + logger.Debug("Hole punching is disabled, ignoring start request") + return fmt.Errorf("hole punching is disabled") + } + if m.running { m.mu.Unlock() logger.Debug("UDP hole punch already running, skipping new request") @@ -359,6 +420,12 @@ func (m *Manager) StartMultipleExitNodes(exitNodes []ExitNode) error { func (m *Manager) Start() error { m.mu.Lock() + if m.disabled { + m.mu.Unlock() + logger.Debug("Hole punching is disabled, ignoring start request") + return fmt.Errorf("hole punching is disabled") + } + if m.running { m.mu.Unlock() logger.Debug("UDP hole punch already running") @@ -408,31 +475,20 @@ func (m *Manager) runMultipleExitNodes() { var resolvedNodes []resolvedExitNode for _, exitNode := range currentExitNodes { - var host string - var err error - if len(m.publicDNS) > 0 { - host, err = util.ResolveDomainUpstream(exitNode.Endpoint, m.publicDNS) - } else { - host, err = util.ResolveDomain(exitNode.Endpoint) - } + remoteAddrs, err := m.resolveExitNodeAddrs(exitNode) if err != nil { logger.Warn("Failed to resolve endpoint %s: %v", exitNode.Endpoint, err) continue } - serverAddr := net.JoinHostPort(host, strconv.Itoa(int(exitNode.RelayPort))) - remoteAddr, err := net.ResolveUDPAddr("udp", serverAddr) - if err != nil { - logger.Error("Failed to resolve UDP address %s: %v", serverAddr, err) - continue + for _, remoteAddr := range remoteAddrs { + resolvedNodes = append(resolvedNodes, resolvedExitNode{ + remoteAddr: remoteAddr, + publicKey: exitNode.PublicKey, + endpointName: exitNode.Endpoint, + }) + logger.Debug("Resolved exit node: %s -> %s", exitNode.Endpoint, remoteAddr.String()) } - - resolvedNodes = append(resolvedNodes, resolvedExitNode{ - remoteAddr: remoteAddr, - publicKey: exitNode.PublicKey, - endpointName: exitNode.Endpoint, - }) - logger.Debug("Resolved exit node: %s -> %s", exitNode.Endpoint, remoteAddr.String()) } return resolvedNodes } diff --git a/util/util.go b/util/util.go index 0ce5dee..a92dbf9 100644 --- a/util/util.go +++ b/util/util.go @@ -15,18 +15,16 @@ import ( "golang.zx2c4.com/wireguard/device" ) -func ResolveDomainUpstream(domain string, publicDNS []string) (string, error) { - // trim whitespace +// splitDomainHostPort strips a protocol prefix/trailing slash from domain and +// separates it into host and port (port may be ""). If host is already a +// literal IP address (v4 or v6, brackets stripped), literalIP is non-nil and +// resolution can be skipped entirely. +func splitDomainHostPort(domain string) (host, port string, literalIP net.IP) { domain = strings.TrimSpace(domain) - - // Remove any protocol prefix if present (do this first, before splitting host/port) domain = strings.TrimPrefix(domain, "http://") domain = strings.TrimPrefix(domain, "https://") - - // if there are any trailing slashes, remove them domain = strings.TrimSuffix(domain, "/") - // Check if there's a port in the domain host, port, err := net.SplitHostPort(domain) if err != nil { // No port found, use the domain as is @@ -38,138 +36,143 @@ func ResolveDomainUpstream(domain string, publicDNS []string) (string, error) { // For IPv6, the host from SplitHostPort will already have brackets stripped // but if there was no port, we need to handle bracketed IPv6 addresses cleanHost := strings.TrimPrefix(strings.TrimSuffix(host, "]"), "[") - if ip := net.ParseIP(cleanHost); ip != nil { - // It's already an IP address, no need to resolve - ipAddr := ip.String() + return host, port, net.ParseIP(cleanHost) +} + +// resolveIPs looks up every address (all families) for host, preferring the +// given upstream DNS servers (each queried directly over UDP) when provided. +// If every upstream server is unreachable - e.g. the only configured/system +// DNS server is only reachable over an address family this process's own +// socket path doesn't currently have a route for (IPv6-only mobile networks +// commonly hand out IPv6-only resolvers) - this falls back to the platform's +// own resolver, which routes independently of our socket path and reliably +// works even then. See https://github.com/fosrl/android/issues/42 and +// https://github.com/fosrl/pangolin/issues/3471. +func resolveIPs(host string, publicDNS []string) ([]net.IP, error) { + if len(publicDNS) == 0 { + return net.LookupIP(host) + } + + var lastErr error + for _, server := range publicDNS { + // Ensure the upstream DNS address has a port + dnsAddr := server + if _, _, err := net.SplitHostPort(dnsAddr); err != nil { + // No port specified, default to 53 + dnsAddr = net.JoinHostPort(server, "53") + } + + resolver := &net.Resolver{ + PreferGo: true, + Dial: func(ctx context.Context, network, address string) (net.Conn, error) { + d := net.Dialer{} + return d.DialContext(ctx, "udp", dnsAddr) + }, + } + ips, err := resolver.LookupIP(context.Background(), "ip", host) + if err == nil { + return ips, nil + } + lastErr = err + } + + if ips, err := net.LookupIP(host); err == nil { + logger.Debug("All upstream DNS servers failed to resolve %s (%v), falling back to platform resolver", host, lastErr) + return ips, nil + } + + return nil, fmt.Errorf("DNS lookup failed using all upstream servers: %v", lastErr) +} + +// pickAddr chooses a single address from ips, preferring IPv4 for +// backward-compatible callers that only ever use one address (e.g. a +// WireGuard peer endpoint). Returns "" if ips is empty. +func pickAddr(ips []net.IP) string { + for _, ip := range ips { + if ipv4 := ip.To4(); ipv4 != nil { + return ipv4.String() + } + } + if len(ips) == 0 { + return "" + } + return ips[0].String() +} + +func ResolveDomainUpstream(domain string, publicDNS []string) (string, error) { + host, port, literalIP := splitDomainHostPort(domain) + if literalIP != nil { if port != "" { - return net.JoinHostPort(ipAddr, port), nil + return net.JoinHostPort(literalIP.String(), port), nil } - return ipAddr, nil + return literalIP.String(), nil } - // Lookup IP addresses using the upstream DNS servers if provided - var ips []net.IP - if len(publicDNS) > 0 { - var lastErr error - for _, server := range publicDNS { - // Ensure the upstream DNS address has a port - dnsAddr := server - if _, _, err := net.SplitHostPort(dnsAddr); err != nil { - // No port specified, default to 53 - dnsAddr = net.JoinHostPort(server, "53") - } - - resolver := &net.Resolver{ - PreferGo: true, - Dial: func(ctx context.Context, network, address string) (net.Conn, error) { - d := net.Dialer{} - return d.DialContext(ctx, "udp", dnsAddr) - }, - } - ips, lastErr = resolver.LookupIP(context.Background(), "ip", host) - if lastErr == nil { - break - } - } - if lastErr != nil { - return "", fmt.Errorf("DNS lookup failed using all upstream servers: %v", lastErr) - } - } else { - ips, err = net.LookupIP(host) - if err != nil { - return "", fmt.Errorf("DNS lookup failed: %v", err) - } + ips, err := resolveIPs(host, publicDNS) + if err != nil { + return "", err } - if len(ips) == 0 { return "", fmt.Errorf("no IP addresses found for domain %s", host) } - // Get the first IPv4 address if available - var ipAddr string - for _, ip := range ips { - if ipv4 := ip.To4(); ipv4 != nil { - ipAddr = ipv4.String() - break - } - } - - // If no IPv4 found, use the first IP (might be IPv6) - if ipAddr == "" { - ipAddr = ips[0].String() - } - - // Add port back if it existed + ipAddr := pickAddr(ips) if port != "" { ipAddr = net.JoinHostPort(ipAddr, port) } - return ipAddr, nil } - func ResolveDomain(domain string) (string, error) { - // trim whitespace - domain = strings.TrimSpace(domain) + return ResolveDomainUpstream(domain, nil) +} - // Remove any protocol prefix if present (do this first, before splitting host/port) - domain = strings.TrimPrefix(domain, "http://") - domain = strings.TrimPrefix(domain, "https://") - - // if there are any trailing slashes, remove them - domain = strings.TrimSuffix(domain, "/") - - // Check if there's a port in the domain - host, port, err := net.SplitHostPort(domain) - if err != nil { - // No port found, use the domain as is - host = domain - port = "" - } - - // Check if host is already an IP address (IPv4 or IPv6) - // For IPv6, the host from SplitHostPort will already have brackets stripped - // but if there was no port, we need to handle bracketed IPv6 addresses - cleanHost := strings.TrimPrefix(strings.TrimSuffix(host, "]"), "[") - if ip := net.ParseIP(cleanHost); ip != nil { - // It's already an IP address, no need to resolve - ipAddr := ip.String() +// ResolveDomainAllUpstream resolves domain to every candidate address (all +// families, deduplicated), each formatted as "ip:port" (or bare ip if domain +// had no port). Unlike ResolveDomainUpstream, which collapses to a single +// IPv4-preferred address, this lets a caller that can try more than one +// candidate (e.g. UDP hole punching) reach the destination over whichever +// address family the local network path actually has a route for, instead of +// always preferring an IPv4 address that may be completely unreachable (e.g. +// on an IPv6-only/NAT64 network). See +// https://github.com/fosrl/olm/issues/108. +func ResolveDomainAllUpstream(domain string, publicDNS []string) ([]string, error) { + host, port, literalIP := splitDomainHostPort(domain) + if literalIP != nil { if port != "" { - return net.JoinHostPort(ipAddr, port), nil + return []string{net.JoinHostPort(literalIP.String(), port)}, nil } - return ipAddr, nil + return []string{literalIP.String()}, nil } - // Lookup IP addresses - ips, err := net.LookupIP(host) + ips, err := resolveIPs(host, publicDNS) if err != nil { - return "", fmt.Errorf("DNS lookup failed: %v", err) + return nil, err } - if len(ips) == 0 { - return "", fmt.Errorf("no IP addresses found for domain %s", host) + return nil, fmt.Errorf("no IP addresses found for domain %s", host) } - // Get the first IPv4 address if available - var ipAddr string + seen := make(map[string]bool, len(ips)) + results := make([]string, 0, len(ips)) for _, ip := range ips { - if ipv4 := ip.To4(); ipv4 != nil { - ipAddr = ipv4.String() - break + s := ip.String() + if seen[s] { + continue } + seen[s] = true + if port != "" { + s = net.JoinHostPort(s, port) + } + results = append(results, s) } + return results, nil +} - // If no IPv4 found, use the first IP (might be IPv6) - if ipAddr == "" { - ipAddr = ips[0].String() - } - - // Add port back if it existed - if port != "" { - ipAddr = net.JoinHostPort(ipAddr, port) - } - - return ipAddr, nil +// ResolveDomainAll is ResolveDomainAllUpstream using only the system/platform +// resolver (no explicit upstream DNS servers). +func ResolveDomainAll(domain string) ([]string, error) { + return ResolveDomainAllUpstream(domain, nil) } func ParseLogLevel(level string) logger.LogLevel {