mirror of
https://github.com/fosrl/newt.git
synced 2026-09-30 01:39:08 +02:00
Resolve exit node endpoints over whichever address family is reachable
Fixes fosrl/android#42 Fixes fosrl/pangolin#3471 Fixes fosrl/olm#108
This commit is contained in:
+93
-37
@@ -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
|
||||
}
|
||||
|
||||
+112
-109
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user