From 6cdb1b78d95648ced01f5a665ba43d9d5202710f Mon Sep 17 00:00:00 2001 From: Yibo Zhuang Date: Wed, 19 Aug 2026 09:21:32 -0700 Subject: [PATCH] Recover controller sessions without stopping active VMs (#466) --- internal/worker/rpc.go | 8 +- internal/worker/rpcv2.go | 7 +- internal/worker/worker.go | 365 ++++++++++++----- internal/worker/worker_test.go | 701 +++++++++++++++++++++++++++++++++ 4 files changed, 986 insertions(+), 95 deletions(-) diff --git a/internal/worker/rpc.go b/internal/worker/rpc.go index dad5850..f3a87ba 100644 --- a/internal/worker/rpc.go +++ b/internal/worker/rpc.go @@ -22,7 +22,7 @@ import ( "github.com/samber/lo" ) -func (worker *Worker) watchRPC(ctx context.Context) error { +func (worker *Worker) watchRPC(ctx context.Context, operationCtx context.Context, onEstablished func()) error { worker.logger.Infof("connecting to %s over gRPC", worker.client.GRPCTarget()) conn, err := grpc.NewClient(worker.client.GRPCTarget(), @@ -40,11 +40,13 @@ func (worker *Worker) watchRPC(ctx context.Context) error { client := rpc.NewControllerClient(conn) ctxWithMetadata := metadata.NewOutgoingContext(ctx, worker.grpcMetadata()) + operationCtxWithMetadata := metadata.NewOutgoingContext(operationCtx, worker.grpcMetadata()) stream, err := client.Watch(ctxWithMetadata, &emptypb.Empty{}) if err != nil { return err } + onEstablished() worker.logger.Infof("running gRPC stream with the controller") @@ -56,11 +58,11 @@ func (worker *Worker) watchRPC(ctx context.Context) error { switch action := watchFromController.Action.(type) { case *rpc.WatchInstruction_PortForwardAction: - go worker.handlePortForward(ctxWithMetadata, client, action.PortForwardAction) + go worker.handlePortForward(operationCtxWithMetadata, client, action.PortForwardAction) case *rpc.WatchInstruction_SyncVmsAction: worker.requestVMSyncing() case *rpc.WatchInstruction_ResolveIpAction: - go worker.handleGetIP(ctxWithMetadata, client, action.ResolveIpAction) + go worker.handleGetIP(operationCtxWithMetadata, client, action.ResolveIpAction) } } } diff --git a/internal/worker/rpcv2.go b/internal/worker/rpcv2.go index f9d714d..7ddee44 100644 --- a/internal/worker/rpcv2.go +++ b/internal/worker/rpcv2.go @@ -11,21 +11,22 @@ import ( "github.com/samber/lo" ) -func (worker *Worker) watchRPCV2(ctx context.Context) error { +func (worker *Worker) watchRPCV2(ctx context.Context, operationCtx context.Context, onEstablished func()) error { watchInstructionCh, watchErrCh, err := worker.client.RPC().Watch(ctx, worker.name) if err != nil { return err } + onEstablished() for { select { case watchInstruction := <-watchInstructionCh: if portForwardAction := watchInstruction.PortForwardAction; portForwardAction != nil { - go worker.handlePortForwardV2(ctx, portForwardAction) + go worker.handlePortForwardV2(operationCtx, portForwardAction) } else if syncVMsAction := watchInstruction.SyncVMsAction; syncVMsAction != nil { worker.requestVMSyncing() } else if resolveIPAction := watchInstruction.ResolveIPAction; resolveIPAction != nil { - go worker.handleGetIPV2(ctx, resolveIPAction) + go worker.handleGetIPV2(operationCtx, resolveIPAction) } case watchErr := <-watchErrCh: return watchErr diff --git a/internal/worker/worker.go b/internal/worker/worker.go index aa1cb79..19e3616 100644 --- a/internal/worker/worker.go +++ b/internal/worker/worker.go @@ -11,7 +11,6 @@ import ( goruntime "runtime" - "github.com/avast/retry-go/v4" "github.com/cirruslabs/orchard/internal/dialer" "github.com/cirruslabs/orchard/internal/opentelemetry" "github.com/cirruslabs/orchard/internal/worker/dhcpleasetime" @@ -38,9 +37,17 @@ import ( const ( pollInterval = 5 * time.Second workerResourceUpdateInterval = 15 * time.Second + recoveredVMProtectionPeriod = 30 * time.Second + rpcWatchReconnectInterval = 100 * time.Millisecond + rpcWatchReconnectMaxInterval = 5 * time.Second + rpcWatchReconnectMultiplier = 2 + rpcWatchHealthyInterval = time.Second ) -var ErrPollFailed = errors.New("failed to poll controller") +var ( + ErrPollFailed = errors.New("failed to poll controller") + errRPCWatchDisconnected = errors.New("RPC watch disconnected") +) type Worker struct { name string @@ -49,6 +56,7 @@ type Worker struct { vmm *vmmanager.VMManager client *client.Client pollTicker *time.Ticker + recoveredVMs map[ondiskname.OnDiskName]time.Time resources v1.Resources labels v1.Labels @@ -68,6 +76,7 @@ func New(client *client.Client, opts ...Option) (*Worker, error) { worker := &Worker{ client: client, pollTicker: time.NewTicker(pollInterval), + recoveredVMs: make(map[ondiskname.OnDiskName]time.Time), vmm: vmmanager.New(), syncRequested: make(chan bool, 1), } @@ -147,8 +156,20 @@ func (worker *Worker) Run(ctx context.Context) error { } } + var reconnectBackoff rpcWatchReconnectBackoff + reconnectBackoff.reset() + for { - if err := worker.runNewSession(ctx); err != nil { + if err := worker.runNewSession(ctx, reconnectBackoff.reset); err != nil { + if errors.Is(err, errRPCWatchDisconnected) { + select { + case <-time.After(reconnectBackoff.next()): + continue + case <-ctx.Done(): + return ctx.Err() + } + } + return err } @@ -161,6 +182,29 @@ func (worker *Worker) Run(ctx context.Context) error { } } +type rpcWatchReconnectBackoff struct { + nextInterval time.Duration +} + +func (backoff *rpcWatchReconnectBackoff) next() time.Duration { + interval := backoff.nextInterval + if interval <= 0 { + interval = rpcWatchReconnectInterval + } + + if interval >= rpcWatchReconnectMaxInterval/rpcWatchReconnectMultiplier { + backoff.nextInterval = rpcWatchReconnectMaxInterval + } else { + backoff.nextInterval = interval * rpcWatchReconnectMultiplier + } + + return min(interval, rpcWatchReconnectMaxInterval) +} + +func (backoff *rpcWatchReconnectBackoff) reset() { + backoff.nextInterval = rpcWatchReconnectInterval +} + func (worker *Worker) Close() error { var result error for _, vm := range worker.vmm.List() { @@ -175,7 +219,7 @@ func (worker *Worker) Close() error { return result } -func (worker *Worker) runNewSession(ctx context.Context) error { +func (worker *Worker) runNewSession(ctx context.Context, onWatchHealthy func()) error { subCtx, cancel := context.WithCancel(ctx) defer cancel() @@ -192,35 +236,25 @@ func (worker *Worker) runNewSession(ctx context.Context) error { return nil } - if info.Capabilities.Has(v1.ControllerCapabilityRPCV2) { - worker.logger.Infof("using WebSocket-based v2 RPC") - - go func() { - _ = retry.Do(func() error { - return worker.watchRPCV2(subCtx) - }, retry.OnRetry(func(n uint, err error) { - worker.logger.Warnf("failed to watch RPC v2: %v", err) - }), retry.Context(subCtx), retry.Attempts(0)) - }() - } else { - worker.logger.Infof("using gRPC-based v1 RPC") - - go func() { - _ = retry.Do(func() error { - return worker.watchRPC(subCtx) - }, retry.OnRetry(func(n uint, err error) { - worker.logger.Warnf("failed to watch RPC v1: %v", err) - }), retry.Context(subCtx), retry.Attempts(0)) - }() - } + group, sessionCtx := errgroup.WithContext(subCtx) + worker.superviseRPCWatch(sessionCtx, ctx, group, info, onWatchHealthy) // Sync on-disk VMs - if err := worker.syncOnDiskVMs(ctx); err != nil { + if err := worker.syncOnDiskVMs(sessionCtx); err != nil { + cancel() + watchErr := group.Wait() + worker.logger.Errorf("failed to sync on-disk VMs: %v", err) + if errors.Is(watchErr, errRPCWatchDisconnected) { + return watchErr + } + return nil } + recoveredVMs := worker.trackRecoveredVMs(time.Now()) + // Backward compatibility with for older Orchard Controllers updateFuncInner := worker.client.VMs().UpdateState @@ -239,17 +273,15 @@ func (worker *Worker) runNewSession(ctx context.Context) error { return err } - group, ctx := errgroup.WithContext(subCtx) - group.Go(func() error { for { - if err := worker.updateWorker(ctx); err != nil { + if err := worker.updateWorker(sessionCtx); err != nil { return fmt.Errorf("failed to update worker resource: %w", err) } select { - case <-ctx.Done(): - return ctx.Err() + case <-sessionCtx.Done(): + return sessionCtx.Err() case <-time.After(workerResourceUpdateInterval): // Proceed } @@ -258,13 +290,13 @@ func (worker *Worker) runNewSession(ctx context.Context) error { group.Go(func() error { for { - if err := worker.syncVMs(ctx, updateFunc); err != nil { + if err := worker.syncVMs(sessionCtx, updateFunc, recoveredVMs); err != nil { return fmt.Errorf("failed to sync VMs: %w", err) } select { - case <-ctx.Done(): - return ctx.Err() + case <-sessionCtx.Done(): + return sessionCtx.Err() case <-worker.syncRequested: case <-worker.pollTicker.C: // Proceed @@ -274,11 +306,102 @@ func (worker *Worker) runNewSession(ctx context.Context) error { if err := group.Wait(); err != nil { worker.logger.Errorf("%v", err) + + if errors.Is(err, errRPCWatchDisconnected) { + return err + } } return nil } +func (worker *Worker) superviseRPCWatch( + sessionCtx context.Context, + operationCtx context.Context, + group *errgroup.Group, + info v1.ControllerInfo, + onWatchHealthy func(), +) { + watchRPC := worker.watchRPC + rpcVersion := "v1" + + if info.Capabilities.Has(v1.ControllerCapabilityRPCV2) { + worker.logger.Infof("using WebSocket-based v2 RPC") + watchRPC = worker.watchRPCV2 + rpcVersion = "v2" + } else { + worker.logger.Infof("using gRPC-based v1 RPC") + } + + watchEstablished := make(chan struct{}) + + group.Go(func() error { + if err := watchRPC(sessionCtx, operationCtx, func() { close(watchEstablished) }); err != nil { + if sessionCtx.Err() != nil { + return sessionCtx.Err() + } + + return fmt.Errorf("%w: failed to watch RPC %s: %w", errRPCWatchDisconnected, rpcVersion, err) + } + + return fmt.Errorf("%w: RPC %s watch closed unexpectedly", errRPCWatchDisconnected, rpcVersion) + }) + + group.Go(func() error { + return monitorRPCWatchHealth(sessionCtx, watchEstablished, rpcWatchHealthyInterval, onWatchHealthy) + }) +} + +func monitorRPCWatchHealth( + ctx context.Context, + established <-chan struct{}, + healthyAfter time.Duration, + onHealthy func(), +) error { + select { + case <-established: + case <-ctx.Done(): + return ctx.Err() + } + + healthTimer := time.NewTimer(healthyAfter) + defer healthTimer.Stop() + + select { + case <-healthTimer.C: + onHealthy() + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (worker *Worker) trackRecoveredVMs(now time.Time) map[ondiskname.OnDiskName]time.Time { + if worker.recoveredVMs == nil { + worker.recoveredVMs = make(map[ondiskname.OnDiskName]time.Time) + } + + for onDiskName := range worker.recoveredVMs { + if !worker.vmm.Exists(onDiskName) { + delete(worker.recoveredVMs, onDiskName) + } + } + + for _, vm := range worker.vmm.List() { + status := vm.Status() + if status != v1.VMStatusPending && status != v1.VMStatusRunning { + continue + } + + onDiskName := vm.OnDiskName() + if _, alreadyTracked := worker.recoveredVMs[onDiskName]; !alreadyTracked { + worker.recoveredVMs[onDiskName] = now.Add(recoveredVMProtectionPeriod) + } + } + + return worker.recoveredVMs +} + func (worker *Worker) registerWorker(ctx context.Context) error { platformUUID, err := platform.MachineID() if err != nil { @@ -326,8 +449,12 @@ func (worker *Worker) updateWorker(ctx context.Context) error { return nil } -//nolint:nestif,gocognit // nested "if" and cognitive complexity is tolerable for now -func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context, v1.VM) error) error { +//nolint:gocognit // VM lifecycle branches are clearest in a single reconciliation loop. +func (worker *Worker) syncVMs( + ctx context.Context, + updateVM func(context.Context, v1.VM) error, + recoveredVMs map[ondiskname.OnDiskName]time.Time, +) error { allKeys := mapset.NewSet[ondiskname.OnDiskName]() remoteVMs, err := worker.client.VMs().FindForWorker(ctx, worker.name) @@ -367,6 +494,8 @@ func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context // we risk violating the scheduler resource assumptions sortNonExistentAndFailedFirst(pairs) + hasUnaccountedRecoveredVM := false + for _, tuple := range pairs { onDiskName, vmResource, vm := lo.Unpack3(tuple) @@ -382,6 +511,14 @@ func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context localConditions = vm.Conditions() } + if shouldPreserveRecoveredVM(recoveredVMs, onDiskName, vmResource, localState, time.Now()) { + hasUnaccountedRecoveredVM = true + worker.logger.Warnf("preserving active VM %s missing from controller during worker session recovery", + onDiskName) + + continue + } + action := transitions[remoteState][localState] worker.logger.Debugf("processing VM: %s, remote state: %s, local state: %s, "+ @@ -391,6 +528,13 @@ func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context switch action { case ActionCreate: // Remote VM was created, but not the local VM + if hasUnaccountedRecoveredVM { + worker.logger.Warnf("deferring VM %s while recovered VMs are missing from controller inventory", + onDiskName) + + continue + } + worker.createVM(onDiskName, *vmResource) case ActionMonitorPending: if vmResource.StatusMessage != vm.StatusMessage() { @@ -418,59 +562,8 @@ func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context return err } case ActionMonitorRunning: - if vmResource.Generation != vm.Resource().Generation { - // VM specification changed, reboot the VM for the changes to take effect - stoppingOrSuspending := v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeStopping) || - v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeSuspending) - - if v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeRunning) && !stoppingOrSuspending { - // VM is running, suspend or stop it first - shouldStop := vmResource.PowerState == v1.PowerStateStopped || !vm.Resource().Suspendable - - if shouldStop { - vm.Stop() - } else { - vm.Suspend() - } - } - - if v1.ConditionIsFalse(vm.Conditions(), v1.ConditionTypeRunning) && !stoppingOrSuspending { - // VM stopped, update its specification - vm.SetResource(*vmResource) - - if vmResource.PowerState == v1.PowerStateRunning { - // Start the VM - eventStreamer := worker.client.VMs().StreamEvents(vmResource.Name) - vm.Start(eventStreamer) - } - } - } - - var updateNeeded bool - - if vmResource.StatusMessage != vm.StatusMessage() { - vmResource.StatusMessage = vm.StatusMessage() - - updateNeeded = true - } - - if vmResource.ObservedGeneration != vm.Resource().ObservedGeneration { - vmResource.ObservedGeneration = vm.Resource().ObservedGeneration - - updateNeeded = true - } - - // Propagate VM's conditions to the Orchard Controller - for _, condition := range vm.Conditions() { - if v1.ConditionsSet(&vmResource.Conditions, condition) { - updateNeeded = true - } - } - - if updateNeeded { - if err := updateVM(ctx, *vmResource); err != nil { - return err - } + if err := worker.monitorRunningVM(ctx, vmResource, vm, updateVM); err != nil { + return err } case ActionStop: // VM has failed on the remote side, stop it locally to prevent incorrect @@ -517,6 +610,100 @@ func (worker *Worker) syncVMs(ctx context.Context, updateVM func(context.Context return nil } +func (worker *Worker) monitorRunningVM( + ctx context.Context, + vmResource *v1.VM, + vm vmmanager.VM, + updateVM func(context.Context, v1.VM) error, +) error { + worker.reconcileRunningVM(vmResource, vm) //nolint:contextcheck // Event streams outlive sync sessions. + + var updateNeeded bool + + if vmResource.StatusMessage != vm.StatusMessage() { + vmResource.StatusMessage = vm.StatusMessage() + + updateNeeded = true + } + + if vmResource.ObservedGeneration != vm.Resource().ObservedGeneration { + vmResource.ObservedGeneration = vm.Resource().ObservedGeneration + + updateNeeded = true + } + + // Propagate VM's conditions to the Orchard Controller + for _, condition := range vm.Conditions() { + if v1.ConditionsSet(&vmResource.Conditions, condition) { + updateNeeded = true + } + } + + if updateNeeded { + return updateVM(ctx, *vmResource) + } + + return nil +} + +func (worker *Worker) reconcileRunningVM(vmResource *v1.VM, vm vmmanager.VM) { + if vmResource.Generation == vm.Resource().Generation { + return + } + + stoppingOrSuspending := v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeStopping) || + v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeSuspending) + if stoppingOrSuspending { + return + } + + if v1.ConditionIsTrue(vm.Conditions(), v1.ConditionTypeRunning) { + // VM is running, suspend or stop it first. + shouldStop := vmResource.PowerState == v1.PowerStateStopped || !vm.Resource().Suspendable + + if shouldStop { + vm.Stop() + } else { + vm.Suspend() + } + } + + if v1.ConditionIsFalse(vm.Conditions(), v1.ConditionTypeRunning) { + // VM stopped, update its specification. + vm.SetResource(*vmResource) + + if vmResource.PowerState == v1.PowerStateRunning { + // Start the VM. + eventStreamer := worker.client.VMs().StreamEvents(vmResource.Name) + vm.Start(eventStreamer) + } + } +} + +func shouldPreserveRecoveredVM( + recoveredVMs map[ondiskname.OnDiskName]time.Time, + onDiskName ondiskname.OnDiskName, + remoteVM *v1.VM, + localState mo.Option[v1.VMStatus], + now time.Time, +) bool { + deadline, recovered := recoveredVMs[onDiskName] + if !recovered { + return false + } + + active := localState == mo.Some(v1.VMStatusPending) || localState == mo.Some(v1.VMStatusRunning) + if remoteVM == nil && active && now.Before(deadline) { + return true + } + + // Once the controller recognizes a recovered VM, or the bounded recovery + // window expires, user-requested deletion follows the normal lifecycle. + delete(recoveredVMs, onDiskName) + + return false +} + //nolint:nestif,gocognit // complexity is tolerable for now func (worker *Worker) syncOnDiskVMs(ctx context.Context) error { if worker.runtime.Synthetic() { diff --git a/internal/worker/worker_test.go b/internal/worker/worker_test.go index 93ee32b..8c4424f 100644 --- a/internal/worker/worker_test.go +++ b/internal/worker/worker_test.go @@ -1,16 +1,717 @@ package worker import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "sync" + "sync/atomic" "testing" + "time" "github.com/cirruslabs/orchard/internal/worker/ondiskname" "github.com/cirruslabs/orchard/internal/worker/vmmanager" "github.com/cirruslabs/orchard/internal/worker/vmmanager/tart" + "github.com/cirruslabs/orchard/pkg/client" v1 "github.com/cirruslabs/orchard/pkg/resource/v1" + "github.com/cirruslabs/orchard/rpc" + "github.com/coder/websocket" "github.com/samber/lo" + "github.com/samber/mo" "github.com/stretchr/testify/require" + "go.uber.org/zap" + "golang.org/x/net/http2" + "golang.org/x/net/http2/h2c" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/types/known/emptypb" ) +const ( + recoveryTestWorkerName = "worker-a" + recoveryTestVMUID = "running-vm-uid" + recoveryTestWorkerPath = "/v1/workers/" + recoveryTestWorkerName + recoveryTestVMsPath = "/v1/vms" + recoveryTestWatchPath = "/v1/rpc/watch" +) + +func TestWorkerRecoversControllerSessionWithoutDeletingRunningVM(t *testing.T) { + firstHeartbeat := make(chan struct{}) + reregistered := make(chan struct{}) + var firstHeartbeatOnce sync.Once + var registrations atomic.Int32 + var watches atomic.Int32 + + vmResource := v1.VM{ + Meta: v1.Meta{Name: "running-vm"}, + UID: recoveryTestVMUID, + Worker: recoveryTestWorkerName, + Status: v1.VMStatusRunning, + } + recoveredVM := &recoveryTestVM{ + resource: vmResource, + conditionsSeen: make(chan struct{}), + } + + controller := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodPost && request.URL.Path == "/v1/workers": + var workerResource v1.Worker + if err := json.NewDecoder(request.Body).Decode(&workerResource); err != nil { + t.Errorf("failed to decode worker registration: %v", err) + return + } + + if registrations.Add(1) == 2 { + close(reregistered) + } + + writeRecoveryTestJSON(t, writer, workerResource) + case request.Method == http.MethodGet && request.URL.Path == "/v1/controller/info": + writeRecoveryTestJSON(t, writer, v1.ControllerInfo{ + Capabilities: v1.ControllerCapabilities{v1.ControllerCapabilityRPCV2}, + }) + case request.Method == http.MethodGet && request.URL.Path == recoveryTestWorkerPath: + writeRecoveryTestJSON(t, writer, v1.Worker{Meta: v1.Meta{Name: recoveryTestWorkerName}}) + case request.Method == http.MethodPut && request.URL.Path == recoveryTestWorkerPath: + var workerResource v1.Worker + if err := json.NewDecoder(request.Body).Decode(&workerResource); err != nil { + t.Errorf("failed to decode worker heartbeat: %v", err) + return + } + + firstHeartbeatOnce.Do(func() { close(firstHeartbeat) }) + writeRecoveryTestJSON(t, writer, workerResource) + case request.Method == http.MethodGet && request.URL.Path == recoveryTestVMsPath: + writeRecoveryTestJSON(t, writer, []v1.VM{}) + case request.URL.Path == recoveryTestWatchPath: + handleRecoveryTestWatch(t, writer, request, &watches, firstHeartbeat, recoveredVM.conditionsSeen) + default: + http.NotFound(writer, request) + } + })) + t.Cleanup(controller.Close) + + controllerClient, err := client.New(client.WithAddress(controller.URL)) + require.NoError(t, err) + + worker, err := New(controllerClient, WithName(recoveryTestWorkerName), WithSynthetic(), WithLogger(zap.NewNop())) + require.NoError(t, err) + + t.Cleanup(worker.pollTicker.Stop) + worker.vmm.Put(ondiskname.NewFromResource(vmResource), recoveredVM) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + runResult := make(chan error, 1) + go func() { + runResult <- worker.Run(ctx) + }() + + select { + case <-reregistered: + case <-time.After(2 * time.Second): + t.Fatalf("worker did not re-register after its controller RPC watch disconnected (registrations=%d, watches=%d)", + registrations.Load(), watches.Load()) + } + + require.True(t, worker.vmm.Exists(ondiskname.NewFromResource(vmResource))) + require.False(t, recoveredVM.stopped.Load(), "controller recovery must not stop a running VM") + require.False(t, recoveredVM.deleted.Load(), "controller recovery must not delete a running VM") + + cancel() + require.ErrorIs(t, <-runResult, context.Canceled) +} + +func handleRecoveryTestWatch( + t *testing.T, + writer http.ResponseWriter, + request *http.Request, + watches *atomic.Int32, + firstHeartbeat <-chan struct{}, + conditionsSeen <-chan struct{}, +) { + t.Helper() + + connection, err := websocket.Accept(writer, request, nil) + if err != nil { + t.Errorf("failed to accept RPC watch: %v", err) + return + } + defer connection.CloseNow() + + if watches.Add(1) == 1 { + select { + case <-firstHeartbeat: + case <-request.Context().Done(): + return + } + + select { + case <-conditionsSeen: + case <-request.Context().Done(): + return + } + + _ = connection.Close(websocket.StatusGoingAway, "controller rollout") + return + } + + <-request.Context().Done() +} + +func TestRPCWatchReconnectBackoff(t *testing.T) { + var backoff rpcWatchReconnectBackoff + backoff.reset() + + expected := []time.Duration{ + 100 * time.Millisecond, + 200 * time.Millisecond, + 400 * time.Millisecond, + 800 * time.Millisecond, + 1600 * time.Millisecond, + 3200 * time.Millisecond, + 5 * time.Second, + 5 * time.Second, + } + for _, interval := range expected { + require.Equal(t, interval, backoff.next()) + } + + backoff.reset() + require.Equal(t, 100*time.Millisecond, backoff.next(), + "a successfully connected RPC watch should restore the fast initial retry") +} + +func TestMonitorRPCWatchHealth(t *testing.T) { + t.Run("closed watch is not healthy", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + established := make(chan struct{}) + close(established) + var markedHealthy atomic.Bool + result := make(chan error, 1) + + go func() { + result <- monitorRPCWatchHealth(ctx, established, time.Second, func() { + markedHealthy.Store(true) + }) + }() + + cancel() + require.ErrorIs(t, <-result, context.Canceled) + require.False(t, markedHealthy.Load()) + }) + + t.Run("stable watch resets backoff", func(t *testing.T) { + established := make(chan struct{}) + close(established) + var backoff rpcWatchReconnectBackoff + backoff.nextInterval = rpcWatchReconnectMaxInterval + + require.NoError(t, monitorRPCWatchHealth(context.Background(), established, + 10*time.Millisecond, backoff.reset)) + require.Equal(t, rpcWatchReconnectInterval, backoff.next()) + }) +} + +func TestWorkerBacksOffPersistentRPCWatchFailures(t *testing.T) { + testCases := []struct { + name string + useRPCV2 bool + closeAfter bool + }{ + {name: "HTTP rejection", useRPCV2: true}, + {name: "WebSocket closes immediately", useRPCV2: true, closeAfter: true}, + {name: "gRPC first receive fails"}, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + testWorkerBacksOffPersistentRPCWatchFailure(t, testCase.useRPCV2, testCase.closeAfter) + }) + } +} + +func testWorkerBacksOffPersistentRPCWatchFailure(t *testing.T, useRPCV2 bool, closeAfter bool) { + t.Helper() + + watchAttempts := make(chan time.Time, 4) + grpcServer := grpc.NewServer() + rpc.RegisterControllerServer(grpcServer, &failingRecoveryRPCServer{watchAttempts: watchAttempts}) + t.Cleanup(grpcServer.Stop) + + handler := http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.Header.Get("Content-Type") == "application/grpc" { + grpcServer.ServeHTTP(writer, request) + return + } + + switch { + case request.Method == http.MethodPost && request.URL.Path == "/v1/workers": + var workerResource v1.Worker + if err := json.NewDecoder(request.Body).Decode(&workerResource); err != nil { + t.Errorf("failed to decode worker registration: %v", err) + return + } + writeRecoveryTestJSON(t, writer, workerResource) + case request.Method == http.MethodGet && request.URL.Path == "/v1/controller/info": + info := v1.ControllerInfo{} + if useRPCV2 { + info.Capabilities = v1.ControllerCapabilities{v1.ControllerCapabilityRPCV2} + } + writeRecoveryTestJSON(t, writer, info) + case request.Method == http.MethodGet && request.URL.Path == recoveryTestWorkerPath: + writeRecoveryTestJSON(t, writer, v1.Worker{Meta: v1.Meta{Name: recoveryTestWorkerName}}) + case request.Method == http.MethodPut && request.URL.Path == recoveryTestWorkerPath: + writeRecoveryTestJSON(t, writer, v1.Worker{Meta: v1.Meta{Name: recoveryTestWorkerName}}) + case request.Method == http.MethodGet && request.URL.Path == recoveryTestVMsPath: + writeRecoveryTestJSON(t, writer, []v1.VM{}) + case request.URL.Path == recoveryTestWatchPath: + select { + case watchAttempts <- time.Now(): + default: + } + + if closeAfter { + connection, err := websocket.Accept(writer, request, nil) + if err != nil { + t.Errorf("failed to accept rapidly closing RPC watch: %v", err) + return + } + + _ = connection.Close(websocket.StatusGoingAway, "upstream unavailable") + return + } + + writer.WriteHeader(http.StatusForbidden) + default: + http.NotFound(writer, request) + } + }) + controller := httptest.NewServer(h2c.NewHandler(handler, &http2.Server{})) + t.Cleanup(controller.Close) + + controllerClient, err := client.New(client.WithAddress(controller.URL)) + require.NoError(t, err) + worker, err := New(controllerClient, WithName(recoveryTestWorkerName), WithSynthetic(), WithLogger(zap.NewNop())) + require.NoError(t, err) + t.Cleanup(worker.pollTicker.Stop) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + runResult := make(chan error, 1) + go func() { + runResult <- worker.Run(ctx) + }() + + attempts := make([]time.Time, 0, 4) + for len(attempts) < 4 { + select { + case attemptedAt := <-watchAttempts: + attempts = append(attempts, attemptedAt) + case <-time.After(3 * time.Second): + t.Fatalf("worker made only %d RPC watch attempts", len(attempts)) + } + } + + expectedMinimums := []time.Duration{100 * time.Millisecond, 200 * time.Millisecond, 400 * time.Millisecond} + for index, minimum := range expectedMinimums { + actual := attempts[index+1].Sub(attempts[index]) + require.GreaterOrEqual(t, actual, minimum-20*time.Millisecond, + "RPC watch retry %d did not increase its backoff", index+1) + } + + cancel() + require.ErrorIs(t, <-runResult, context.Canceled) +} + +type failingRecoveryRPCServer struct { + rpc.UnimplementedControllerServer + + watchAttempts chan time.Time +} + +func (server *failingRecoveryRPCServer) Watch( + _ *emptypb.Empty, + _ rpc.Controller_WatchServer, +) error { + select { + case server.watchAttempts <- time.Now(): + default: + } + + return status.Error(codes.Unavailable, "RPC watch upstream unavailable") +} + +func TestShouldPreserveRecoveredVM(t *testing.T) { + onDiskName := ondiskname.New("running-vm", recoveryTestVMUID, 0) + now := time.Unix(1_000, 0) + deadline := now.Add(recoveredVMProtectionPeriod) + + t.Run("missing running VM stays protected", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.True(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusRunning), now)) + require.Contains(t, recoveredVMs, onDiskName) + }) + + t.Run("missing pending VM stays protected", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.True(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusPending), now)) + require.Contains(t, recoveredVMs, onDiskName) + }) + + t.Run("missing running VM loses protection after recovery deadline", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusRunning), deadline)) + require.NotContains(t, recoveredVMs, onDiskName) + }) + + t.Run("missing pending VM loses protection after recovery deadline", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusPending), deadline)) + require.NotContains(t, recoveredVMs, onDiskName) + }) + + t.Run("recognized VM returns to normal deletion behavior", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, + &v1.VM{Status: v1.VMStatusRunning}, mo.Some(v1.VMStatusRunning), now)) + require.NotContains(t, recoveredVMs, onDiskName) + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusRunning), now)) + }) + + t.Run("failed recovered VM is not protected", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{onDiskName: deadline} + + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusFailed), now)) + require.NotContains(t, recoveredVMs, onDiskName) + }) + + t.Run("new VM is not protected", func(t *testing.T) { + recoveredVMs := map[ondiskname.OnDiskName]time.Time{} + + require.False(t, shouldPreserveRecoveredVM(recoveredVMs, onDiskName, nil, + mo.Some(v1.VMStatusRunning), now)) + }) +} + +func TestTrackRecoveredVMsPreservesDeadlinesAcrossSessions(t *testing.T) { + firstVMResource := v1.VM{ + Meta: v1.Meta{Name: "first-running-vm"}, + UID: "first-running-vm-uid", + Worker: recoveryTestWorkerName, + Status: v1.VMStatusRunning, + } + secondVMResource := v1.VM{ + Meta: v1.Meta{Name: "second-pending-vm"}, + UID: "second-pending-vm-uid", + Worker: recoveryTestWorkerName, + Status: v1.VMStatusPending, + } + + worker := &Worker{vmm: vmmanager.New()} + firstOnDiskName := ondiskname.NewFromResource(firstVMResource) + worker.vmm.Put(firstOnDiskName, &recoveryTestVM{resource: firstVMResource}) + + firstSession := time.Unix(1_000, 0) + firstDeadline := firstSession.Add(recoveredVMProtectionPeriod) + require.Equal(t, firstDeadline, worker.trackRecoveredVMs(firstSession)[firstOnDiskName]) + + secondSession := firstSession.Add(20 * time.Second) + secondOnDiskName := ondiskname.NewFromResource(secondVMResource) + worker.vmm.Put(secondOnDiskName, &recoveryTestVM{resource: secondVMResource}) + recoveredVMs := worker.trackRecoveredVMs(secondSession) + + require.Equal(t, firstDeadline, recoveredVMs[firstOnDiskName], + "controller reconnects must not extend an existing VM's protection") + require.Equal(t, secondSession.Add(recoveredVMProtectionPeriod), recoveredVMs[secondOnDiskName], + "newly recovered pending VMs receive their own bounded protection window") + + worker.vmm.Delete(secondOnDiskName) + require.NotContains(t, worker.trackRecoveredVMs(secondSession.Add(time.Second)), secondOnDiskName) +} + +func TestSyncVMsDeletesRecoveredVMAfterProtectionExpires(t *testing.T) { + controller := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.URL.Path != recoveryTestVMsPath { + http.NotFound(writer, request) + return + } + + writeRecoveryTestJSON(t, writer, []v1.VM{}) + })) + t.Cleanup(controller.Close) + + controllerClient, err := client.New(client.WithAddress(controller.URL)) + require.NoError(t, err) + + worker, err := New(controllerClient, WithName(recoveryTestWorkerName), WithSynthetic(), WithLogger(zap.NewNop())) + require.NoError(t, err) + t.Cleanup(worker.pollTicker.Stop) + + vmResource := v1.VM{ + Meta: v1.Meta{Name: "deleted-vm"}, + UID: "deleted-vm-uid", + Worker: recoveryTestWorkerName, + Status: v1.VMStatusRunning, + } + deletedVM := &recoveryTestVM{ + resource: vmResource, + conditionsSeen: make(chan struct{}), + } + onDiskName := ondiskname.NewFromResource(vmResource) + worker.vmm.Put(onDiskName, deletedVM) + + recoveredVMs := map[ondiskname.OnDiskName]time.Time{ + onDiskName: time.Now().Add(-time.Second), + } + err = worker.syncVMs(context.Background(), func(context.Context, v1.VM) error { + return nil + }, recoveredVMs) + + require.NoError(t, err) + require.True(t, deletedVM.stopped.Load(), "expired recovery protection must not prevent VM shutdown") + require.True(t, deletedVM.deleted.Load(), "expired recovery protection must not prevent VM deletion") + require.False(t, worker.vmm.Exists(onDiskName)) + require.NotContains(t, recoveredVMs, onDiskName) +} + +func TestSyncVMsDefersNewVMWhileRecoveredVMIsUnaccounted(t *testing.T) { + for _, existingStatus := range []v1.VMStatus{v1.VMStatusPending, v1.VMStatusRunning} { + t.Run(string(existingStatus), func(t *testing.T) { + existing := v1.VM{ + Meta: v1.Meta{Name: "missing-vm"}, + UID: "missing-vm-uid", + Worker: recoveryTestWorkerName, + Status: existingStatus, + Resources: v1.Resources{v1.ResourceTartVMs: 1}, + } + replacement := v1.VM{ + Meta: v1.Meta{Name: "replacement-vm"}, + UID: "replacement-vm-uid", + Worker: recoveryTestWorkerName, + Status: v1.VMStatusPending, + Resources: v1.Resources{v1.ResourceTartVMs: 1}, + } + + controller := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.URL.Path != recoveryTestVMsPath { + http.NotFound(writer, request) + return + } + + writeRecoveryTestJSON(t, writer, []v1.VM{replacement}) + })) + t.Cleanup(controller.Close) + + controllerClient, err := client.New(client.WithAddress(controller.URL)) + require.NoError(t, err) + worker, err := New(controllerClient, + WithName(recoveryTestWorkerName), + WithSynthetic(), + WithResources(v1.Resources{v1.ResourceTartVMs: 1}), + WithLogger(zap.NewNop()), + ) + require.NoError(t, err) + t.Cleanup(worker.pollTicker.Stop) + + existingOnDiskName := ondiskname.NewFromResource(existing) + existingVM := &recoveryTestVM{ + resource: existing, + conditionsSeen: make(chan struct{}), + } + worker.vmm.Put(existingOnDiskName, existingVM) + recoveredVMs := map[ondiskname.OnDiskName]time.Time{ + existingOnDiskName: time.Now().Add(recoveredVMProtectionPeriod), + } + updateVM := func(context.Context, v1.VM) error { return nil } + + require.NoError(t, worker.syncVMs(context.Background(), updateVM, recoveredVMs)) + require.True(t, worker.vmm.Exists(existingOnDiskName)) + require.False(t, worker.vmm.Exists(ondiskname.NewFromResource(replacement)), + "a one-slot worker must not start another VM while an unaccounted VM is active") + require.Len(t, worker.vmm.List(), 1) + + recoveredVMs[existingOnDiskName] = time.Now().Add(-time.Second) + require.NoError(t, worker.syncVMs(context.Background(), updateVM, recoveredVMs)) + require.True(t, existingVM.stopped.Load()) + require.True(t, existingVM.deleted.Load()) + require.True(t, worker.vmm.Exists(ondiskname.NewFromResource(replacement)), + "the queued VM should start after the expired recovered VM is removed") + t.Cleanup(func() { require.NoError(t, worker.Close()) }) + }) + } +} + +func TestWatchRPCV2PreservesActiveOperationsAfterWatchCloses(t *testing.T) { + operationStarted := make(chan struct{}) + releaseOperation := make(chan struct{}) + resolvedIP := make(chan string, 1) + + controller := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case recoveryTestWatchPath: + connection, err := websocket.Accept(writer, request, nil) + if err != nil { + t.Errorf("failed to accept RPC watch: %v", err) + return + } + defer connection.CloseNow() + + instruction := v1.WatchInstruction{ResolveIPAction: &v1.ResolveIPAction{ + Session: "ip-session", + VMUID: recoveryTestVMUID, + }} + payload, err := json.Marshal(instruction) + if err != nil { + t.Errorf("failed to encode RPC instruction: %v", err) + return + } + payload = append(payload, '\n') + + if err := connection.Write(request.Context(), websocket.MessageBinary, payload); err != nil { + t.Errorf("failed to send RPC instruction: %v", err) + return + } + + <-request.Context().Done() + case "/v1/rpc/resolve-ip": + resolvedIP <- request.URL.Query().Get("ip") + writer.WriteHeader(http.StatusOK) + default: + http.NotFound(writer, request) + } + })) + t.Cleanup(controller.Close) + + controllerClient, err := client.New(client.WithAddress(controller.URL)) + require.NoError(t, err) + + worker, err := New(controllerClient, WithName(recoveryTestWorkerName), WithSynthetic(), WithLogger(zap.NewNop())) + require.NoError(t, err) + t.Cleanup(worker.pollTicker.Stop) + + vmResource := v1.VM{Meta: v1.Meta{Name: "running-vm"}, UID: recoveryTestVMUID} + worker.vmm.Put(ondiskname.NewFromResource(vmResource), &recoveryIPTestVM{ + recoveryTestVM: recoveryTestVM{resource: vmResource}, + started: operationStarted, + release: releaseOperation, + }) + + operationCtx, cancelOperation := context.WithCancel(context.Background()) + t.Cleanup(cancelOperation) + watchCtx, cancelWatch := context.WithCancel(operationCtx) + t.Cleanup(cancelWatch) + + watchResult := make(chan error, 1) + go func() { + watchResult <- worker.watchRPCV2(watchCtx, operationCtx, func() {}) + }() + + select { + case <-operationStarted: + case err := <-watchResult: + t.Fatalf("RPC watch terminated before its operation started: %v", err) + case <-time.After(2 * time.Second): + t.Fatal("RPC operation did not start") + } + + cancelWatch() + require.ErrorIs(t, <-watchResult, context.Canceled) + close(releaseOperation) + + select { + case ip := <-resolvedIP: + require.Equal(t, "192.0.2.10", ip) + case <-time.After(2 * time.Second): + t.Fatal("active RPC operation was canceled with its watch session") + } +} + +type recoveryTestVM struct { + vmmanager.VM + + resource v1.VM + conditionsSeen chan struct{} + conditionsOnce sync.Once + stopped atomic.Bool + deleted atomic.Bool +} + +type recoveryIPTestVM struct { + recoveryTestVM + + started chan struct{} + release chan struct{} +} + +func (vm *recoveryIPTestVM) IP(ctx context.Context) (string, error) { + close(vm.started) + + select { + case <-vm.release: + return "192.0.2.10", nil + case <-ctx.Done(): + return "", ctx.Err() + } +} + +func (vm *recoveryTestVM) Resource() v1.VM { + return vm.resource +} + +func (vm *recoveryTestVM) OnDiskName() ondiskname.OnDiskName { + return ondiskname.NewFromResource(vm.resource) +} + +func (vm *recoveryTestVM) Status() v1.VMStatus { + return vm.resource.Status +} + +func (vm *recoveryTestVM) Conditions() []v1.Condition { + vm.conditionsOnce.Do(func() { close(vm.conditionsSeen) }) + + return nil +} + +func (vm *recoveryTestVM) Stop() <-chan error { + vm.stopped.Store(true) + + result := make(chan error, 1) + result <- nil + + return result +} + +func (vm *recoveryTestVM) Delete() error { + vm.deleted.Store(true) + + return nil +} + +func writeRecoveryTestJSON(t *testing.T, writer http.ResponseWriter, value any) { + t.Helper() + + writer.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(writer).Encode(value); err != nil { + t.Errorf("failed to encode controller response: %v", err) + } +} + func TestSortNonExistentAndFailedFirst(t *testing.T) { newVMTuple := func(name string, vmResource *v1.VM) lo.Tuple3[ondiskname.OnDiskName, *v1.VM, vmmanager.VM] { return lo.T3[ondiskname.OnDiskName, *v1.VM, vmmanager.VM](