diff --git a/cmd/containerboot/main.go b/cmd/containerboot/main.go index fe56b6960..5d249cc00 100644 --- a/cmd/containerboot/main.go +++ b/cmd/containerboot/main.go @@ -156,6 +156,7 @@ import ( "tailscale.com/tailcfg" "tailscale.com/types/logger" "tailscale.com/types/views" + "tailscale.com/util/backoff" "tailscale.com/util/deephash" "tailscale.com/util/def" "tailscale.com/util/dnsname" @@ -173,9 +174,13 @@ func getAutoAdvertiseBool() bool { return def.Bool(os.Getenv("TS_EXPERIMENTAL_SERVICE_AUTO_ADVERTISEMENT"), true) } -const containerbootWatchMask = ipn.NotifyInitialStatus | - ipn.NotifyPeerChanges | - ipn.NotifyNoNetMap +func containerbootWatchMask(cfg *settings) ipn.NotifyWatchOpt { + mask := ipn.NotifyInitialStatus | ipn.NotifyNoNetMap + if cfg.TailnetTargetFQDN != "" || cfg.EgressProxiesCfgPath != "" { + mask |= ipn.NotifyPeerChanges + } + return mask +} func notifyState(n ipn.Notify) (_ ipn.State, ok bool) { if n.State != nil { @@ -235,21 +240,105 @@ func (s netmapState) processNotify(ctx context.Context, client *local.Client, n } func (s netmapState) updateFromStatus(st *ipnstate.Status) netmapState { - s.certDomains = views.SliceOf(st.CertDomains) - s.dnsExtraRecords = views.SliceOf(st.ExtraRecords) + s = netmapState{ + certDomains: views.SliceOf(st.CertDomains), + dnsExtraRecords: views.SliceOf(st.ExtraRecords), + } if st.Self != nil { s.self = nodeFromPeerStatus(st.Self).View() } - if len(st.Peer) != 0 { - s.peersByID = nil - s.peersByName = nil - for _, ps := range st.Peer { - s = s.upsertPeer(nodeFromPeerStatus(ps).View()) - } + for _, ps := range st.Peer { + s = s.upsertPeer(nodeFromPeerStatus(ps).View()) } return s } +type watchIPNBusFunc func(context.Context, ipn.NotifyWatchOpt) (klc.IPNBusWatcher, error) + +// maxIPNBusDialFailure is how long reconnectingIPNBusWatcher keeps failing to +// open a new watch before giving up. A closed stream is always retried, but a +// watch that cannot be opened at all usually means tailscaled is gone. +const maxIPNBusDialFailure = time.Minute + +type reconnectingIPNBusWatcher struct { + ctx context.Context + watch watchIPNBusFunc + mask ipn.NotifyWatchOpt + bo *backoff.Backoff + watcher klc.IPNBusWatcher + startedAt time.Time + maxDialFailure time.Duration + dialFailSince time.Time +} + +func newReconnectingIPNBusWatcher(ctx context.Context, name string, watch watchIPNBusFunc, mask ipn.NotifyWatchOpt, maxBackoff time.Duration) *reconnectingIPNBusWatcher { + return &reconnectingIPNBusWatcher{ + ctx: ctx, + watch: watch, + mask: mask, + bo: backoff.NewBackoff(name, log.Printf, maxBackoff), + maxDialFailure: maxIPNBusDialFailure, + } +} + +// Next returns the next notification, reconnecting after the watch stream +// closes. Each new watch starts with an authoritative InitialStatus snapshot. +func (w *reconnectingIPNBusWatcher) Next() (ipn.Notify, error) { + for { + if err := w.ctx.Err(); err != nil { + return ipn.Notify{}, err + } + if w.watcher == nil { + watcher, err := w.watch(w.ctx, w.mask) + if err != nil { + if w.dialFailSince.IsZero() { + w.dialFailSince = time.Now() + } else if time.Since(w.dialFailSince) >= w.maxDialFailure { + return ipn.Notify{}, err + } + w.bo.BackOff(w.ctx, err) + continue + } + w.watcher = watcher + w.startedAt = time.Now() + w.dialFailSince = time.Time{} + } + + n, err := w.watcher.Next() + if err == nil { + if n.ErrMessage != nil { + log.Printf("tailscaled IPN bus error: %s", *n.ErrMessage) + } + return n, nil + } + w.watcher.Close() + w.watcher = nil + if ctxErr := w.ctx.Err(); ctxErr != nil { + return ipn.Notify{}, ctxErr + } + log.Printf("IPN bus watch ended; reconnecting: %v", err) + if time.Since(w.startedAt) >= 30*time.Second { + w.bo.Reset() + } + w.bo.BackOff(w.ctx, err) + } +} + +func (w *reconnectingIPNBusWatcher) SetMask(mask ipn.NotifyWatchOpt) { + w.Close() + w.mask = mask + w.bo.Reset() +} + +func (w *reconnectingIPNBusWatcher) Close() error { + if w.watcher == nil { + return nil + } + err := w.watcher.Close() + w.watcher = nil + return err +} + func (s netmapState) upsertPeer(n tailcfg.NodeView) netmapState { if !n.Valid() { return s @@ -461,10 +550,10 @@ func run() error { } } - w, err := client.WatchIPNBus(bootCtx, containerbootWatchMask|ipn.NotifyInitialPrefs|ipn.NotifyInitialHealthState) - if err != nil { - return fmt.Errorf("failed to watch tailscaled for updates: %w", err) - } + watchMask := containerbootWatchMask(cfg) + authWatchMask := watchMask | ipn.NotifyInitialPrefs | ipn.NotifyInitialHealthState + localClient := klc.New(client) + w := newReconnectingIPNBusWatcher(bootCtx, "containerboot-auth-ipn-watch", localClient.WatchIPNBus, authWatchMask, 5*time.Second) // Now that we've started tailscaled, we can symlink the socket to the // default location if needed. @@ -501,10 +590,8 @@ func run() error { if err := tailscaleUp(bootCtx, cfg); err != nil { return fmt.Errorf("failed to auth tailscale: %w", err) } - w, err = client.WatchIPNBus(bootCtx, containerbootWatchMask) - if err != nil { - return fmt.Errorf("rewatching tailscaled for updates after auth: %w", err) - } + authWatchMask = watchMask + w.SetMask(authWatchMask) return nil } @@ -518,7 +605,7 @@ authLoop: for { n, err := w.Next() if err != nil { - return fmt.Errorf("failed to read from tailscaled: %w", err) + return fmt.Errorf("reading tailscaled IPN bus: %w", err) } if state, ok := notifyState(n); ok { @@ -613,11 +700,6 @@ authLoop: } } - w, err = client.WatchIPNBus(ctx, containerbootWatchMask) - if err != nil { - return fmt.Errorf("rewatching tailscaled for updates after auth: %w", err) - } - // If tailscaled config was read from a mounted file, watch the file for updates and reload. cfgWatchErrChan := make(chan error) cfgWatchCtx, cfgWatchCancel := context.WithCancel(ctx) @@ -694,15 +776,21 @@ authLoop: var egressSvcsNotify chan netmapState notifyChan := make(chan ipn.Notify) - errChan := make(chan error) + errChan := make(chan error, 1) + steadyWatch := newReconnectingIPNBusWatcher(ctx, "containerboot-ipn-watch", localClient.WatchIPNBus, watchMask, 30*time.Second) go func() { for { - n, err := w.Next() + n, err := steadyWatch.Next() if err != nil { - errChan <- err - break - } else { - notifyChan <- n + if ctx.Err() == nil { + errChan <- err + } + return + } + select { + case notifyChan <- n: + case <-ctx.Done(): + return } } }() @@ -721,7 +809,7 @@ runLoop: killTailscaled() break runLoop case err := <-errChan: - return fmt.Errorf("failed to read from tailscaled: %w", err) + return fmt.Errorf("reading tailscaled IPN bus: %w", err) case err := <-cfgWatchErrChan: return fmt.Errorf("failed to watch tailscaled config: %w", err) case n := <-notifyChan: @@ -908,7 +996,7 @@ runLoop: return fmt.Errorf("autoadvertisement: failed to get serve config: %w", err) } - err = refreshAdvertiseServices(ctx, prevServeConfig, klc.New(client)) + err = refreshAdvertiseServices(ctx, prevServeConfig, localClient) if err != nil { return fmt.Errorf("autoadvertisement: failed to refresh advertise services: %w", err) } diff --git a/cmd/containerboot/main_test.go b/cmd/containerboot/main_test.go index 0013bb8e4..8c4a079d0 100644 --- a/cmd/containerboot/main_test.go +++ b/cmd/containerboot/main_test.go @@ -41,10 +41,12 @@ import ( "tailscale.com/kube/egressservices" "tailscale.com/kube/kubeclient" "tailscale.com/kube/kubetypes" + klc "tailscale.com/kube/localclient" "tailscale.com/net/memnet" "tailscale.com/tailcfg" "tailscale.com/tstest" "tailscale.com/types/key" + "tailscale.com/types/views" ) const configFileAuthKey = "some-auth-key" @@ -1290,12 +1292,6 @@ func TestContainerBoot(t *testing.T) { t.Fatalf("phase %d: updating mtime for %q: %v", i, path, err) } } - if p.Notify != nil && p.Notify.InitialStatus == nil { - // Shallow-copy before mutating to avoid a race with - // parallel subtests that share the same *ipn.Notify. - p.Notify = new(*p.Notify) - p.Notify.InitialStatus = statusFromNotify(p.Notify) - } env.lapi.Notify(p.Notify) if p.Signal != nil { cmd.Process.Signal(*p.Signal) @@ -1488,6 +1484,7 @@ type localAPI struct { sync.Mutex cond *sync.Cond notify *ipn.Notify + status *ipnstate.Status } func (lc *localAPI) Start() error { @@ -1521,6 +1518,32 @@ func (lc *localAPI) Notify(n *ipn.Notify) { lc.Lock() defer lc.Unlock() lc.notify = n + if n.InitialStatus != nil { + lc.status = n.InitialStatus + } else if lc.status == nil { + lc.status = statusFromNotify(n) + } else { + if n.State != nil { + lc.status.BackendState = n.State.String() + } + if n.SelfChange != nil { + lc.status.Self = peerStatusFromNode(n.SelfChange.View()) + } + for _, p := range n.PeersChanged { + if lc.status.Peer == nil { + lc.status.Peer = map[key.NodePublic]*ipnstate.PeerStatus{} + } + pv := p.View() + lc.status.Peer[pv.Key()] = peerStatusFromNode(pv) + } + for _, id := range n.PeersRemoved { + for k, p := range lc.status.Peer { + if p.NodeID == id { + delete(lc.status.Peer, k) + } + } + } + } lc.cond.Broadcast() } @@ -1617,9 +1640,15 @@ func (lc *localAPI) ServeHTTP(w http.ResponseWriter, r *http.Request) { enc := json.NewEncoder(w) lc.Lock() defer lc.Unlock() + first := true for { if lc.notify != nil { - if err := enc.Encode(lc.notify); err != nil { + n := *lc.notify + if first { + n.InitialStatus = lc.status + first = false + } + if err := enc.Encode(&n); err != nil { // Usually broken pipe as the test client disconnects. return } @@ -2002,3 +2031,137 @@ func TestProcessNotifyRefreshesDNSOnSelfChange(t *testing.T) { t.Errorf("certDomains = %v, want [node.tailnet.ts.net]", got.certDomains.AsSlice()) } } + +func TestContainerbootWatchMask(t *testing.T) { + const base = ipn.NotifyInitialStatus | ipn.NotifyNoNetMap + tests := []struct { + name string + cfg settings + want ipn.NotifyWatchOpt + }{ + {name: "subnet_router", cfg: settings{Routes: new("10.0.0.0/8")}, want: base}, + {name: "tailnet_target_IP", cfg: settings{TailnetTargetIP: "100.64.0.1"}, want: base}, + {name: "tailnet_target_FQDN", cfg: settings{TailnetTargetFQDN: "target.example.ts.net"}, want: base | ipn.NotifyPeerChanges}, + {name: "egress_services", cfg: settings{EgressProxiesCfgPath: "/etc/egress-services"}, want: base | ipn.NotifyPeerChanges}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := containerbootWatchMask(&tt.cfg); got != tt.want { + t.Errorf("containerbootWatchMask() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestUpdateFromStatusReplacesState(t *testing.T) { + oldPeer := &tailcfg.Node{ID: 1, StableID: "old", Name: "old.example.ts.net."} + s := netmapState{ + self: oldPeer.View(), + certDomains: views.SliceOf([]string{"old.example.ts.net"}), + dnsExtraRecords: views.SliceOf([]tailcfg.DNSRecord{{Name: "old.example.ts.net."}}), + }.upsertPeer(oldPeer.View()) + + got := s.updateFromStatus(new(ipnstate.Status)) + if got.self.Valid() { + t.Error("self remained valid after an empty authoritative status") + } + if got.peersByID != nil && got.peersByID.Len() != 0 { + t.Errorf("peer count = %d, want 0", got.peersByID.Len()) + } + if got.certDomains.Len() != 0 || got.dnsExtraRecords.Len() != 0 { + t.Error("DNS state remained after an empty authoritative status") + } +} + +type scriptedIPNBusWatcher struct { + notifies []ipn.Notify + err error +} + +func (w *scriptedIPNBusWatcher) Close() error { return nil } + +func (w *scriptedIPNBusWatcher) Next() (ipn.Notify, error) { + if len(w.notifies) == 0 { + return ipn.Notify{}, w.err + } + n := w.notifies[0] + w.notifies = w.notifies[1:] + return n, nil +} + +func TestReconnectingIPNBusWatcher(t *testing.T) { + terminalMessage := "IPN bus consumer fell behind; closing watch" + watches := []*scriptedIPNBusWatcher{ + {notifies: []ipn.Notify{{ErrMessage: &terminalMessage}}, err: io.EOF}, + {notifies: []ipn.Notify{{InitialStatus: &ipnstate.Status{BackendState: ipn.Running.String()}}}, err: io.EOF}, + } + var ( + mu sync.Mutex + masks []ipn.NotifyWatchOpt + ) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + watch := func(_ context.Context, mask ipn.NotifyWatchOpt) (klc.IPNBusWatcher, error) { + mu.Lock() + defer mu.Unlock() + masks = append(masks, mask) + if len(watches) == 0 { + return nil, errors.New("no scripted watcher") + } + w := watches[0] + watches = watches[1:] + return w, nil + } + mask := ipn.NotifyInitialStatus | ipn.NotifyNoNetMap + w := newReconnectingIPNBusWatcher(ctx, "test", watch, mask, time.Millisecond) + defer w.Close() + + n, err := w.Next() + if err != nil { + t.Fatal(err) + } + if n.ErrMessage == nil || *n.ErrMessage != terminalMessage { + t.Fatalf("first notification = %+v, want terminal error", n) + } + + n, err = w.Next() + if err != nil { + t.Fatal(err) + } + if n.InitialStatus == nil || n.InitialStatus.BackendState != ipn.Running.String() { + t.Fatalf("notification = %+v, want replacement running status", n) + } + cancel() + if _, err := w.Next(); !errors.Is(err, context.Canceled) { + t.Fatalf("Next error = %v, want context.Canceled", err) + } + + mu.Lock() + defer mu.Unlock() + if len(masks) < 2 { + t.Fatalf("watch attempts = %d, want at least 2", len(masks)) + } + for i, got := range masks[:2] { + if got != mask { + t.Errorf("watch attempt %d mask = %v, want %v", i, got, mask) + } + } +} + +func TestReconnectingIPNBusWatcherGivesUpOnDialFailure(t *testing.T) { + dialErr := errors.New("connection refused") + var attempts int + watch := func(context.Context, ipn.NotifyWatchOpt) (klc.IPNBusWatcher, error) { + attempts++ + return nil, dialErr + } + w := newReconnectingIPNBusWatcher(t.Context(), "test", watch, ipn.NotifyInitialStatus, time.Millisecond) + w.maxDialFailure = 20 * time.Millisecond + + if _, err := w.Next(); !errors.Is(err, dialErr) { + t.Fatalf("Next error = %v, want %v", err, dialErr) + } + if attempts < 2 { + t.Errorf("watch attempts = %d, want at least 2", attempts) + } +}