mirror of
https://github.com/fosrl/newt.git
synced 2026-09-30 01:39:08 +02:00
+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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
package netstack2
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// With an exit-node rule (0.0.0.0/0) installed, traffic addressed to newt's own
|
||||
// tunnel IP - e.g. olm's connection-status probe to the wgtester - must be left
|
||||
// for the main stack rather than matched by the catch-all and proxied out.
|
||||
func TestHandleIncomingPacket_LocalTunnelAddressBypassesCatchAll(t *testing.T) {
|
||||
ph, err := NewProxyHandler(ProxyHandlerOptions{EnableICMP: true, MTU: 1500})
|
||||
if err != nil {
|
||||
t.Fatalf("NewProxyHandler: %v", err)
|
||||
}
|
||||
if err := ph.Initialize(noopNotification{}); err != nil {
|
||||
t.Fatalf("Initialize: %v", err)
|
||||
}
|
||||
defer ph.Close()
|
||||
|
||||
tunnelIP := netip.MustParseAddr("100.90.128.1")
|
||||
clientIP := netip.MustParseAddr("100.90.128.5")
|
||||
internetIP := netip.MustParseAddr("203.0.113.50")
|
||||
|
||||
ph.SetLocalAddresses([]netip.Addr{tunnelIP})
|
||||
ph.AddSubnetRule(SubnetRule{
|
||||
SourcePrefix: netip.MustParsePrefix("100.90.128.0/24"),
|
||||
DestPrefix: netip.MustParsePrefix("0.0.0.0/0"),
|
||||
})
|
||||
|
||||
if ph.HandleIncomingPacket(buildICMPEchoRequest(t, clientIP, tunnelIP)) {
|
||||
t.Error("packet to the local tunnel IP was proxied; expected it to be left for the main stack")
|
||||
}
|
||||
if !ph.HandleIncomingPacket(buildICMPEchoRequest(t, clientIP, internetIP)) {
|
||||
t.Error("packet to an internet address should still match the exit-node rule")
|
||||
}
|
||||
}
|
||||
@@ -136,6 +136,13 @@ type ProxyHandler struct {
|
||||
accessLogger *AccessLogger // Access logger for tracking sessions
|
||||
httpRequestLogger *HTTPRequestLogger // HTTP request logger for proxied HTTP/HTTPS requests
|
||||
blocked atomic.Bool // when true, all new connections are dropped
|
||||
|
||||
// localAddrs are the addresses owned by the main netstack (the tunnel IP).
|
||||
// Traffic addressed to them terminates on the main stack (wgtester, SSH,
|
||||
// ...) and must never be proxied out to the host network, even when a
|
||||
// catch-all rule such as an exit node's 0.0.0.0/0 would otherwise match it.
|
||||
// Written once during setup, before any packet is processed.
|
||||
localAddrs map[netip.Addr]struct{}
|
||||
}
|
||||
|
||||
// ProxyHandlerOptions configures the proxy handler
|
||||
@@ -494,6 +501,19 @@ func (p *ProxyHandler) Initialize(notifiable channel.Notification) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetLocalAddresses registers the addresses owned by the main netstack so
|
||||
// packets destined to them are left for the main stack instead of being proxied.
|
||||
// Must be called before the device starts processing packets.
|
||||
func (p *ProxyHandler) SetLocalAddresses(addrs []netip.Addr) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
p.localAddrs = make(map[netip.Addr]struct{}, len(addrs))
|
||||
for _, addr := range addrs {
|
||||
p.localAddrs[addr.Unmap()] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
// HandleIncomingPacket processes incoming packets and determines if they should
|
||||
// be injected into the proxy stack
|
||||
func (p *ProxyHandler) HandleIncomingPacket(packet []byte) bool {
|
||||
@@ -522,6 +542,13 @@ func (p *ProxyHandler) HandleIncomingPacket(packet []byte) bool {
|
||||
dstBytes := dstIP.As4()
|
||||
dstAddr := netip.AddrFrom4(dstBytes)
|
||||
|
||||
// Traffic for our own tunnel IP (e.g. the olm connection-status probe to the
|
||||
// wgtester, or SSH) is served by the main stack. Without this, an exit node's
|
||||
// 0.0.0.0/0 rule matches it and forwards it out to the host network instead.
|
||||
if _, isLocal := p.localAddrs[dstAddr]; isLocal {
|
||||
return false
|
||||
}
|
||||
|
||||
// Parse transport layer to get destination port
|
||||
var dstPort uint16
|
||||
protocol := ipv4Header.TransportProtocol()
|
||||
|
||||
@@ -175,13 +175,27 @@ func (sl *SubnetLookup) Match(srcIP, dstIP netip.Addr, port uint16, proto tcpip.
|
||||
continue
|
||||
}
|
||||
|
||||
// Supernets() yields longest-prefix-match first, then progressively
|
||||
// less specific. Once a more specific, non-catch-all destination
|
||||
// rule has been seen and rejected (wrong port/protocol), a
|
||||
// 0.0.0.0/0 (or ::/0) exit-node rule must not be allowed to rescue
|
||||
// it - "whole subnet" routing only applies when no more specific
|
||||
// resource covers this destination at all. Fallthrough between two
|
||||
// specific (non-catch-all) rules is intentional and unaffected.
|
||||
sawRejectedSpecificDest := false
|
||||
|
||||
// Step 2: Find all destination prefixes that contain dstIP
|
||||
// This is also O(log n) for each matching source prefix
|
||||
for _, rules := range destTriePtr.trie.Supernets(dstPrefix) {
|
||||
for destPrefix, rules := range destTriePtr.trie.Supernets(dstPrefix) {
|
||||
if rules == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
isCatchAll := destPrefix.Bits() == 0
|
||||
if isCatchAll && sawRejectedSpecificDest {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Step 3: Check each rule for ICMP and port restrictions
|
||||
for _, rule := range rules {
|
||||
// Handle ICMP before port range check — ICMP has no ports
|
||||
@@ -216,6 +230,10 @@ func (sl *SubnetLookup) Match(srcIP, dstIP netip.Addr, port uint16, proto tcpip.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !isCatchAll {
|
||||
sawRejectedSpecificDest = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
package netstack2
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
)
|
||||
|
||||
// clientPrefix is the shared SourcePrefix used across these tests, mirroring
|
||||
// how the server always assigns a /32 per client (server/lib/ip.ts).
|
||||
var clientPrefix = netip.MustParsePrefix("10.0.0.5/32")
|
||||
var clientIP = clientPrefix.Addr()
|
||||
|
||||
func TestMatch_SpecificResourceRejectsPort_DoesNotFallThroughToExitNode(t *testing.T) {
|
||||
sl := NewSubnetLookup()
|
||||
|
||||
// A specific /24 resource restricted to port 443 only.
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("192.168.1.0/24"),
|
||||
PortRanges: []PortRange{{Min: 443, Max: 443, Protocol: "tcp"}},
|
||||
ResourceId: 100,
|
||||
})
|
||||
|
||||
// A 0.0.0.0/0 exit-node rule with no port restriction.
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("0.0.0.0/0"),
|
||||
ResourceId: 999,
|
||||
})
|
||||
|
||||
inCIDR := netip.MustParseAddr("192.168.1.50")
|
||||
|
||||
// Disallowed port on the in-CIDR destination must be denied outright,
|
||||
// not rescued by the exit node's permissive catch-all.
|
||||
if rule := sl.Match(clientIP, inCIDR, 22, header.TCPProtocolNumber); rule != nil {
|
||||
t.Fatalf("expected deny for disallowed port via specific resource, got rule with ResourceId=%d", rule.ResourceId)
|
||||
}
|
||||
|
||||
// Allowed port on the in-CIDR destination must match the specific resource.
|
||||
rule := sl.Match(clientIP, inCIDR, 443, header.TCPProtocolNumber)
|
||||
if rule == nil {
|
||||
t.Fatal("expected match for allowed port on specific resource, got nil")
|
||||
}
|
||||
if rule.ResourceId != 100 {
|
||||
t.Fatalf("expected ResourceId=100 (specific resource), got %d", rule.ResourceId)
|
||||
}
|
||||
|
||||
// A destination outside the /24 has no specific resource covering it,
|
||||
// so the exit node must still catch it normally.
|
||||
outsideCIDR := netip.MustParseAddr("8.8.8.8")
|
||||
rule = sl.Match(clientIP, outsideCIDR, 22, header.TCPProtocolNumber)
|
||||
if rule == nil {
|
||||
t.Fatal("expected exit-node match for destination outside the specific resource, got nil")
|
||||
}
|
||||
if rule.ResourceId != 999 {
|
||||
t.Fatalf("expected ResourceId=999 (exit node), got %d", rule.ResourceId)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatch_NonCatchAllFallthroughStillWorks(t *testing.T) {
|
||||
sl := NewSubnetLookup()
|
||||
|
||||
// A very specific /32 limited to SSH only.
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("192.168.1.50/32"),
|
||||
PortRanges: []PortRange{{Min: 22, Max: 22, Protocol: "tcp"}},
|
||||
ResourceId: 1,
|
||||
})
|
||||
|
||||
// A broader /24 (non-catch-all) that allows HTTP.
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("192.168.1.0/24"),
|
||||
PortRanges: []PortRange{{Min: 80, Max: 80, Protocol: "tcp"}},
|
||||
ResourceId: 2,
|
||||
})
|
||||
|
||||
ip := netip.MustParseAddr("192.168.1.50")
|
||||
|
||||
// Port 80 doesn't match the /32's SSH-only rule, so it must still fall
|
||||
// through to the broader /24 rule that allows it (non-catch-all
|
||||
// fallthrough is preserved).
|
||||
rule := sl.Match(clientIP, ip, 80, header.TCPProtocolNumber)
|
||||
if rule == nil {
|
||||
t.Fatal("expected fallthrough match on broader /24 rule, got nil")
|
||||
}
|
||||
if rule.ResourceId != 2 {
|
||||
t.Fatalf("expected ResourceId=2 (broader /24 resource), got %d", rule.ResourceId)
|
||||
}
|
||||
|
||||
// Port 22 matches the /32 directly.
|
||||
rule = sl.Match(clientIP, ip, 22, header.TCPProtocolNumber)
|
||||
if rule == nil {
|
||||
t.Fatal("expected match on specific /32 rule, got nil")
|
||||
}
|
||||
if rule.ResourceId != 1 {
|
||||
t.Fatalf("expected ResourceId=1 (specific /32 resource), got %d", rule.ResourceId)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatch_ICMPDisabledOnSpecificResource_HardDeniesRegardlessOfExitNode(t *testing.T) {
|
||||
sl := NewSubnetLookup()
|
||||
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("192.168.1.0/24"),
|
||||
DisableIcmp: true,
|
||||
ResourceId: 100,
|
||||
})
|
||||
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("0.0.0.0/0"),
|
||||
ResourceId: 999,
|
||||
})
|
||||
|
||||
ip := netip.MustParseAddr("192.168.1.50")
|
||||
|
||||
if rule := sl.Match(clientIP, ip, 0, header.ICMPv4ProtocolNumber); rule != nil {
|
||||
t.Fatalf("expected ICMP deny on specific resource, got rule with ResourceId=%d", rule.ResourceId)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatch_ExitNodeOnlyMatchesWhenNoSpecificResourceCovers(t *testing.T) {
|
||||
sl := NewSubnetLookup()
|
||||
|
||||
sl.AddSubnet(SubnetRule{
|
||||
SourcePrefix: clientPrefix,
|
||||
DestPrefix: netip.MustParsePrefix("0.0.0.0/0"),
|
||||
ResourceId: 999,
|
||||
})
|
||||
|
||||
ip := netip.MustParseAddr("1.2.3.4")
|
||||
rule := sl.Match(clientIP, ip, 443, header.TCPProtocolNumber)
|
||||
if rule == nil {
|
||||
t.Fatal("expected exit-node match when no specific resource exists, got nil")
|
||||
}
|
||||
if rule.ResourceId != 999 {
|
||||
t.Fatalf("expected ResourceId=999, got %d", rule.ResourceId)
|
||||
}
|
||||
}
|
||||
@@ -139,6 +139,9 @@ func CreateNetTUNWithOptions(localAddresses, dnsServers []netip.Addr, mtu int, o
|
||||
dev.hasV6 = true
|
||||
}
|
||||
}
|
||||
// Packets to our own addresses belong to the main stack, not the proxy.
|
||||
dev.proxyHandler.SetLocalAddresses(localAddresses)
|
||||
|
||||
if dev.hasV4 {
|
||||
dev.stack.AddRoute(tcpip.Route{Destination: header.IPv4EmptySubnet, NIC: 1})
|
||||
}
|
||||
|
||||
@@ -209,6 +209,193 @@ func LinuxRemoveRoute(destination string, interfaceName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// LinuxAddBypassRoute adds an explicit /32 host route for destIP via
|
||||
// whatever gateway/interface the kernel currently uses to reach it, so a
|
||||
// broader route added afterward (e.g. a gateway/full-tunnel default route)
|
||||
// can never capture this destination - see AddBypassRouteForDestination.
|
||||
func LinuxAddBypassRoute(destIP string) error {
|
||||
if runtime.GOOS != "linux" {
|
||||
return nil
|
||||
}
|
||||
|
||||
ip := net.ParseIP(destIP)
|
||||
if ip == nil {
|
||||
return fmt.Errorf("invalid destination address: %s", destIP)
|
||||
}
|
||||
|
||||
routes, err := netlink.RouteGet(ip)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to look up current route to %s: %v", destIP, err)
|
||||
}
|
||||
if len(routes) == 0 {
|
||||
return fmt.Errorf("no route found to %s", destIP)
|
||||
}
|
||||
current := routes[0]
|
||||
|
||||
link, err := netlink.LinkByIndex(current.LinkIndex)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to resolve interface for route to %s: %v", destIP, err)
|
||||
}
|
||||
|
||||
route := &netlink.Route{
|
||||
Dst: &net.IPNet{IP: ip, Mask: net.CIDRMask(32, 32)},
|
||||
Gw: current.Gw,
|
||||
LinkIndex: link.Attrs().Index,
|
||||
}
|
||||
|
||||
logger.Info("Adding bypass route to %s via %s (interface %s)", destIP, current.Gw, link.Attrs().Name)
|
||||
|
||||
if err := netlink.RouteAdd(route); err != nil {
|
||||
return fmt.Errorf("failed to add bypass route to %s: %v", destIP, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// LinuxRemoveBypassRoute removes a route previously added by
|
||||
// LinuxAddBypassRoute. It deliberately does not re-derive the route via
|
||||
// RouteGet - by the time this runs, our own /32 bypass route is the most
|
||||
// specific match for destIP and RouteGet would just find itself - so it
|
||||
// instead deletes by destination alone.
|
||||
func LinuxRemoveBypassRoute(destIP string) error {
|
||||
if runtime.GOOS != "linux" {
|
||||
return nil
|
||||
}
|
||||
|
||||
ip := net.ParseIP(destIP)
|
||||
if ip == nil {
|
||||
return fmt.Errorf("invalid destination address: %s", destIP)
|
||||
}
|
||||
|
||||
route := &netlink.Route{
|
||||
Dst: &net.IPNet{IP: ip, Mask: net.CIDRMask(32, 32)},
|
||||
}
|
||||
|
||||
if err := netlink.RouteDel(route); err != nil {
|
||||
return fmt.Errorf("failed to remove bypass route to %s: %v", destIP, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DarwinAddBypassRoute adds an explicit /32 host route for destIP via
|
||||
// whatever gateway/interface the kernel currently uses to reach it (parsed
|
||||
// from `route -n get`), so a broader route added afterward can never capture
|
||||
// this destination - see AddBypassRouteForDestination.
|
||||
func DarwinAddBypassRoute(destIP string) error {
|
||||
if runtime.GOOS != "darwin" {
|
||||
return nil
|
||||
}
|
||||
if NativeConfigDisabled {
|
||||
return nil
|
||||
}
|
||||
|
||||
cmd := exec.Command("route", "-n", "get", destIP)
|
||||
logger.Info("Running command: %v", cmd)
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("route get command failed: %v, output: %s", err, out)
|
||||
}
|
||||
|
||||
var gateway, iface string
|
||||
for _, line := range strings.Split(string(out), "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
switch {
|
||||
case strings.HasPrefix(line, "gateway:"):
|
||||
gateway = strings.TrimSpace(strings.TrimPrefix(line, "gateway:"))
|
||||
case strings.HasPrefix(line, "interface:"):
|
||||
iface = strings.TrimSpace(strings.TrimPrefix(line, "interface:"))
|
||||
}
|
||||
}
|
||||
if gateway == "" && iface == "" {
|
||||
return fmt.Errorf("could not determine current route to %s from `route get` output: %s", destIP, out)
|
||||
}
|
||||
|
||||
return DarwinAddRouteWithSource(destIP+"/32", gateway, iface, "")
|
||||
}
|
||||
|
||||
// DarwinRemoveBypassRoute removes a route previously added by
|
||||
// DarwinAddBypassRoute.
|
||||
func DarwinRemoveBypassRoute(destIP string) error {
|
||||
if runtime.GOOS != "darwin" {
|
||||
return nil
|
||||
}
|
||||
return DarwinRemoveRoute(destIP + "/32")
|
||||
}
|
||||
|
||||
// AddGatewayDefaultRoute installs the OS-level "route everything" equivalent
|
||||
// for a full-tunnel/gateway peer. NetworkSettings is always populated first
|
||||
// (regardless of GOOS - mobile packet-tunnel providers read it independent
|
||||
// of platform, see AddRouteForServerIPWithSource), via an IsDefault included
|
||||
// route. On desktop platforms, where olm manages the OS routing table
|
||||
// directly, this then also installs the standard wg-quick split-default-route
|
||||
// technique (0.0.0.0/1 + 128.0.0.0/1) instead of a literal 0.0.0.0/0, so the
|
||||
// host's real default route is never replaced or raced with - it is only
|
||||
// outranked by two strictly more-specific halves. PreferLocalRoutes (if set)
|
||||
// still applies to these routes exactly as it does to any other tunnel
|
||||
// route, so an overlapping local/LAN route continues to win even in gateway
|
||||
// mode.
|
||||
func AddGatewayDefaultRoute(interfaceName, sourceIP string) error {
|
||||
AddIPv4IncludedRoute(IPv4Route{DestinationAddress: "0.0.0.0", SubnetMask: "0.0.0.0", IsDefault: true})
|
||||
|
||||
if runtime.GOOS == "android" || runtime.GOOS == "ios" {
|
||||
return nil
|
||||
}
|
||||
return AddRoutesWithSource([]string{"0.0.0.0/1", "128.0.0.0/1"}, interfaceName, sourceIP)
|
||||
}
|
||||
|
||||
// RemoveGatewayDefaultRoute reverses AddGatewayDefaultRoute.
|
||||
func RemoveGatewayDefaultRoute(interfaceName string) error {
|
||||
RemoveIPv4IncludedRoute(IPv4Route{DestinationAddress: "0.0.0.0", SubnetMask: "0.0.0.0", IsDefault: true})
|
||||
|
||||
if runtime.GOOS == "android" || runtime.GOOS == "ios" {
|
||||
return nil
|
||||
}
|
||||
return RemoveRoutes([]string{"0.0.0.0/1", "128.0.0.0/1"}, interfaceName)
|
||||
}
|
||||
|
||||
// AddBypassRouteForDestination installs an explicit /32 host route for destIP
|
||||
// using whatever gateway/interface the OS routing table currently uses to
|
||||
// reach it - i.e. the physical/original path, not the tunnel. It must be
|
||||
// called BEFORE AddGatewayDefaultRoute so the destination's own path is
|
||||
// pinned down first and can never be captured by the more general gateway
|
||||
// route. This is the same technique wg-quick uses (set_endpoint_direct_route)
|
||||
// to keep a WireGuard peer's own UDP traffic from being captured by the
|
||||
// gateway route it is itself responsible for installing.
|
||||
//
|
||||
// NetworkSettings is always populated (an excluded route, for mobile
|
||||
// packet-tunnel providers), regardless of GOOS; the OS routing table is only
|
||||
// touched on desktop platforms, which have no equivalent of
|
||||
// NEIPv4Settings.excludedRoutes/VpnService.Builder.excludeRoute.
|
||||
func AddBypassRouteForDestination(destIP string) error {
|
||||
AddIPv4ExcludedRoute(IPv4Route{DestinationAddress: destIP, SubnetMask: "255.255.255.255"})
|
||||
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
return LinuxAddBypassRoute(destIP)
|
||||
case "darwin":
|
||||
return DarwinAddBypassRoute(destIP)
|
||||
case "windows":
|
||||
return WindowsAddBypassRoute(destIP)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveBypassRouteForDestination reverses AddBypassRouteForDestination.
|
||||
func RemoveBypassRouteForDestination(destIP string) error {
|
||||
RemoveIPv4ExcludedRoute(IPv4Route{DestinationAddress: destIP, SubnetMask: "255.255.255.255"})
|
||||
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
return LinuxRemoveBypassRoute(destIP)
|
||||
case "darwin":
|
||||
return DarwinRemoveBypassRoute(destIP)
|
||||
case "windows":
|
||||
return WindowsRemoveBypassRoute(destIP)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// addRouteForServerIP adds an OS-specific route for the server IP
|
||||
func AddRouteForServerIP(serverIP, interfaceName string) error {
|
||||
return AddRouteForServerIPWithSource(serverIP, interfaceName, "")
|
||||
|
||||
@@ -9,3 +9,11 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
func WindowsRemoveRoute(destination string, interfaceName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func WindowsAddBypassRoute(destIP string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func WindowsRemoveBypassRoute(destIP string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -100,6 +100,94 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
// WindowsAddBypassRoute adds an explicit /32 host route for destIP via
|
||||
// whatever gateway/interface the OS routing table currently uses to reach
|
||||
// it, so a broader route added afterward (e.g. a gateway/full-tunnel default
|
||||
// route) can never capture this destination - see
|
||||
// network.AddBypassRouteForDestination.
|
||||
func WindowsAddBypassRoute(destIP string) error {
|
||||
addr, err := netip.ParseAddr(destIP)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid destination address: %v", err)
|
||||
}
|
||||
|
||||
var family winipcfg.AddressFamily
|
||||
if addr.Is4() {
|
||||
family = 2 // AF_INET
|
||||
} else {
|
||||
family = 23 // AF_INET6
|
||||
}
|
||||
|
||||
routes, err := winipcfg.GetIPForwardTable2(family)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get route table: %v", err)
|
||||
}
|
||||
|
||||
var best *winipcfg.MibIPforwardRow2
|
||||
bestBits := -1
|
||||
for i := range routes {
|
||||
route := &routes[i]
|
||||
prefix := route.DestinationPrefix.Prefix()
|
||||
if !prefix.Contains(addr) {
|
||||
continue
|
||||
}
|
||||
if prefix.Bits() > bestBits || (prefix.Bits() == bestBits && best != nil && route.Metric < best.Metric) {
|
||||
bestBits = prefix.Bits()
|
||||
best = route
|
||||
}
|
||||
}
|
||||
if best == nil {
|
||||
return fmt.Errorf("no route found to %s", destIP)
|
||||
}
|
||||
|
||||
prefix := netip.PrefixFrom(addr, addr.BitLen())
|
||||
logger.Info("Adding bypass route to %s via interface LUID %v", destIP, best.InterfaceLUID)
|
||||
|
||||
if err := best.InterfaceLUID.AddRoute(prefix, best.NextHop.Addr(), 0); err != nil {
|
||||
return fmt.Errorf("failed to add bypass route: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// WindowsRemoveBypassRoute removes a route previously added by
|
||||
// WindowsAddBypassRoute. It deliberately does not re-derive the route via a
|
||||
// longest-prefix-match lookup - by the time this runs, our own /32 bypass
|
||||
// route is the most specific match for destIP and the lookup would just find
|
||||
// itself - so it instead deletes by exact destination prefix alone.
|
||||
func WindowsRemoveBypassRoute(destIP string) error {
|
||||
addr, err := netip.ParseAddr(destIP)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid destination address: %v", err)
|
||||
}
|
||||
prefix := netip.PrefixFrom(addr, addr.BitLen())
|
||||
|
||||
var family winipcfg.AddressFamily
|
||||
if addr.Is4() {
|
||||
family = 2
|
||||
} else {
|
||||
family = 23
|
||||
}
|
||||
|
||||
routes, err := winipcfg.GetIPForwardTable2(family)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get route table: %v", err)
|
||||
}
|
||||
|
||||
for _, route := range routes {
|
||||
if route.DestinationPrefix.Prefix() != prefix {
|
||||
continue
|
||||
}
|
||||
logger.Info("Removing bypass route to %s on interface LUID %v", destIP, route.InterfaceLUID)
|
||||
if err := route.Delete(); err != nil {
|
||||
return fmt.Errorf("failed to delete bypass route: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("bypass route to %s not found", destIP)
|
||||
}
|
||||
|
||||
func WindowsRemoveRoute(destination string, interfaceName string) error {
|
||||
// Parse destination CIDR
|
||||
_, ipNet, err := net.ParseCIDR(destination)
|
||||
|
||||
@@ -181,6 +181,43 @@ func SetIPv4ExcludedRoutes(routes []IPv4Route) {
|
||||
logger.Info("Set IPv4 excluded routes: %d routes", len(routes))
|
||||
}
|
||||
|
||||
// AddIPv4ExcludedRoute adds a single excluded route, e.g. so mobile
|
||||
// (iOS/Android) packet-tunnel providers keep a specific destination (a site's
|
||||
// live endpoint, the control-plane server) out of an otherwise-broad included
|
||||
// route such as a full-tunnel/gateway default route. Mirrors
|
||||
// AddIPv4IncludedRoute's dedup-by-equality behavior.
|
||||
func AddIPv4ExcludedRoute(route IPv4Route) {
|
||||
networkSettingsMutex.Lock()
|
||||
defer networkSettingsMutex.Unlock()
|
||||
|
||||
for _, r := range networkSettings.IPv4ExcludedRoutes {
|
||||
if r == route {
|
||||
logger.Info("IPv4 excluded route already exists: %+v", route)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
networkSettings.IPv4ExcludedRoutes = append(networkSettings.IPv4ExcludedRoutes, route)
|
||||
incrementor++
|
||||
logger.Info("Added IPv4 excluded route: %+v", route)
|
||||
}
|
||||
|
||||
// RemoveIPv4ExcludedRoute reverses AddIPv4ExcludedRoute.
|
||||
func RemoveIPv4ExcludedRoute(route IPv4Route) {
|
||||
networkSettingsMutex.Lock()
|
||||
defer networkSettingsMutex.Unlock()
|
||||
routes := networkSettings.IPv4ExcludedRoutes
|
||||
for i, r := range routes {
|
||||
if r == route {
|
||||
networkSettings.IPv4ExcludedRoutes = append(routes[:i], routes[i+1:]...)
|
||||
incrementor++
|
||||
logger.Info("Removed IPv4 excluded route: %+v", route)
|
||||
return
|
||||
}
|
||||
}
|
||||
logger.Info("IPv4 excluded route not found for removal: %+v", route)
|
||||
}
|
||||
|
||||
// SetIPv6Settings sets IPv6 addresses and network prefixes
|
||||
func SetIPv6Settings(addresses []string, networkPrefixes []string) {
|
||||
networkSettingsMutex.Lock()
|
||||
|
||||
@@ -34,6 +34,13 @@ const (
|
||||
fmtErrParsingTargetData = "Error parsing target data: %v"
|
||||
)
|
||||
|
||||
// NewtErrorData represents a warning/error message sent down from the server,
|
||||
// e.g. when it was unable to complete part of the site's registration.
|
||||
type NewtErrorData struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
func (n *Newt) registerHandlers(ctx context.Context) {
|
||||
//TODO: MOVE MORE OF THESE HANDLERS TO STANDALONE FUNCTIONS IN THE DATA.GO AND CONNECT.GO FILES
|
||||
|
||||
@@ -41,6 +48,23 @@ func (n *Newt) registerHandlers(ctx context.Context) {
|
||||
n.handleConnect(ctx, msg)
|
||||
})
|
||||
|
||||
n.client.RegisterHandler("newt/error", func(msg websocket.WSMessage) {
|
||||
var errorData NewtErrorData
|
||||
|
||||
jsonData, err := json.Marshal(msg.Data)
|
||||
if err != nil {
|
||||
logger.Error(fmtErrMarshaling, err)
|
||||
return
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(jsonData, &errorData); err != nil {
|
||||
logger.Error("Error unmarshaling newt error data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
logger.Warn("Site warning (code: %s): %s", errorData.Code, errorData.Message)
|
||||
})
|
||||
|
||||
n.client.RegisterHandler("newt/wg/reconnect", func(msg websocket.WSMessage) {
|
||||
logger.Info("Received reconnect message")
|
||||
if n.wgData.PublicKey != "" {
|
||||
|
||||
+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