mirror of
https://github.com/fosrl/newt.git
synced 2026-10-07 05:09:07 +02:00
Route watch working
This commit is contained in:
@@ -0,0 +1,116 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
package network
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func ipCmd(t *testing.T, args ...string) string {
|
||||||
|
t.Helper()
|
||||||
|
out, err := exec.Command("ip", args...).CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ip %v: %v %s", args, err, out)
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(string(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitFor(t *testing.T, what string, cond func() bool) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(8 * time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if cond() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("timed out waiting for %s", what)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBypassReconcileNetns exercises bypass routes against a real kernel
|
||||||
|
// routing table, with a gateway route (0.0.0.0/1 + 128.0.0.0/1) on a fake
|
||||||
|
// tunnel interface: adding a bypass while the gateway route is installed, and
|
||||||
|
// WatchRouteChanges + ReconcileBypassRoute following the physical network
|
||||||
|
// through an interface switch, a default route move, and going offline. It
|
||||||
|
// modifies the routing table, so it only runs inside a throwaway network
|
||||||
|
// namespace:
|
||||||
|
//
|
||||||
|
// NETNS_TEST=1 unshare -rn go test ./network -run TestBypassReconcileNetns -v
|
||||||
|
func TestBypassReconcileNetns(t *testing.T) {
|
||||||
|
if os.Getenv("NETNS_TEST") == "" {
|
||||||
|
t.Skip()
|
||||||
|
}
|
||||||
|
ipCmd(t, "link", "add", "phys0", "type", "dummy")
|
||||||
|
ipCmd(t, "link", "set", "phys0", "up")
|
||||||
|
ipCmd(t, "addr", "add", "10.0.0.2/24", "dev", "phys0")
|
||||||
|
ipCmd(t, "route", "add", "default", "via", "10.0.0.1", "dev", "phys0")
|
||||||
|
ipCmd(t, "link", "add", "tun0", "type", "dummy")
|
||||||
|
ipCmd(t, "link", "set", "tun0", "up")
|
||||||
|
ipCmd(t, "addr", "add", "100.64.0.2/32", "dev", "tun0")
|
||||||
|
ipCmd(t, "route", "add", "0.0.0.0/1", "dev", "tun0")
|
||||||
|
ipCmd(t, "route", "add", "128.0.0.0/1", "dev", "tun0")
|
||||||
|
|
||||||
|
// Added while the gateway route is already installed.
|
||||||
|
if err := LinuxAddBypassRoute("1.1.1.1", "tun0"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Logf("initial: %s", ipCmd(t, "route", "get", "1.1.1.1"))
|
||||||
|
|
||||||
|
var calls atomic.Int32
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
if err := WatchRouteChanges(ctx, func() {
|
||||||
|
calls.Add(1)
|
||||||
|
changed, err := ReconcileBypassRoute("1.1.1.1", "tun0")
|
||||||
|
t.Logf("reconcile #%d: changed=%v err=%v", calls.Load(), changed, err)
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
routeVia := func(want string) func() bool {
|
||||||
|
return func() bool {
|
||||||
|
out, _ := exec.Command("ip", "route", "show", "1.1.1.1/32").CombinedOutput()
|
||||||
|
return strings.Contains(string(out), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1. Physical interface disappears; a new network comes up elsewhere.
|
||||||
|
ipCmd(t, "link", "del", "phys0")
|
||||||
|
ipCmd(t, "link", "add", "phys1", "type", "dummy")
|
||||||
|
ipCmd(t, "link", "set", "phys1", "up")
|
||||||
|
ipCmd(t, "addr", "add", "192.168.5.2/24", "dev", "phys1")
|
||||||
|
ipCmd(t, "route", "add", "default", "via", "192.168.5.1", "dev", "phys1")
|
||||||
|
waitFor(t, "bypass via phys1", routeVia("via 192.168.5.1 dev phys1"))
|
||||||
|
t.Logf("after interface switch: %s", ipCmd(t, "route", "get", "1.1.1.1"))
|
||||||
|
|
||||||
|
// 2. Default moves to another interface while the old one stays up.
|
||||||
|
ipCmd(t, "link", "add", "phys2", "type", "dummy")
|
||||||
|
ipCmd(t, "link", "set", "phys2", "up")
|
||||||
|
ipCmd(t, "addr", "add", "10.9.0.2/24", "dev", "phys2")
|
||||||
|
ipCmd(t, "route", "replace", "default", "via", "10.9.0.1", "dev", "phys2")
|
||||||
|
waitFor(t, "bypass via phys2", routeVia("via 10.9.0.1 dev phys2"))
|
||||||
|
t.Logf("after default moved: %s", ipCmd(t, "route", "get", "1.1.1.1"))
|
||||||
|
|
||||||
|
// 3. No feedback loop: our own route changes settle to a no-op.
|
||||||
|
time.Sleep(2500 * time.Millisecond)
|
||||||
|
settled := calls.Load()
|
||||||
|
time.Sleep(3 * time.Second)
|
||||||
|
if calls.Load() != settled {
|
||||||
|
t.Fatalf("reconcile keeps firing: %d -> %d", settled, calls.Load())
|
||||||
|
}
|
||||||
|
t.Logf("settled after %d reconcile calls", settled)
|
||||||
|
|
||||||
|
// 4. Offline: reported as ErrNoPhysicalRoute.
|
||||||
|
cancel()
|
||||||
|
time.Sleep(200 * time.Millisecond)
|
||||||
|
ipCmd(t, "route", "del", "default")
|
||||||
|
if _, err := ReconcileBypassRoute("1.1.1.1", "tun0"); !errors.Is(err, ErrNoPhysicalRoute) {
|
||||||
|
t.Fatalf("expected ErrNoPhysicalRoute, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+251
-89
@@ -1,6 +1,7 @@
|
|||||||
package network
|
package network
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
@@ -218,63 +219,39 @@ func LinuxRemoveRoute(destination string, interfaceName string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ErrNoPhysicalRoute is returned (wrapped) when a bypass route can't be
|
||||||
|
// installed or reconciled because there is currently no route to the
|
||||||
|
// destination outside the tunnel - typically because the host is offline or
|
||||||
|
// between networks. Callers can treat it as transient: the next reconcile
|
||||||
|
// after the network comes back installs the route.
|
||||||
|
var ErrNoPhysicalRoute = errors.New("no route outside the tunnel")
|
||||||
|
|
||||||
// LinuxAddBypassRoute adds an explicit /32 host route for destIP via
|
// LinuxAddBypassRoute adds an explicit /32 host route for destIP via
|
||||||
// whatever gateway/interface the kernel currently uses to reach it, so a
|
// whatever gateway/interface the kernel currently uses to reach it, so a
|
||||||
// broader route added afterward (e.g. a gateway/full-tunnel default route)
|
// broader route added afterward (e.g. a gateway/full-tunnel default route)
|
||||||
// can never capture this destination - see AddBypassRouteForDestination.
|
// can never capture this destination - see AddBypassRouteForDestination.
|
||||||
// Routes on tunnelInterface are ignored when picking that path - see
|
// Routes on tunnelInterface are ignored when picking that path - see
|
||||||
// linuxBypassNextHop.
|
// linuxBypassNextHop. Replaces any existing /32 route to destIP.
|
||||||
func LinuxAddBypassRoute(destIP string, tunnelInterface string) error {
|
func LinuxAddBypassRoute(destIP string, tunnelInterface string) error {
|
||||||
if runtime.GOOS != "linux" {
|
if runtime.GOOS != "linux" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
_, err := linuxEnsureBypassRoute(destIP, tunnelInterface)
|
||||||
ip := net.ParseIP(destIP)
|
return err
|
||||||
if ip == nil {
|
|
||||||
return fmt.Errorf("invalid destination address: %s", destIP)
|
|
||||||
}
|
|
||||||
|
|
||||||
gw, linkIndex, err := linuxBypassNextHop(ip, tunnelInterface)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
link, err := netlink.LinkByIndex(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: gw,
|
|
||||||
LinkIndex: link.Attrs().Index,
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.Info("Adding bypass route to %s via %s (interface %s)", destIP, 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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// linuxBypassNextHop returns the gateway and interface the kernel uses to
|
// linuxEnsureBypassRoute makes the /32 route to destIP match the current
|
||||||
// reach ip, ignoring any route on tunnelInterface. While a gateway route
|
// physical path (see linuxBypassNextHop), replacing or adding it as needed.
|
||||||
// (0.0.0.0/1 + 128.0.0.0/1 on the tunnel) is installed, the kernel's own
|
// Returns whether the routing table was changed.
|
||||||
// answer for any public address is the tunnel itself, which would make a
|
func linuxEnsureBypassRoute(destIP string, tunnelInterface string) (bool, error) {
|
||||||
// bypass route added at that point useless. In that case this falls back to
|
ip := net.ParseIP(destIP).To4()
|
||||||
// the most specific non-tunnel unicast route in the main table that contains
|
if ip == nil {
|
||||||
// ip - normally the untouched physical 0.0.0.0/0 default route, which the
|
return false, fmt.Errorf("invalid IPv4 destination address: %s", destIP)
|
||||||
// gateway route deliberately leaves in place. tunnelInterface may be "" to
|
|
||||||
// just use the kernel's answer.
|
|
||||||
func linuxBypassNextHop(ip net.IP, tunnelInterface string) (net.IP, int, error) {
|
|
||||||
routes, err := netlink.RouteGet(ip)
|
|
||||||
if err != nil {
|
|
||||||
return nil, 0, fmt.Errorf("failed to look up current route to %s: %v", ip, err)
|
|
||||||
}
|
}
|
||||||
if len(routes) == 0 {
|
|
||||||
return nil, 0, fmt.Errorf("no route found to %s", ip)
|
routes, err := netlink.RouteList(nil, familyV4)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("failed to list routes: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
tunnelIndex := -1
|
tunnelIndex := -1
|
||||||
@@ -283,20 +260,54 @@ func linuxBypassNextHop(ip net.IP, tunnelInterface string) (net.IP, int, error)
|
|||||||
tunnelIndex = link.Attrs().Index
|
tunnelIndex = link.Attrs().Index
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if routes[0].LinkIndex != tunnelIndex {
|
|
||||||
return routes[0].Gw, routes[0].LinkIndex, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
candidates, err := netlink.RouteList(nil, familyV4)
|
gw, linkIndex, err := linuxBypassNextHop(routes, ip, tunnelIndex)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, fmt.Errorf("failed to list routes: %v", err)
|
return false, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
for _, r := range routes {
|
||||||
|
if isHostRouteTo(r.Dst, ip) && r.LinkIndex == linkIndex && r.Gw.Equal(gw) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
link, err := netlink.LinkByIndex(linkIndex)
|
||||||
|
if err != nil {
|
||||||
|
return false, 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: gw,
|
||||||
|
LinkIndex: linkIndex,
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Info("Setting bypass route to %s via %s (interface %s)", destIP, gw, link.Attrs().Name)
|
||||||
|
|
||||||
|
if err := netlink.RouteReplace(route); err != nil {
|
||||||
|
return false, fmt.Errorf("failed to set bypass route to %s: %v", destIP, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// linuxBypassNextHop picks, from the main-table routes, the gateway and
|
||||||
|
// interface the kernel would use to reach ip if neither the tunnel nor our
|
||||||
|
// own bypass route existed: the most specific unicast route containing ip,
|
||||||
|
// lowest metric first, skipping any route on tunnelIndex and any /32 route to
|
||||||
|
// ip itself. While a gateway route (0.0.0.0/1 + 128.0.0.0/1 on the tunnel)
|
||||||
|
// is installed, that is normally the untouched physical 0.0.0.0/0 default
|
||||||
|
// route, which the gateway route deliberately leaves in place. Skipping our
|
||||||
|
// own /32 lets the same lookup tell whether an existing bypass route still
|
||||||
|
// matches the current physical path (see linuxEnsureBypassRoute). Policy
|
||||||
|
// routing rules are not considered.
|
||||||
|
func linuxBypassNextHop(routes []netlink.Route, ip net.IP, tunnelIndex int) (net.IP, int, error) {
|
||||||
var best *netlink.Route
|
var best *netlink.Route
|
||||||
bestBits := -1
|
bestBits := -1
|
||||||
for i := range candidates {
|
for i := range routes {
|
||||||
r := &candidates[i]
|
r := &routes[i]
|
||||||
if r.Type != rtnUnicast {
|
if r.Type != rtnUnicast || isHostRouteTo(r.Dst, ip) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
bits := 0
|
bits := 0
|
||||||
@@ -321,11 +332,20 @@ func linuxBypassNextHop(ip net.IP, tunnelInterface string) (net.IP, int, error)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if best == nil {
|
if best == nil {
|
||||||
return nil, 0, fmt.Errorf("no route to %s outside tunnel interface %s", ip, tunnelInterface)
|
return nil, 0, fmt.Errorf("%w: %s", ErrNoPhysicalRoute, ip)
|
||||||
}
|
}
|
||||||
return best.Gw, best.LinkIndex, nil
|
return best.Gw, best.LinkIndex, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isHostRouteTo reports whether dst is exactly ip/32.
|
||||||
|
func isHostRouteTo(dst *net.IPNet, ip net.IP) bool {
|
||||||
|
if dst == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
ones, bits := dst.Mask.Size()
|
||||||
|
return ones == 32 && bits == 32 && dst.IP.Equal(ip)
|
||||||
|
}
|
||||||
|
|
||||||
// LinuxRemoveBypassRoute removes a route previously added by
|
// LinuxRemoveBypassRoute removes a route previously added by
|
||||||
// LinuxAddBypassRoute. It deliberately does not re-derive the route via
|
// LinuxAddBypassRoute. It deliberately does not re-derive the route via
|
||||||
// RouteGet - by the time this runs, our own /32 bypass route is the most
|
// RouteGet - by the time this runs, our own /32 bypass route is the most
|
||||||
@@ -358,7 +378,7 @@ func LinuxRemoveBypassRoute(destIP string) error {
|
|||||||
// this destination - see AddBypassRouteForDestination. If that route is on
|
// this destination - see AddBypassRouteForDestination. If that route is on
|
||||||
// tunnelInterface - i.e. a gateway route (0.0.0.0/1 + 128.0.0.0/1) is already
|
// tunnelInterface - i.e. a gateway route (0.0.0.0/1 + 128.0.0.0/1) is already
|
||||||
// capturing destIP - the physical default route, which the gateway route
|
// capturing destIP - the physical default route, which the gateway route
|
||||||
// deliberately leaves in place, is used instead.
|
// deliberately leaves in place, is used instead (see darwinBypassNextHop).
|
||||||
func DarwinAddBypassRoute(destIP string, tunnelInterface string) error {
|
func DarwinAddBypassRoute(destIP string, tunnelInterface string) error {
|
||||||
if runtime.GOOS != "darwin" {
|
if runtime.GOOS != "darwin" {
|
||||||
return nil
|
return nil
|
||||||
@@ -366,47 +386,151 @@ func DarwinAddBypassRoute(destIP string, tunnelInterface string) error {
|
|||||||
if NativeConfigDisabled {
|
if NativeConfigDisabled {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
_, err := darwinEnsureBypassRoute(destIP, tunnelInterface)
|
||||||
gateway, iface, err := darwinRouteGet(destIP)
|
return err
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if tunnelInterface != "" && iface == tunnelInterface {
|
|
||||||
gateway, iface, err = darwinRouteGet("default")
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if iface == tunnelInterface {
|
|
||||||
return fmt.Errorf("no route to %s outside tunnel interface %s", destIP, tunnelInterface)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return DarwinAddRouteWithSource(destIP+"/32", gateway, iface, "")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// darwinRouteGet returns the gateway and interface `route -n get` reports for
|
// darwinEnsureBypassRoute makes the /32 route to destIP match the current
|
||||||
// destination (an address, or "default").
|
// physical path, adding it if missing and replacing it if it points
|
||||||
func darwinRouteGet(destination string) (gateway, iface string, err error) {
|
// somewhere else. Returns whether the routing table was changed.
|
||||||
cmd := exec.Command("route", "-n", "get", destination)
|
func darwinEnsureBypassRoute(destIP string, tunnelInterface string) (bool, error) {
|
||||||
logger.Info("Running command: %v", cmd)
|
ip := net.ParseIP(destIP).To4()
|
||||||
out, err := cmd.CombinedOutput()
|
if ip == nil {
|
||||||
if err != nil {
|
return false, fmt.Errorf("invalid IPv4 destination address: %s", destIP)
|
||||||
return "", "", fmt.Errorf("route get command failed: %v, output: %s", err, out)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, line := range strings.Split(string(out), "\n") {
|
current, err := darwinRouteGet(destIP)
|
||||||
line = strings.TrimSpace(line)
|
if err != nil {
|
||||||
switch {
|
return false, err
|
||||||
case strings.HasPrefix(line, "gateway:"):
|
}
|
||||||
gateway = strings.TrimSpace(strings.TrimPrefix(line, "gateway:"))
|
installed := current.isStaticHostRouteTo(ip)
|
||||||
case strings.HasPrefix(line, "interface:"):
|
|
||||||
iface = strings.TrimSpace(strings.TrimPrefix(line, "interface:"))
|
var gateway, iface string
|
||||||
|
if installed {
|
||||||
|
// `route get` now just finds our own route, so work the physical
|
||||||
|
// path out without it.
|
||||||
|
gateway, iface, err = darwinBypassNextHop(ip, tunnelInterface)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
if iface == current.iface && (gateway == "" || gateway == current.gateway) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err := DarwinRemoveRoute(destIP + "/32"); err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
gateway, iface = current.gateway, current.iface
|
||||||
|
if net.ParseIP(gateway) == nil {
|
||||||
|
// On-link (e.g. an ARP entry's link-layer "gateway"): route
|
||||||
|
// via the interface itself.
|
||||||
|
gateway = ""
|
||||||
|
}
|
||||||
|
if tunnelInterface != "" && iface == tunnelInterface {
|
||||||
|
gateway, iface, err = darwinBypassNextHop(ip, tunnelInterface)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if gateway == "" && iface == "" {
|
|
||||||
return "", "", fmt.Errorf("could not determine current route to %s from `route get` output: %s", destination, out)
|
if err := DarwinAddRouteWithSource(destIP+"/32", gateway, iface, ""); err != nil {
|
||||||
|
return false, err
|
||||||
}
|
}
|
||||||
return gateway, iface, nil
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// darwinBypassNextHop returns the physical path to ip without consulting
|
||||||
|
// any more specific route to it (ours, or the tunnel's): directly on-link
|
||||||
|
// through a non-tunnel interface whose subnet contains ip, otherwise the
|
||||||
|
// unscoped default route - which a gateway route (0.0.0.0/1 + 128.0.0.0/1)
|
||||||
|
// deliberately leaves in place. Other, more specific non-default routes are
|
||||||
|
// not considered.
|
||||||
|
func darwinBypassNextHop(ip net.IP, tunnelInterface string) (gateway, iface string, err error) {
|
||||||
|
if ifaces, err := net.Interfaces(); err == nil {
|
||||||
|
for _, ifc := range ifaces {
|
||||||
|
if ifc.Flags&net.FlagUp == 0 || ifc.Flags&net.FlagLoopback != 0 || ifc.Name == tunnelInterface {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
addrs, err := ifc.Addrs()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, addr := range addrs {
|
||||||
|
if ipNet, ok := addr.(*net.IPNet); ok && ipNet.IP.To4() != nil && ipNet.Contains(ip) {
|
||||||
|
return "", ifc.Name, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
def, err := darwinRouteGet("default")
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
if def.iface == tunnelInterface {
|
||||||
|
return "", "", fmt.Errorf("%w: %s", ErrNoPhysicalRoute, ip)
|
||||||
|
}
|
||||||
|
return def.gateway, def.iface, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// darwinRoute is the relevant part of `route -n get` output.
|
||||||
|
type darwinRoute struct {
|
||||||
|
destination, mask, gateway, iface string
|
||||||
|
flags []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// isStaticHostRouteTo reports whether r is a static /32 route to ip - i.e.
|
||||||
|
// one we added, as opposed to an ARP-cloned host entry or a broader route.
|
||||||
|
func (r darwinRoute) isStaticHostRouteTo(ip net.IP) bool {
|
||||||
|
if !net.ParseIP(r.destination).Equal(ip) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if r.mask != "" && r.mask != "255.255.255.255" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, f := range r.flags {
|
||||||
|
if f == "STATIC" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// darwinRouteGet returns what `route -n get` reports for destination (an
|
||||||
|
// address, or "default").
|
||||||
|
func darwinRouteGet(destination string) (darwinRoute, error) {
|
||||||
|
cmd := exec.Command("route", "-n", "get", destination)
|
||||||
|
logger.Debug("Running command: %v", cmd)
|
||||||
|
out, err := cmd.CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
return darwinRoute{}, fmt.Errorf("%w: route get %s failed: %v, output: %s", ErrNoPhysicalRoute, destination, err, out)
|
||||||
|
}
|
||||||
|
|
||||||
|
var r darwinRoute
|
||||||
|
for _, line := range strings.Split(string(out), "\n") {
|
||||||
|
key, value, ok := strings.Cut(strings.TrimSpace(line), ":")
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
switch key {
|
||||||
|
case "destination":
|
||||||
|
r.destination = value
|
||||||
|
case "mask":
|
||||||
|
r.mask = value
|
||||||
|
case "gateway":
|
||||||
|
r.gateway = value
|
||||||
|
case "interface":
|
||||||
|
r.iface = value
|
||||||
|
case "flags":
|
||||||
|
r.flags = strings.Split(strings.Trim(value, "<>"), ",")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if r.gateway == "" && r.iface == "" {
|
||||||
|
return darwinRoute{}, fmt.Errorf("could not determine current route to %s from `route get` output: %s", destination, out)
|
||||||
|
}
|
||||||
|
return r, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// DarwinRemoveBypassRoute removes a route previously added by
|
// DarwinRemoveBypassRoute removes a route previously added by
|
||||||
@@ -482,6 +606,44 @@ func AddBypassRouteForDestination(destIP string, tunnelInterface string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ReconcileBypassRoute makes sure the host route AddBypassRouteForDestination
|
||||||
|
// installed for destIP is still present and still follows the current
|
||||||
|
// physical path, re-adding or moving it if not. The OS drops these routes
|
||||||
|
// along with the interface or address they were on (e.g. switching Wi-Fi
|
||||||
|
// networks, or Wi-Fi to Ethernet), and if the default route moves to a
|
||||||
|
// different interface they would otherwise keep using the old one. Routes on
|
||||||
|
// tunnelInterface are ignored, as in AddBypassRouteForDestination. Returns
|
||||||
|
// whether anything changed; a no-op where ManagesHostRoutes is false.
|
||||||
|
// NetworkSettings is not touched - platforms that apply excluded routes from
|
||||||
|
// it track the physical network themselves.
|
||||||
|
func ReconcileBypassRoute(destIP string, tunnelInterface string) (bool, error) {
|
||||||
|
if !ManagesHostRoutes() {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "linux":
|
||||||
|
return linuxEnsureBypassRoute(destIP, tunnelInterface)
|
||||||
|
case "darwin":
|
||||||
|
return darwinEnsureBypassRoute(destIP, tunnelInterface)
|
||||||
|
case "windows":
|
||||||
|
return windowsEnsureBypassRoute(destIP, tunnelInterface)
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ManagesHostRoutes reports whether this package installs routes in the host
|
||||||
|
// OS routing table itself (desktop platforms), as opposed to only populating
|
||||||
|
// NetworkSettings for a mobile/NetworkExtension host app to apply.
|
||||||
|
func ManagesHostRoutes() bool {
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "linux", "windows":
|
||||||
|
return true
|
||||||
|
case "darwin":
|
||||||
|
return !NativeConfigDisabled
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// RemoveBypassRouteForDestination reverses AddBypassRouteForDestination.
|
// RemoveBypassRouteForDestination reverses AddBypassRouteForDestination.
|
||||||
func RemoveBypassRouteForDestination(destIP string) error {
|
func RemoveBypassRouteForDestination(destIP string) error {
|
||||||
RemoveIPv4ExcludedRoute(IPv4Route{DestinationAddress: destIP, SubnetMask: "255.255.255.255"})
|
RemoveIPv4ExcludedRoute(IPv4Route{DestinationAddress: destIP, SubnetMask: "255.255.255.255"})
|
||||||
|
|||||||
@@ -10,6 +10,10 @@ func WindowsRemoveRoute(destination string, interfaceName string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func windowsEnsureBypassRoute(destIP string, tunnelInterface string) (bool, error) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
func WindowsAddBypassRoute(destIP string, tunnelInterface string) error {
|
func WindowsAddBypassRoute(destIP string, tunnelInterface string) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
+64
-27
@@ -9,6 +9,7 @@ import (
|
|||||||
"runtime"
|
"runtime"
|
||||||
|
|
||||||
"github.com/fosrl/newt/logger"
|
"github.com/fosrl/newt/logger"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
"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
|
// (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.
|
// physical default route rather than back into the tunnel.
|
||||||
func WindowsAddBypassRoute(destIP string, tunnelInterface string) error {
|
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)
|
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 {
|
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
|
var tunnelLUID winipcfg.LUID
|
||||||
hasTunnelLUID := false
|
hasTunnelLUID := false
|
||||||
if tunnelInterface != "" {
|
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
|
var best *winipcfg.MibIPforwardRow2
|
||||||
bestBits := -1
|
bestBits := -1
|
||||||
|
var bestMetric uint32
|
||||||
for i := range routes {
|
for i := range routes {
|
||||||
route := &routes[i]
|
route := &routes[i]
|
||||||
prefix := route.DestinationPrefix.Prefix()
|
prefix := route.DestinationPrefix.Prefix()
|
||||||
if !prefix.Contains(addr) {
|
if !prefix.Contains(addr) || prefix.Bits() == addr.BitLen() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if hasTunnelLUID && route.InterfaceLUID == tunnelLUID {
|
if hasTunnelLUID && route.InterfaceLUID == tunnelLUID {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if prefix.Bits() > bestBits || (prefix.Bits() == bestBits && best != nil && route.Metric < best.Metric) {
|
iface, err := route.InterfaceLUID.IPInterface(windows.AF_INET)
|
||||||
bestBits = prefix.Bits()
|
if err != nil || !iface.Connected {
|
||||||
best = route
|
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 {
|
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)
|
||||||
}
|
}
|
||||||
|
return best, nil
|
||||||
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
|
// WindowsRemoveBypassRoute removes a route previously added by
|
||||||
|
|||||||
@@ -0,0 +1,75 @@
|
|||||||
|
package network
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// routeChangeSettle is how long the routing table must be quiet after a
|
||||||
|
// change before onChange runs, so the burst of updates a single network
|
||||||
|
// switch produces (link, addresses, routes) is handled once, after it has
|
||||||
|
// finished.
|
||||||
|
routeChangeSettle = time.Second
|
||||||
|
// routeChangeMaxDelay bounds how long continuous churn can postpone
|
||||||
|
// onChange.
|
||||||
|
routeChangeMaxDelay = 5 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
// WatchRouteChanges calls onChange whenever the host's routing table,
|
||||||
|
// interface addresses or link state change - debounced, so one network switch
|
||||||
|
// results in one call - until ctx is done. Only events from the OS are
|
||||||
|
// watched: changes this process makes itself (e.g. adding bypass routes) also
|
||||||
|
// trigger it, so onChange must be idempotent. A no-op where ManagesHostRoutes
|
||||||
|
// is false, since there are no host routes to keep up to date there.
|
||||||
|
func WatchRouteChanges(ctx context.Context, onChange func()) error {
|
||||||
|
if !ManagesHostRoutes() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
trigger := make(chan struct{}, 1)
|
||||||
|
notify := func() {
|
||||||
|
select {
|
||||||
|
case trigger <- struct{}{}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := watchRouteEvents(ctx, notify); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-trigger:
|
||||||
|
}
|
||||||
|
|
||||||
|
settle := time.NewTimer(routeChangeSettle)
|
||||||
|
maxDelay := time.NewTimer(routeChangeMaxDelay)
|
||||||
|
wait:
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
settle.Stop()
|
||||||
|
maxDelay.Stop()
|
||||||
|
return
|
||||||
|
case <-trigger:
|
||||||
|
settle.Reset(routeChangeSettle)
|
||||||
|
case <-settle.C:
|
||||||
|
break wait
|
||||||
|
case <-maxDelay.C:
|
||||||
|
break wait
|
||||||
|
}
|
||||||
|
}
|
||||||
|
settle.Stop()
|
||||||
|
maxDelay.Stop()
|
||||||
|
|
||||||
|
onChange()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
package network
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"golang.org/x/net/route"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchRouteEvents calls notify for routing socket messages that can change
|
||||||
|
// the physical path: routes added/removed/changed, interface addresses
|
||||||
|
// added/removed, and interface state changes. ARP/NDP neighbor entries
|
||||||
|
// (cloned host routes) and RTM_GET replies - which include those to our own
|
||||||
|
// `route get` calls - are ignored, as they would otherwise trigger constantly.
|
||||||
|
func watchRouteEvents(ctx context.Context, notify func()) error {
|
||||||
|
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Non-blocking so os.NewFile registers it with the runtime poller and
|
||||||
|
// Close interrupts a pending Read.
|
||||||
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
|
unix.Close(fd)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
f := os.NewFile(uintptr(fd), "route")
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
f.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
for {
|
||||||
|
n, err := f.Read(buf)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if isRelevantRouteMessage(buf[:n]) {
|
||||||
|
notify()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func isRelevantRouteMessage(b []byte) bool {
|
||||||
|
if len(b) < 4 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch int(b[3]) { // rtm_type, common to every routing message header
|
||||||
|
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
|
||||||
|
return true
|
||||||
|
case unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_CHANGE:
|
||||||
|
msgs, err := route.ParseRIB(route.RIBTypeRoute, b)
|
||||||
|
if err != nil || len(msgs) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if m, ok := msgs[0].(*route.RouteMessage); ok && m.Flags&(unix.RTF_LLINFO|unix.RTF_WASCLONED) != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
package network
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fosrl/newt/logger"
|
||||||
|
"github.com/vishvananda/netlink"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchRouteEvents calls notify for every netlink route, address and link
|
||||||
|
// update. netlink ends a subscription on any receive error (e.g. the socket
|
||||||
|
// buffer overflowing during a burst), so it resubscribes when that happens,
|
||||||
|
// notifying once in case the lost messages mattered.
|
||||||
|
func watchRouteEvents(ctx context.Context, notify func()) error {
|
||||||
|
done, routes, addrs, links, err := subscribeRouteEvents()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
stopRouteEvents(done, routes, addrs, links)
|
||||||
|
return
|
||||||
|
case _, ok := <-routes:
|
||||||
|
if ok {
|
||||||
|
notify()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
case _, ok := <-addrs:
|
||||||
|
if ok {
|
||||||
|
notify()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
case _, ok := <-links:
|
||||||
|
if ok {
|
||||||
|
notify()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A subscription ended without ctx being done.
|
||||||
|
logger.Warn("Route change subscription ended, resubscribing")
|
||||||
|
stopRouteEvents(done, routes, addrs, links)
|
||||||
|
notify()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
}
|
||||||
|
done, routes, addrs, links, err = subscribeRouteEvents()
|
||||||
|
if err == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
logger.Warn("Failed to resubscribe to route changes: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func subscribeRouteEvents() (chan struct{}, chan netlink.RouteUpdate, chan netlink.AddrUpdate, chan netlink.LinkUpdate, error) {
|
||||||
|
done := make(chan struct{})
|
||||||
|
routes := make(chan netlink.RouteUpdate, 64)
|
||||||
|
addrs := make(chan netlink.AddrUpdate, 64)
|
||||||
|
links := make(chan netlink.LinkUpdate, 64)
|
||||||
|
|
||||||
|
if err := netlink.RouteSubscribe(routes, done); err != nil {
|
||||||
|
close(done)
|
||||||
|
return nil, nil, nil, nil, err
|
||||||
|
}
|
||||||
|
if err := netlink.AddrSubscribe(addrs, done); err != nil {
|
||||||
|
stopRouteEvents(done, routes, nil, nil)
|
||||||
|
return nil, nil, nil, nil, err
|
||||||
|
}
|
||||||
|
if err := netlink.LinkSubscribe(links, done); err != nil {
|
||||||
|
stopRouteEvents(done, routes, addrs, nil)
|
||||||
|
return nil, nil, nil, nil, err
|
||||||
|
}
|
||||||
|
return done, routes, addrs, links, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopRouteEvents ends the subscriptions and drains their channels in the
|
||||||
|
// background: netlink's receive goroutines block sending to a full channel
|
||||||
|
// and only exit (closing it) once they observe the closed socket.
|
||||||
|
func stopRouteEvents(done chan struct{}, routes chan netlink.RouteUpdate, addrs chan netlink.AddrUpdate, links chan netlink.LinkUpdate) {
|
||||||
|
close(done)
|
||||||
|
go func() {
|
||||||
|
if routes != nil {
|
||||||
|
for range routes {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if addrs != nil {
|
||||||
|
for range addrs {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if links != nil {
|
||||||
|
for range links {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
//go:build !linux && !darwin && !windows
|
||||||
|
|
||||||
|
package network
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
func watchRouteEvents(ctx context.Context, notify func()) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
package network
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchRouteEvents calls notify for every route, interface and unicast
|
||||||
|
// address change notification.
|
||||||
|
func watchRouteEvents(ctx context.Context, notify func()) error {
|
||||||
|
routeCb, err := winipcfg.RegisterRouteChangeCallback(func(winipcfg.MibNotificationType, *winipcfg.MibIPforwardRow2) {
|
||||||
|
notify()
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ifaceCb, err := winipcfg.RegisterInterfaceChangeCallback(func(winipcfg.MibNotificationType, *winipcfg.MibIPInterfaceRow) {
|
||||||
|
notify()
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
_ = routeCb.Unregister()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
addrCb, err := winipcfg.RegisterUnicastAddressChangeCallback(func(winipcfg.MibNotificationType, *winipcfg.MibUnicastIPAddressRow) {
|
||||||
|
notify()
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
_ = routeCb.Unregister()
|
||||||
|
_ = ifaceCb.Unregister()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
_ = routeCb.Unregister()
|
||||||
|
_ = ifaceCb.Unregister()
|
||||||
|
_ = addrCb.Unregister()
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user