Merge pull request #457 from fosrl/dev

1.18.0
This commit is contained in:
Owen Schwartz
2026-09-29 14:49:23 -04:00
committed by GitHub
12 changed files with 779 additions and 147 deletions
+93 -37
View File
@@ -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
}
+37
View File
@@ -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")
}
}
+27
View File
@@ -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()
+19 -1
View File
@@ -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
}
}
}
+144
View File
@@ -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)
}
}
+3
View File
@@ -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})
}
+187
View File
@@ -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, "")
+8
View File
@@ -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
}
+88
View File
@@ -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)
+37
View File
@@ -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()
+24
View File
@@ -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
View File
@@ -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 {