From 5f0cf87429dc93e9882f12ad2d97c5e00b6d8a15 Mon Sep 17 00:00:00 2001 From: Raj Singh Date: Wed, 23 Sep 2026 04:36:01 -0500 Subject: [PATCH] cmd/containerboot: recover from IPN watch closure (#21374) When containerboot falls behind on the IPN bus, tailscaled closes the watch. containerboot treated the EOF as fatal and SIGTERMed a healthy tailscaled, which is easy to hit on large, churny tailnets. Instead, reconnect and rebuild state from the new watch's initial status, and only request peer changes in modes that use them. If the watch can't be reopened for a minute, exit so a dead tailscaled still restarts the container. Fixes #21373 Change-Id: Iad7749e4fd0f43eabdb471d6e64bb43f37ff70ff Signed-off-by: Raj Singh --- cmd/containerboot/main.go | 154 ++++++++++++++++++++++------ cmd/containerboot/main_test.go | 177 +++++++++++++++++++++++++++++++-- 2 files changed, 291 insertions(+), 40 deletions(-) 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) + } +}