mirror of
https://github.com/fosrl/newt.git
synced 2026-10-04 11:49:06 +02:00
Route watch working
This commit is contained in:
+64
-27
@@ -9,6 +9,7 @@ import (
|
||||
"runtime"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
"golang.org/x/sys/windows"
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
)
|
||||
|
||||
@@ -109,11 +110,62 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
// (0.0.0.0/1 + 128.0.0.0/1 on the tunnel) is installed still resolves to the
|
||||
// physical default route rather than back into the tunnel.
|
||||
func WindowsAddBypassRoute(destIP string, tunnelInterface string) error {
|
||||
_, err := windowsEnsureBypassRoute(destIP, tunnelInterface)
|
||||
return err
|
||||
}
|
||||
|
||||
// windowsEnsureBypassRoute makes the /32 route to destIP match the current
|
||||
// physical path (see windowsBypassNextHop), adding it if missing and
|
||||
// replacing it if it points somewhere else. Returns whether the routing table
|
||||
// was changed.
|
||||
func windowsEnsureBypassRoute(destIP string, tunnelInterface string) (bool, error) {
|
||||
addr, err := netip.ParseAddr(destIP)
|
||||
if err != nil || !addr.Is4() {
|
||||
return false, fmt.Errorf("invalid IPv4 destination address: %s", destIP)
|
||||
}
|
||||
host := netip.PrefixFrom(addr, addr.BitLen())
|
||||
|
||||
routes, err := winipcfg.GetIPForwardTable2(windows.AF_INET)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid destination address: %v", err)
|
||||
return false, fmt.Errorf("failed to get route table: %v", err)
|
||||
}
|
||||
|
||||
best, err := windowsBypassNextHop(routes, addr, tunnelInterface)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
var existing *winipcfg.MibIPforwardRow2
|
||||
for i := range routes {
|
||||
if routes[i].DestinationPrefix.Prefix() == host {
|
||||
existing = &routes[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if existing != nil {
|
||||
if existing.InterfaceLUID == best.InterfaceLUID && existing.NextHop.Addr() == best.NextHop.Addr() {
|
||||
return false, nil
|
||||
}
|
||||
if err := existing.Delete(); err != nil {
|
||||
return false, fmt.Errorf("failed to remove stale bypass route to %s: %v", destIP, err)
|
||||
}
|
||||
}
|
||||
|
||||
logger.Info("Setting bypass route to %s via %s (interface LUID %v)", destIP, best.NextHop.Addr(), best.InterfaceLUID)
|
||||
|
||||
if err := best.InterfaceLUID.AddRoute(host, best.NextHop.Addr(), 0); err != nil {
|
||||
return false, fmt.Errorf("failed to add bypass route: %v", err)
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// windowsBypassNextHop picks the route Windows would use to reach addr if
|
||||
// neither the tunnel nor our own bypass route existed: the longest prefix
|
||||
// containing addr, lowest effective metric (route + interface metric) first,
|
||||
// skipping routes on tunnelInterface, on disconnected interfaces, and the /32
|
||||
// route to addr itself.
|
||||
func windowsBypassNextHop(routes []winipcfg.MibIPforwardRow2, addr netip.Addr, tunnelInterface string) (*winipcfg.MibIPforwardRow2, error) {
|
||||
var tunnelLUID winipcfg.LUID
|
||||
hasTunnelLUID := false
|
||||
if tunnelInterface != "" {
|
||||
@@ -124,46 +176,31 @@ func WindowsAddBypassRoute(destIP string, tunnelInterface string) error {
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
var bestMetric uint32
|
||||
for i := range routes {
|
||||
route := &routes[i]
|
||||
prefix := route.DestinationPrefix.Prefix()
|
||||
if !prefix.Contains(addr) {
|
||||
if !prefix.Contains(addr) || prefix.Bits() == addr.BitLen() {
|
||||
continue
|
||||
}
|
||||
if hasTunnelLUID && route.InterfaceLUID == tunnelLUID {
|
||||
continue
|
||||
}
|
||||
if prefix.Bits() > bestBits || (prefix.Bits() == bestBits && best != nil && route.Metric < best.Metric) {
|
||||
bestBits = prefix.Bits()
|
||||
best = route
|
||||
iface, err := route.InterfaceLUID.IPInterface(windows.AF_INET)
|
||||
if err != nil || !iface.Connected {
|
||||
continue
|
||||
}
|
||||
metric := route.Metric + iface.Metric
|
||||
if prefix.Bits() > bestBits || (prefix.Bits() == bestBits && metric < bestMetric) {
|
||||
best, bestBits, bestMetric = route, prefix.Bits(), metric
|
||||
}
|
||||
}
|
||||
if best == nil {
|
||||
return fmt.Errorf("no route found to %s outside tunnel interface %q", destIP, tunnelInterface)
|
||||
return nil, fmt.Errorf("%w: %s", ErrNoPhysicalRoute, addr)
|
||||
}
|
||||
|
||||
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
|
||||
return best, nil
|
||||
}
|
||||
|
||||
// WindowsRemoveBypassRoute removes a route previously added by
|
||||
|
||||
Reference in New Issue
Block a user