diff --git a/feature/routecheck/routecheck.go b/feature/routecheck/routecheck.go index c22332f81..d09bc09ca 100644 --- a/feature/routecheck/routecheck.go +++ b/feature/routecheck/routecheck.go @@ -13,12 +13,15 @@ package routecheck import ( + "context" + "errors" "fmt" + "sync" "tailscale.com/ipn/ipnext" "tailscale.com/net/routecheck" + "tailscale.com/tailcfg" "tailscale.com/types/logger" - "tailscale.com/types/netmap" ) // FeatureName is the name of the feature implemented by this package. @@ -42,6 +45,14 @@ type Extension struct { backend ipnext.SafeBackend nb nodeBackender nm routecheck.NetMapper + routers *RouterTracker + + reconcile struct { + sync.Mutex + args chan tailcfg.NodeView // pending arguments for StartStopWatcher + closed bool + done chan struct{} + } } var _ ipnext.Extension = new(Extension) @@ -65,6 +76,11 @@ func (e *Extension) Init(h ipnext.Host) error { } e.nm = nm + ipnbus, ok := e.backend.(ipnext.NotifyWatcher) + if !ok { + return fmt.Errorf("backend %T does not implement ipnext.NotifyWatcher", e.backend) + } + pinger := e.backend.Sys().Engine.Get() c, err := routecheck.NewClient(e.logf, e.nb, e.nm, pinger) @@ -73,20 +89,77 @@ func (e *Extension) Init(h ipnext.Host) error { } e.Client = c - h.Hooks().OnNetMapToggle.Add(e.onNetMapToggle) + e.routers = TrackRouters(context.Background(), e.logf, ipnbus) + e.routers.OnNetMapAvailable = e.Client.NotifyNetMapAvailable + e.routers.OnRoutersChange = e.incrementalRefresh + + // Watch for changes to the self node that would toggle the routecheck feature. + e.reconcile.args = make(chan tailcfg.NodeView, 1) + e.reconcile.done = make(chan struct{}) + go e.reconcileLoop() + h.Hooks().OnSelfChange.Add(e.reconcileWatcher) return nil } // Shutdown implements the [ipnext.Extension.Shutdown] interface method. func (e *Extension) Shutdown() error { - err := e.Client.Close() - return err + e.reconcile.Lock() + e.reconcile.closed = true + close(e.reconcile.args) // lock prevents reconcileWatcher from writing to this channel + e.reconcile.Unlock() + + e.routers.Close() // stop the watcher before waiting for reconcile.done + <-e.reconcile.done + return e.Client.Close() } -func (e *Extension) onNetMapToggle(nm *netmap.NetworkMap) { - if nm == nil { +func (e *Extension) needsRefresh() { + // TODO(sfllaw): Call e.Client.NeedsRefresh() after implementing it. +} + +func (e *Extension) incrementalRefresh(added, modified, removed []tailcfg.NodeID) { + // TODO(sfllaw): This refresh should be incremental, + // based on the added, modified, and removed nodes. + // Currently it refreshes everything. + e.needsRefresh() +} + +// reconcileWatcher is called whenever e.routers should start, stop, or restart its watcher. +// It may trigger a restart when self indicates that we have switched to a different tailnet or user, +// in order to reset the internal state of e.routers and start tracking from scratch. +// This work is performed by [Extension.reconcileLoop]. +// +// This function must never block, because it’s called from +// [ipnlocal.LocalBackend.SetControlClientStatus], which locks LocalBackend.mu. +// This lock is also acquired when unwinding [ipnlocal.LocalBackend.WatchNotificationsAs] +// which is what [RouterTracker.stopWatcherLocked] is waiting for. +func (e *Extension) reconcileWatcher(self tailcfg.NodeView) { + e.reconcile.Lock() + defer e.reconcile.Unlock() + if e.reconcile.closed { return } - e.Client.NotifyNetMapAvailable(nm) + select { + case <-e.reconcile.args: // drain stale args so StartStopWatcher is always called with the latest + default: + } + e.reconcile.args <- self +} + +// reconcileLoop starts, stops, or restarts its watcher after calls to [Extension.reconcileWatcher]. +func (e *Extension) reconcileLoop() { + defer close(e.reconcile.done) + for self := range e.reconcile.args { + started, err := e.routers.StartStopWatcher(self) + if err != nil { + if !errors.Is(err, ErrRouteCheckNotEnabled) { + e.logf("error tracking routers: %v", err) + } + continue // can be started by toggling the nodeattr + } + if started { + e.needsRefresh() + } + } } diff --git a/feature/routecheck/routertracker.go b/feature/routecheck/routertracker.go new file mode 100644 index 000000000..4f180f1f3 --- /dev/null +++ b/feature/routecheck/routertracker.go @@ -0,0 +1,227 @@ +// Copyright (c) Tailscale Inc & contributors +// SPDX-License-Identifier: BSD-3-Clause + +package routecheck + +import ( + "context" + "errors" + "fmt" + "sync" + + "tailscale.com/ipn" + "tailscale.com/ipn/ipnext" + "tailscale.com/net/routecheck" + "tailscale.com/syncs" + "tailscale.com/tailcfg" + "tailscale.com/types/logger" + "tailscale.com/util/set" +) + +var ErrRouteCheckNotEnabled = errors.New("routecheck not enabled") + +type RouterTracker struct { + // OnNetMapAvailable is called when the initial network map is received + // or is loaded from its cache. + OnNetMapAvailable func() + + // OnRoutersChange is called when one or more peer nodes, which function as routers, + // have been added, removed, or change their routes. + OnRoutersChange func(added, modified, removed []tailcfg.NodeID) + + ctx context.Context // root context + logf logger.Logf + ipnbus ipnext.NotifyWatcher + + mu sync.Mutex + closed bool + self tailcfg.NodeView // self node for the IPN bus being watched + cancel context.CancelFunc // non-nil iff the watcher is running + done chan struct{} // closed by the watcher goroutine when it exits +} + +// TrackRouters returns a tracker for keeping track of which nodes are routers +// by watching the IPN bus for netmap changes. +func TrackRouters(ctx context.Context, logf logger.Logf, ipnbus ipnext.NotifyWatcher) *RouterTracker { + return &RouterTracker{ + ctx: ctx, + logf: logf, + ipnbus: ipnbus, + } +} + +// Close implements the [io.Closer] interface. +func (rt *RouterTracker) Close() error { + if rt == nil { + return nil + } + + rt.mu.Lock() + defer rt.mu.Unlock() + + if rt.closed { + return nil + } + rt.closed = true + + rt.stopWatcherLocked() + return nil +} + +// StartStopWatcher starts or stops watching the IPN bus based on +// the state of the "client-side-reachability-routecheck" node attribute: +// see [tailcfg.NodeAttrClientSideReachabilityRouteCheck]. +// +// StartStopWatcher considers stopping and then restarting the watcher goroutine +// if the self node and user differ from the ones that started the current watcher. +// It stops or starts the watcher when routecheck is disabled or enabled, respectively. +// +// StartStopWatcher reports whether the watcher goroutine was started, +// either because it was previously stopped or because it needed restarting. +func (rt *RouterTracker) StartStopWatcher(self tailcfg.NodeView) (started bool, _ error) { + rt.mu.Lock() + defer rt.mu.Unlock() + + orig := rt.self + rt.self = self + + toggled := routecheck.IsEnabled(orig) != routecheck.IsEnabled(self) + if toggled || !sameNode(orig, self) { + rt.stopWatcherLocked() + if err := rt.startWatcherLocked(self); err != nil { + return false, err + } + return true, nil + } + return false, nil +} + +// sameNode reports whether a and b have the same [tailcfg.NodeView.ID]s. +func sameNode(a, b tailcfg.NodeView) bool { + var aID, bID tailcfg.NodeID + if a.Valid() { + aID = a.ID() + } + if b.Valid() { + bID = b.ID() + } + return aID == bID +} + +// startWatcherLocked launches the goroutine that watches the IPN bus. +// rt.mu must be held and the watcher must not already be running. +func (rt *RouterTracker) startWatcherLocked(self tailcfg.NodeView) error { + syncs.RequiresMutex(&rt.mu) + if rt.closed { + return fmt.Errorf("cannot start, tracker was closed") + } + if rt.cancel != nil || rt.done != nil { + return fmt.Errorf("cannot start, already watching IPN bus") + } + + if !routecheck.IsEnabled(self) { + if !self.Valid() { + return ErrRouteCheckNotEnabled + } + return fmt.Errorf("%w for %v on %v", ErrRouteCheckNotEnabled, self.User(), self.ID()) + } + rt.self = self + + ctx, cancel := context.WithCancel(rt.ctx) + rt.cancel = cancel + rt.done = make(chan struct{}) + + go rt.watchIPNBus(ctx, rt.done, self) + return nil +} + +// stopWatcherLocked cancels the watcher goroutine and waits for it to exit. +// rt.mu must be held. It is a no-op when no watcher is running. +// +// Blocking while locked is safe because the watcher’s callback never locks rt.mu. +func (rt *RouterTracker) stopWatcherLocked() { + syncs.RequiresMutex(&rt.mu) + var stopped bool + if rt.cancel != nil { + rt.cancel() + rt.cancel = nil + stopped = true + } + if rt.done != nil { + <-rt.done + rt.done = nil + stopped = true + } + rt.self = tailcfg.NodeView{} + if stopped { + rt.logf("stopped tracking routers") + } +} + +// watchIPNBus subscribes to the IPN bus to learn about changes to the peer map, +// so that it can keep track of which nodes are routers by dead-reckoning. +// The set of routers is tracked internally to process peer churn without locking. +// +// When routers are added, removed, or change their routes, +// it fires the [RouterTracker.OnRoutersChange] hook. +// See tailscale/tailscale#12542. +// +// When the client gets the initial netmap after connecting to the control plane, +// it fires the [RouterTracker.OnNetMapAvailable] hook. +// +// To avoid stalls, these notifications must be processed promptly +// because we enabled [ipn.NotifyInProcessNoDisconnect] which blocks the caller. +func (rt *RouterTracker) watchIPNBus(ctx context.Context, done chan<- struct{}, self tailcfg.NodeView) { + defer close(done) + + routers := make(set.Set[tailcfg.NodeID]) + const mask = ipn.NotifyInProcessNoDisconnect | ipn.NotifyInitialStatus | ipn.NotifyPeerChanges + rt.ipnbus.WatchNotifications(ctx, mask, nil, func(n *ipn.Notify) bool { + var added, modified, removed []tailcfg.NodeID + if s := n.InitialStatus; s != nil { + if rt.OnNetMapAvailable != nil { + rt.OnNetMapAvailable() + } + // Bootstrap the router set from the initial Status. + // This will trigger the initial probe for all routers. + for _, ps := range s.Peer { + if ps.IsRouter() { + nid := ps.NodeID + routers.Add(nid) + added = append(added, nid) + } + } + } + for _, p := range n.PeersChanged { + nid := p.ID + wasRouter := routers.Contains(p.ID) + isRouter := p.IsRouter() + switch { + case !wasRouter && isRouter: + routers.Add(nid) + added = append(added, nid) + case wasRouter && isRouter: + // TODO(sfllaw): Tune this to ignore changes + // that don’t affect this node’s status as a router. + modified = append(modified, nid) + case wasRouter && !isRouter: + routers.Delete(nid) + removed = append(removed, nid) + } + } + for _, nid := range n.PeersRemoved { + if routers.Contains(nid) { + routers.Delete(nid) + removed = append(removed, nid) + } + } + + if added != nil || modified != nil || removed != nil { + if rt.OnRoutersChange != nil { + rt.OnRoutersChange(added, modified, removed) + } + } + return true + }) + rt.logf("stopped tracking routers") +} diff --git a/feature/routecheck/routertracker_test.go b/feature/routecheck/routertracker_test.go new file mode 100644 index 000000000..753fb6bf4 --- /dev/null +++ b/feature/routecheck/routertracker_test.go @@ -0,0 +1,267 @@ +// Copyright (c) Tailscale Inc & contributors +// SPDX-License-Identifier: BSD-3-Clause + +package routecheck_test + +import ( + "context" + "fmt" + "net/netip" + "testing" + "testing/synctest" + + gcmp "github.com/google/go-cmp/cmp" + + "tailscale.com/feature/routecheck" + "tailscale.com/ipn" + "tailscale.com/ipn/ipnstate" + netroutecheck "tailscale.com/net/routecheck" + "tailscale.com/tailcfg" + "tailscale.com/types/opt" +) + +func TestRouterTracker(t *testing.T) { + self := makeSelfNodeWithRouteCheckEnabled(t) + + for _, tc := range []struct { + name string + bus ipnBus + wantAdded []tailcfg.NodeID + wantModified []tailcfg.NodeID + wantRemoved []tailcfg.NodeID + }{ + { + name: "initial", + bus: ipnBus{}, + wantAdded: nil, + wantModified: nil, + wantRemoved: nil, + }, + { + name: "added", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(1, withRoutes(netip.MustParsePrefix("192.168.1.0/24")))}}, + }, + wantAdded: []tailcfg.NodeID{1}, + wantModified: nil, + wantRemoved: nil, + }, + { + name: "modified", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(2, withRoutes(netip.MustParsePrefix("192.168.1.0/24")))}}, + + {PeersChanged: []*tailcfg.Node{makeNode(2, withRoutes(netip.MustParsePrefix("192.168.2.0/24")))}}, + }, + wantAdded: nil, + wantModified: []tailcfg.NodeID{2}, + wantRemoved: nil, + }, + { + name: "removed", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(3, withRoutes(netip.MustParsePrefix("192.168.3.0/24")))}}, + {PeersRemoved: []tailcfg.NodeID{3}}, + }, + wantAdded: nil, + wantModified: nil, + wantRemoved: []tailcfg.NodeID{3}, + }, + { + name: "removed-already", + bus: ipnBus{ + {PeersRemoved: []tailcfg.NodeID{3}}, + }, + wantAdded: nil, + wantModified: nil, + wantRemoved: nil, + }, + { + name: "plain-node", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(4)}}, + + {PeersChanged: []*tailcfg.Node{makeNode(4)}}, + }, + wantAdded: nil, + wantModified: nil, + wantRemoved: nil, + }, + { + name: "authorized", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(5)}}, + + {PeersChanged: []*tailcfg.Node{makeNode(5, withRoutes(netip.MustParsePrefix("192.168.5.0/24")))}}, + }, + wantAdded: []tailcfg.NodeID{5}, + wantModified: nil, + wantRemoved: nil, + }, + { + name: "unauthorized", + bus: ipnBus{ + {PeersChanged: []*tailcfg.Node{makeNode(6, withRoutes(netip.MustParsePrefix("192.168.6.0/24")))}}, + + {PeersChanged: []*tailcfg.Node{makeNode(6)}}, + }, + wantAdded: nil, + wantModified: nil, + wantRemoved: []tailcfg.NodeID{6}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var gotAdded, gotModified, gotRemoved []tailcfg.NodeID + rt := routecheck.TrackRouters(t.Context(), t.Logf, &tc.bus) + rt.OnRoutersChange = func(added, modified, removed []tailcfg.NodeID) { + gotAdded, gotModified, gotRemoved = added, modified, removed + } + defer rt.Close() + + if started, err := rt.StartStopWatcher(self); err != nil { + t.Fatalf("error starting watcher: %v", err) + } else if !started { + t.Fatalf("failed to start watcher") + } + + synctest.Wait() + if diff := gcmp.Diff(tc.wantAdded, gotAdded); diff != "" { + t.Errorf("mismatched added: -want, +got:\n%s", diff) + } + if diff := gcmp.Diff(tc.wantModified, gotModified); diff != "" { + t.Errorf("mismatched modified: -want, +got:\n%s", diff) + } + if diff := gcmp.Diff(tc.wantRemoved, gotRemoved); diff != "" { + t.Errorf("mismatched removed: -want, +got:\n%s", diff) + } + }) + }) + } +} + +func TestRouterTrackerRaisesOnNetMapAvailable(t *testing.T) { + self := makeSelfNodeWithRouteCheckEnabled(t) + routers := []*tailcfg.Node{makeNode(1, withRoutes(netip.MustParsePrefix("192.168.1.0/24")))} + + for _, tc := range []struct { + name string + bus ipnBus + want opt.Bool + }{ + { + name: "empty", + bus: ipnBus{}, + want: opt.ExplicitlyUnset, + }, + { + name: "initial-status", + bus: ipnBus{ + { + InitialStatus: &ipnstate.Status{}, + PeersChanged: routers, + }, + }, + want: opt.True, + }, + { + name: "no-initial-status", + bus: ipnBus{ + { + InitialStatus: nil, + PeersChanged: routers, + }, + }, + want: opt.False, + }, + } { + t.Run(tc.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + got := opt.ExplicitlyUnset + rt := routecheck.TrackRouters(t.Context(), t.Logf, &tc.bus) + rt.OnNetMapAvailable = func() { + t.Logf("OnNetMapAvailable") + got = opt.True + } + rt.OnRoutersChange = func(_, _, _ []tailcfg.NodeID) { + t.Logf("OnRoutersChange") + if got == opt.ExplicitlyUnset { + got = opt.False + } + } + defer rt.Close() + + if started, err := rt.StartStopWatcher(self); err != nil { + t.Fatalf("error starting watcher: %v", err) + } else if !started { + t.Fatalf("failed to start watcher") + } + + synctest.Wait() + if got != tc.want { + t.Errorf("got %v, want %v", got, tc.want) + } + }) + }) + } +} + +func makeSelfNodeWithRouteCheckEnabled(t *testing.T) tailcfg.NodeView { + t.Helper() + self := (&tailcfg.Node{ + CapMap: tailcfg.NodeCapMap{ + tailcfg.NodeAttrClientSideReachability: nil, + tailcfg.NodeAttrClientSideReachabilityRouteCheck: nil, + }, + }).View() + if !netroutecheck.IsEnabled(self) { + t.Fatalf("routecheck not enabled for self node: %v", self) + } + return self +} + +type ipnBus []ipn.Notify + +func (b *ipnBus) WatchNotifications(ctx context.Context, mask ipn.NotifyWatchOpt, onWatchAdded func(), fn func(roNotify *ipn.Notify) (keepGoing bool)) { + if ctx.Err() != nil { + return + } + if onWatchAdded != nil { + onWatchAdded() + } + for _, n := range *b { + if !fn(&n) { + return + } + } + +} + +type nodeOptFunc func(*tailcfg.Node) + +func makeNode(id tailcfg.NodeID, opts ...nodeOptFunc) *tailcfg.Node { + addresses := []netip.Prefix{ + netip.MustParsePrefix(fmt.Sprintf("192.168.0.%d/32", id)), + netip.MustParsePrefix(fmt.Sprintf("fd7a:115c:a1e0::%d/128", id)), + } + node := &tailcfg.Node{ + ID: id, + StableID: tailcfg.StableNodeID(fmt.Sprintf("stable%d", id)), + Name: fmt.Sprintf("node%d", id), + Online: new(true), + MachineAuthorized: true, + HomeDERP: int(id), + Addresses: addresses, + AllowedIPs: addresses, + } + for _, opt := range opts { + opt(node) + } + return node +} + +func withRoutes(routes ...netip.Prefix) nodeOptFunc { + return func(n *tailcfg.Node) { + n.AllowedIPs = append(n.AllowedIPs, routes...) + } +} diff --git a/ipn/ipnext/ipnext.go b/ipn/ipnext/ipnext.go index 1af259aac..fd11fc204 100644 --- a/ipn/ipnext/ipnext.go +++ b/ipn/ipnext/ipnext.go @@ -6,6 +6,7 @@ package ipnext import ( + "context" "errors" "fmt" "iter" @@ -226,6 +227,33 @@ type SafeBackend interface { TailscaleVarRoot() string } +// NotifyWatcher is a subset of [tailscale.com/ipn/ipnlocal.LocalBackend] +// for extensions that subscribe to the IPN notification bus from within tailscaled. +// +// Unlike [SafeBackend], its methods acquire LocalBackend’s internal mutex +// and must not be called from extension hooks, +// instead call them from a goroutine started by [Extension.Init]. +type NotifyWatcher interface { + // WatchNotifications subscribes to the ipn.Notify message bus notification + // messages. + // + // WatchNotifications blocks until ctx is done. + // + // The provided onWatchAdded, if non-nil, will be called once the watcher + // is installed. + // + // The provided fn will be called for each notification. It will only be + // called with non-nil pointers. The caller must not modify roNotify. If + // fn returns false, the watch also stops. + // + // Failure to consume many notifications in a row will result in one final + // notification with ErrMessage set, followed by the watch closing, unless mask + // includes ipn.NotifyInProcessNoDisconnect. Watchers using + // NotifyInProcessNoDisconnect must not call back into LocalBackend from fn or + // wait on work that might call back into LocalBackend. + WatchNotifications(ctx context.Context, mask ipn.NotifyWatchOpt, onWatchAdded func(), fn func(roNotify *ipn.Notify) (keepGoing bool)) +} + // ExtensionServices provides access to the [Host]'s extension management services, // such as fetching active extensions. type ExtensionServices interface { diff --git a/net/routecheck/routecheck.go b/net/routecheck/routecheck.go index 641c4c0fb..b8c44108e 100644 --- a/net/routecheck/routecheck.go +++ b/net/routecheck/routecheck.go @@ -138,8 +138,8 @@ func NewClient(logf logger.Logf, nb NodeBackender, nm NetMapper, pinger Pinger) // NotifyNetMapAvailable wakes up goroutines that have been waiting for the // non-nil network map that the control plane sends after reconnecting. -func (c *Client) NotifyNetMapAvailable(nm *netmap.NetworkMap) { - if nm == nil { +func (c *Client) NotifyNetMapAvailable() { + if nm := c.nm.NetMapNoPeers(); nm == nil { return // client disconnected } var nextCh *chan struct{} diff --git a/net/routecheck/routecheck_test.go b/net/routecheck/routecheck_test.go index 682e4ca67..3984aea48 100644 --- a/net/routecheck/routecheck_test.go +++ b/net/routecheck/routecheck_test.go @@ -153,7 +153,7 @@ func TestRefresh(t *testing.T) { if tt.init { // This callback simulates the delay between // connecting to the backend and receiving the NetMap. - donef := func() { c.NotifyNetMapAvailable(b.NetMapWithPeers()) } + donef := func() { c.NotifyNetMapAvailable() } b.donef.Store(&donef) }