diff --git a/network/route.go b/network/route.go index 8aae063e..767878a1 100644 --- a/network/route.go +++ b/network/route.go @@ -1,6 +1,7 @@ package network import ( + "errors" "fmt" "net" "os/exec" @@ -59,69 +60,112 @@ func LinuxAddRoute(destination string, gateway string, interfaceName string) err return nil } - // Parse destination CIDR + metric := linuxDefaultTunnelRouteMetric + if gateway == "" && interfaceName != "" { + if m, overlap, err := linuxTunnelRouteMetric(destination, interfaceName); err != nil { + logger.Warn("Failed to check local subnet overlap for %s: %v; using high metric", destination, err) + metric = linuxOverlapTunnelRouteMetric + } else { + metric = m + if overlap { + logger.Warn("Remote subnet %s overlaps a local LAN subnet; adding tunnel route with metric %d", destination, metric) + } + } + } + _, ipNet, err := net.ParseCIDR(destination) if err != nil { return fmt.Errorf("invalid destination address: %v", err) } - // Create route - route := &netlink.Route{ - Dst: ipNet, - } + route := &netlink.Route{Dst: ipNet, Priority: metric} if gateway != "" { - // Route with specific gateway gw := net.ParseIP(gateway) if gw == nil { return fmt.Errorf("invalid gateway address: %s", gateway) } route.Gw = gw - logger.Info("Adding route to %s via gateway %s", destination, gateway) + logger.Info("Adding route to %s via gateway %s (metric %d)", destination, gateway, metric) } else if interfaceName != "" { - // Route via interface link, err := netlink.LinkByName(interfaceName) if err != nil { return fmt.Errorf("failed to get interface %s: %v", interfaceName, err) } route.LinkIndex = link.Attrs().Index - logger.Info("Adding route to %s via interface %s", destination, interfaceName) + logger.Info("Adding route to %s via interface %s (metric %d)", destination, interfaceName, metric) } else { return fmt.Errorf("either gateway or interface must be specified") } - // Add the route if err := netlink.RouteAdd(route); err != nil { return fmt.Errorf("failed to add route: %v", err) } - return nil } -func LinuxRemoveRoute(destination string) error { +func LinuxRemoveRoute(destination string, interfaceName ...string) error { if runtime.GOOS != "linux" { return nil } - // Parse destination CIDR _, ipNet, err := net.ParseCIDR(destination) if err != nil { return fmt.Errorf("invalid destination address: %v", err) } - // Create route to delete - route := &netlink.Route{ - Dst: ipNet, + routes, err := netlink.RouteList(nil, linuxRouteFamily(ipNet)) + if err != nil { + return fmt.Errorf("failed to list routes: %w", err) } - logger.Info("Removing route to %s", destination) + targetLinkIndex := 0 + if len(interfaceName) > 0 && interfaceName[0] != "" { + link, err := netlink.LinkByName(interfaceName[0]) + if err != nil { + return fmt.Errorf("failed to get interface %s: %w", interfaceName[0], err) + } + targetLinkIndex = link.Attrs().Index + } - // Delete the route - if err := netlink.RouteDel(route); err != nil { - return fmt.Errorf("failed to delete route: %v", err) + logger.Info("Removing route to %s", destination) + var matches []netlink.Route + linkIndexes := make(map[int]struct{}) + for _, route := range routes { + if route.Dst == nil || route.Dst.String() != ipNet.String() { + continue + } + if route.Priority != linuxDefaultTunnelRouteMetric && route.Priority != linuxOverlapTunnelRouteMetric { + continue + } + if targetLinkIndex != 0 && route.LinkIndex != targetLinkIndex { + continue + } + matches = append(matches, route) + linkIndexes[route.LinkIndex] = struct{}{} + } + if targetLinkIndex == 0 && len(linkIndexes) > 1 { + return fmt.Errorf("multiple tunnel routes to %s found; interface name is required", destination) } - return nil + var delErr error + removed := 0 + for _, route := range matches { + del := &netlink.Route{ + Dst: route.Dst, + LinkIndex: route.LinkIndex, + Priority: route.Priority, + } + if err := netlink.RouteDel(del); err != nil { + delErr = errors.Join(delErr, fmt.Errorf("failed to delete route: %w", err)) + continue + } + removed++ + } + if removed == 0 && delErr == nil { + return fmt.Errorf("tunnel route to %s not found", destination) + } + return delErr } // addRouteForServerIP adds an OS-specific route for the server IP @@ -243,7 +287,7 @@ func AddRoutes(remoteSubnets []string, interfaceName string) error { } // removeRoutesForRemoteSubnets removes routes for each subnet in RemoteSubnets -func RemoveRoutes(remoteSubnets []string) error { +func RemoveRoutes(remoteSubnets []string, interfaceName ...string) error { if len(remoteSubnets) == 0 { return nil } @@ -271,7 +315,7 @@ func RemoveRoutes(remoteSubnets []string) error { logger.Error("Failed to remove Windows route for subnet %s: %v", subnet, err) } case "linux": - if err := LinuxRemoveRoute(subnet); err != nil { + if err := LinuxRemoveRoute(subnet, interfaceName...); err != nil { logger.Error("Failed to remove Linux route for subnet %s: %v", subnet, err) } case "android", "ios": diff --git a/network/route_overlap.go b/network/route_overlap.go new file mode 100644 index 00000000..b5d3c026 --- /dev/null +++ b/network/route_overlap.go @@ -0,0 +1,84 @@ +package network + +import ( + "fmt" + "net" + "syscall" +) + +const ( + linuxDefaultTunnelRouteMetric = 50 // NetworkManager VPN route priority + linuxOverlapTunnelRouteMetric = 1025 // above typical LAN (100/600) and systemd-networkd (1024) +) + +func subnetsOverlap(a, b *net.IPNet) bool { + a4, b4 := a.IP.To4(), b.IP.To4() + if a4 == nil || b4 == nil { + return false + } + onesA, bitsA := a.Mask.Size() + onesB, bitsB := b.Mask.Size() + if bitsA != 32 || bitsB != 32 { + return false + } + minOnes := onesA + if onesB < minOnes { + minOnes = onesB + } + mask := net.CIDRMask(minOnes, 32) + return a4.Mask(mask).Equal(b4.Mask(mask)) +} + +func linuxTunnelRouteMetric(remoteSubnet, excludeIface string) (int, bool, error) { + localNets, err := localIPv4Subnets(excludeIface) + if err != nil { + return linuxDefaultTunnelRouteMetric, false, err + } + return metricForRemoteSubnet(remoteSubnet, localNets) +} + +func metricForRemoteSubnet(remoteSubnet string, localSubnets []*net.IPNet) (int, bool, error) { + _, remoteNet, err := net.ParseCIDR(remoteSubnet) + if err != nil { + return 0, false, fmt.Errorf("invalid remote subnet %s: %w", remoteSubnet, err) + } + for _, localNet := range localSubnets { + if subnetsOverlap(remoteNet, localNet) { + return linuxOverlapTunnelRouteMetric, true, nil + } + } + return linuxDefaultTunnelRouteMetric, false, nil +} + +func linuxRouteFamily(ipNet *net.IPNet) int { + if ipNet.IP.To4() == nil { + return syscall.AF_INET6 + } + return syscall.AF_INET +} + +func localIPv4Subnets(excludeIface string) ([]*net.IPNet, error) { + ifaces, err := net.Interfaces() + if err != nil { + return nil, err + } + + var subnets []*net.IPNet + for _, iface := range ifaces { + if iface.Name == excludeIface || iface.Flags&net.FlagLoopback != 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + return nil, err + } + for _, addr := range addrs { + ipNet, ok := addr.(*net.IPNet) + if !ok || ipNet.IP.To4() == nil { + continue + } + subnets = append(subnets, ipNet) + } + } + return subnets, nil +} diff --git a/network/route_overlap_test.go b/network/route_overlap_test.go new file mode 100644 index 00000000..21ffb9d7 --- /dev/null +++ b/network/route_overlap_test.go @@ -0,0 +1,58 @@ +package network + +import ( + "net" + "syscall" + "testing" +) + +func TestSubnetsOverlap(t *testing.T) { + tests := []struct { + a, b string + want bool + }{ + {"10.200.4.0/23", "10.200.4.0/23", true}, + {"10.200.4.0/23", "10.200.4.0/24", true}, + {"10.200.4.0/24", "10.200.5.0/24", false}, + {"192.168.1.0/24", "10.0.0.0/8", false}, + } + for _, tt := range tests { + _, a, _ := net.ParseCIDR(tt.a) + _, b, _ := net.ParseCIDR(tt.b) + if got := subnetsOverlap(a, b); got != tt.want { + t.Errorf("subnetsOverlap(%s, %s) = %v, want %v", tt.a, tt.b, got, tt.want) + } + } +} + +func TestMetricForRemoteSubnet(t *testing.T) { + local := mustCIDR(t, "10.200.4.0/23") + + metric, overlap, err := metricForRemoteSubnet("10.200.4.0/23", []*net.IPNet{local}) + if err != nil || !overlap || metric != linuxOverlapTunnelRouteMetric { + t.Fatalf("overlap: metric=%d overlap=%v err=%v", metric, overlap, err) + } + + metric, overlap, err = metricForRemoteSubnet("10.10.0.0/16", []*net.IPNet{local}) + if err != nil || overlap || metric != linuxDefaultTunnelRouteMetric { + t.Fatalf("no overlap: metric=%d overlap=%v err=%v", metric, overlap, err) + } +} + +func TestLinuxRouteFamily(t *testing.T) { + if family := linuxRouteFamily(mustCIDR(t, "10.200.4.0/23")); family != syscall.AF_INET { + t.Fatalf("IPv4 family = %d, want %d", family, syscall.AF_INET) + } + if family := linuxRouteFamily(mustCIDR(t, "fd00::/64")); family != syscall.AF_INET6 { + t.Fatalf("IPv6 family = %d, want %d", family, syscall.AF_INET6) + } +} + +func mustCIDR(t *testing.T, cidr string) *net.IPNet { + t.Helper() + _, ipNet, err := net.ParseCIDR(cidr) + if err != nil { + t.Fatal(err) + } + return ipNet +} diff --git a/newt/handlers.go b/newt/handlers.go index 0a448f14..55c7f4e4 100644 --- a/newt/handlers.go +++ b/newt/handlers.go @@ -430,7 +430,7 @@ func (n *Newt) registerHandlers(ctx context.Context) { } if n.config.UseNativeMainInterface { - if err := network.RemoveRoutes(data.Subnets); err != nil { + if err := network.RemoveRoutes(data.Subnets, n.config.NativeMainInterfaceName); err != nil { logger.Warn("Failed to remove routes for subnets: %v", err) } } diff --git a/newt/tunnel.go b/newt/tunnel.go index b396cce8..fdacbf7b 100644 --- a/newt/tunnel.go +++ b/newt/tunnel.go @@ -24,7 +24,7 @@ func (n *Newt) updateRemoteExitNodeSubnets(subnets []string) { } } if len(toRemove) > 0 { - if err := network.RemoveRoutes(toRemove); err != nil { + if err := network.RemoveRoutes(toRemove, n.config.NativeMainInterfaceName); err != nil { logger.Warn("Failed to remove old subnet routes: %v", err) } } @@ -88,7 +88,7 @@ func (n *Newt) closeWgTunnel() { } toRemove = append(toRemove, n.activeRemoteSubnets...) if len(toRemove) > 0 { - if err := network.RemoveRoutes(toRemove); err != nil { + if err := network.RemoveRoutes(toRemove, n.config.NativeMainInterfaceName); err != nil { logger.Warn("Failed to remove native main tunnel routes: %v", err) } }