diff --git a/internal/advancer/advancer.go b/internal/advancer/advancer.go index fb64f4337..546f40ab3 100644 --- a/internal/advancer/advancer.go +++ b/internal/advancer/advancer.go @@ -313,7 +313,16 @@ func (s *Service) processInputs( // Store the result in the database err = s.repository.StoreAdvanceResult(ctx, input.EpochApplicationID, result) if err != nil { - if errors.Is(err, repository.ErrApplicationNotRunnable) { + var errCause string + switch { + case errors.Is(err, context.Canceled) && errors.Is(ctx.Err(), context.Canceled): + errCause = "canceled advance result persistence" + // Shutdown interrupted persistence after the machine advanced. + // Discard that runtime; restart must recover from persisted state. + s.Logger.Debug("Advance result persistence canceled during shutdown; closing machine", + "application", app.Name, "index", input.Index, "error", err) + case errors.Is(err, repository.ErrApplicationNotRunnable): + errCause = "application status race" // Another service durably fenced this application after the // machine began its advance. Discard only this stale runtime; the // database already contains the authoritative app-local outcome. @@ -322,41 +331,35 @@ func (s *Service) processInputs( "epoch", input.EpochIndex, "index", input.Index, "error", err) - closeErr := machine.Close() - if closeErr != nil { - s.Logger.Error("Could not close stale machine after application status race", - "application", app.Name, - "error", closeErr) - } - return processed, false, errors.Join(err, closeErr) + default: + errCause = "its advance result was not confirmed saved; service shutdown is still required" + // Advance has already changed the live machine, but the transaction + // did not confirm that its result was saved. The database may still + // show this input as pending. Reusing this machine could then execute + // the input again from the wrong state, so StoreAdvanceResult is not + // retried against this live machine. + s.Logger.Error( + "Could not confirm that the advance result was saved; "+ + "the live machine has already advanced, so services will stop; "+ + "after the node is restarted, execution will use persisted state", + "application", app.Name, + "epoch", input.EpochIndex, + "index", input.Index, + "error", err) + s.supervisor.Fatal(fmt.Errorf("unconfirmed advance result for %s: %w", app.Name, err)) } - // Advance has already changed the live machine, but the transaction - // did not confirm that its result was saved. The database may still - // show this input as pending. Reusing this machine could then execute - // the input again from the wrong state, so StoreAdvanceResult is not - // retried against this live machine. - s.Logger.Error( - "Could not confirm that the advance result was saved; "+ - "the live machine has already advanced, so services will stop; "+ - "after the node is restarted, execution will use persisted state", - "application", app.Name, - "epoch", input.EpochIndex, - "index", input.Index, - "error", err) - // Try to close the machine now so the already-advanced runtime cannot // be used again. Cancel services even if Close fails. After the node // is restarted, the machine is rebuilt from persisted state, and the // database decides whether this input is still pending and needs a // safe retry. closeErr := machine.Close() - s.supervisor.Fatal(fmt.Errorf("unconfirmed advance result for %s: %w", app.Name, errors.Join(err, closeErr))) if closeErr != nil { - s.Logger.Error("Could not close the machine after its advance result "+ - "was not confirmed saved; service shutdown is still required", + s.Logger.Error("Could not close the machine after "+errCause, "application", app.Name, "error", closeErr) + s.supervisor.Fatal(fmt.Errorf("close machine for %s after %s: %w", app.Name, errCause, closeErr)) } return processed, false, errors.Join(err, closeErr) } diff --git a/internal/advancer/advancer_test.go b/internal/advancer/advancer_test.go index 40efc5ff1..2e93c7655 100644 --- a/internal/advancer/advancer_test.go +++ b/internal/advancer/advancer_test.go @@ -2336,6 +2336,7 @@ type MockMachineInstance struct { machineImpl *MockMachineImpl createSnapshotError error destroyAfterSnapshotError bool + closeError error closeCalls int advanceCalls int } @@ -2386,7 +2387,7 @@ func (m *MockMachineInstance) Hash(ctx context.Context) ([32]byte, error) { // Close implements the MachineInstance interface for testing func (m *MockMachineInstance) Close() error { m.closeCalls++ - return nil + return m.closeError } // ------------------------------------------------------------------------------------------------ @@ -2398,6 +2399,7 @@ type MockRepository struct { GetInputsReturn map[common.Address][]*Input GetInputsError error GetInputsBlock bool + StoreAdvanceHook func(context.Context) error StoreAdvanceError error StoreAdvanceCommitError error StoreAdvanceFailCount int @@ -2502,6 +2504,9 @@ func (mock *MockRepository) StoreAdvanceResult( appID int64, res *AdvanceResult, ) error { + if mock.StoreAdvanceHook != nil { + return mock.StoreAdvanceHook(ctx) + } // Check for context cancellation if ctx.Err() != nil { return ctx.Err() @@ -2750,3 +2755,62 @@ func marshal(res *AdvanceResult) []byte { } return data } + +func (s *AdvancerSuite) TestStoreAdvanceShutdownClassification() { + storageErr := errors.New("storage failed") + closeErr := errors.New("machine close failed") + for _, tc := range []struct { + name string + cancel bool + storeErr error + closeErr error + fatalErr error + }{ + {"shutdown cancellation", true, fmt.Errorf("store: %w", context.Canceled), nil, nil}, + {"independent cancellation", false, context.Canceled, nil, context.Canceled}, + {"storage failure", false, storageErr, nil, storageErr}, + {"storage failure during shutdown", true, storageErr, nil, storageErr}, + {"close failure during shutdown", true, context.Canceled, closeErr, closeErr}, + } { + s.Run(tc.name, func() { + require := s.Require() + env := s.setupOneApp() + ctx, cancel := context.WithCancel(s.T().Context()) + defer cancel() + machine := env.mm.Map[env.app.Application.ID] + machine.closeError = tc.closeErr + env.repo.StoreAdvanceHook = func(storeCtx context.Context) error { + require.NoError(storeCtx.Err(), "cancellation must happen during persistence") + require.Equal(1, machine.advanceCalls) + if tc.cancel { + cancel() + } + return tc.storeErr + } + pending := newInput(env.app.Application.ID, 0, 0, marshal(randomAdvanceResult(0))) + address := env.app.Application.IApplicationAddress + env.repo.GetInputsReturn = map[common.Address][]*Input{address: {pending}} + processed, stopped, err := env.service.processInputs(ctx, env.app.Application, []*Input{ + pending, newInput(env.app.Application.ID, 0, 1, []byte("unreachable")), + }) + require.ErrorIs(err, tc.storeErr) + require.Zero(processed) + require.False(stopped) + require.Equal(1, machine.closeCalls) + require.Equal(1, machine.advanceCalls) + require.Empty(env.repo.StoredResults) + require.Equal([]*Input{pending}, env.repo.GetInputsReturn[address]) + fatal := env.supervisor.FatalError.Load() + if tc.fatalErr == nil { + require.Nil(fatal, "shutdown cancellation must not become fatal") + } else { + require.NotNil(fatal) + require.ErrorIs(*fatal, tc.fatalErr) + if tc.closeErr != nil { + require.ErrorIs(err, tc.closeErr) + require.NotErrorIs(*fatal, context.Canceled) + } + } + }) + } +} diff --git a/internal/evmreader/error_paths_test.go b/internal/evmreader/error_paths_test.go index fc785f2cb..291239548 100644 --- a/internal/evmreader/error_paths_test.go +++ b/internal/evmreader/error_paths_test.go @@ -561,7 +561,7 @@ func (s *EvmReaderSuite) TestEpochLengthZeroSetsAppCorrupted() { s.evmReader.repository = repo err := s.evmReader.readAndStoreInputs(s.ctx, 100, 110, apps) - s.Require().ErrorIs(err, errScanIncomplete) + s.Require().NoError(err) // App must be set inoperable repo.AssertNumberOfCalls(s.T(), "UpdateApplicationStatus", 1) diff --git a/internal/evmreader/evmreader.go b/internal/evmreader/evmreader.go index ea171e409..8eaf35863 100644 --- a/internal/evmreader/evmreader.go +++ b/internal/evmreader/evmreader.go @@ -177,9 +177,10 @@ func (r *Service) processBlockHead( return r.runBlockScanners(ctx, apps, blockNumber) && len(apps) == len(observableApps) } -// runBlockScanners reports whether all scheduled observation work completed without -// errors. An idle scan is healthy. Always run the remaining scanners after a -// failure so one application's stall does not prevent progress elsewhere. +// runBlockScanners reports whether observation completed without shared failures. +// Known input corruption with a persisted integrity status is application-local +// degradation; observation still runs on later ticks. An idle scan is healthy. +// Always run remaining scanners so one application's stall cannot block others. func (r *Service) runBlockScanners( ctx context.Context, apps []appContracts, diff --git a/internal/evmreader/input.go b/internal/evmreader/input.go index f76b2a84c..c9846eb97 100644 --- a/internal/evmreader/input.go +++ b/internal/evmreader/input.go @@ -300,6 +300,28 @@ func indexInputsIntoEpochs( return epochInputMap, nil } +// recordInputCorruption separates a known application-local integrity failure +// from failure to persist its status. Status helpers update app.Status only after +// a successful write and preserve previously persisted integrity terminals. +// Observation must continue on later ticks, including for terminal applications. +func (r *Service) recordInputCorruption( + ctx context.Context, + app *Application, + reasonFmt string, + args ...any, +) bool { + reason := fmt.Sprintf(reasonFmt, args...) + // setApplicationCorrupted always returns non-nil (the reason text itself). + // The DB error case is already logged inside setApplicationStatus. + // TODO: consider returning only DB errors instead of always returning an error. + _ = r.setApplicationCorrupted(ctx, app, "%s", reason) + if app.Status == ApplicationStatus_Corrupted || app.Status == ApplicationStatus_Diverged { + r.Logger.Warn("Input observation degraded for application", "application", app.Name, "reason", reason) + return true + } + return false +} + // readAndStoreInputs reads, inputs from the InputSource given specific filter options, indexes // them into epochs and store the indexed inputs and epochs func (r *Service) readAndStoreInputs( @@ -308,8 +330,6 @@ func (r *Service) readAndStoreInputs( mostRecentBlockNumber uint64, apps []appContracts, ) error { - var scanErr error - if len(apps) == 0 { r.Logger.Warn("No valid running applications") return nil @@ -325,8 +345,10 @@ func (r *Service) readAndStoreInputs( err) } + var scanIncomplete bool + if len(appInputsMap) != len(apps) { - scanErr = errScanIncomplete + scanIncomplete = true } addrToApp := mapAddressToApp(apps) @@ -337,19 +359,15 @@ func (r *Service) readAndStoreInputs( if !exists { r.Logger.Error("Application address on input not found", "address", address) - scanErr = errScanIncomplete + scanIncomplete = true continue } epochLength := app.application.EpochLength if epochLength == 0 { - // setApplicationCorrupted always returns non-nil (the reason text itself). - // The DB error case is already logged inside setApplicationStatus. - // On DB success the app is marked inoperable and won't reappear next tick. - // On DB failure the app reappears as Enabled next tick, retrying this path. - _ = r.setApplicationCorrupted(ctx, app.application, + ok := r.recordInputCorruption(ctx, app.application, "Application has epoch length of zero") - scanErr = errScanIncomplete + scanIncomplete = scanIncomplete || !ok continue } @@ -372,7 +390,7 @@ func (r *Service) readAndStoreInputs( "error", err, ) } - scanErr = errScanIncomplete + scanIncomplete = true continue } @@ -381,8 +399,10 @@ func (r *Service) readAndStoreInputs( epochLength, currentEpoch, inputs, mostRecentBlockNumber) if err != nil { if errors.Is(err, ErrInputForNonOpenEpoch) { - return r.setApplicationCorrupted(ctx, app.application, + ok := r.recordInputCorruption(ctx, app.application, "Should never happen. %v", err) + scanIncomplete = scanIncomplete || !ok + continue } return fmt.Errorf("error indexing inputs: %w", err) } @@ -416,15 +436,10 @@ func (r *Service) readAndStoreInputs( ) if err != nil { if errors.Is(err, repository.ErrInputLogIdentityConflict) { - // A stored input's L1 log identity disagrees with rescanned - // chain data. Retrying the same insert every tick cannot - // succeed; without escalation the app would stall silently - // with Status OK. See setApplicationCorrupted contract in - // the epochLength == 0 branch above. - _ = r.setApplicationCorrupted(ctx, app.application, + ok := r.recordInputCorruption(ctx, app.application, "stored input L1 log identity conflicts with rescanned chain data"+ " (possible reorg past the input cursor); operator reset required. %v", err) - scanErr = errScanIncomplete + scanIncomplete = scanIncomplete || !ok continue } r.Logger.Error("Error storing inputs and epochs", @@ -432,7 +447,7 @@ func (r *Service) readAndStoreInputs( "address", address, "error", err, ) - scanErr = errScanIncomplete + scanIncomplete = true continue } r.Logger.Debug("Inputs and epochs stored successfully", @@ -486,7 +501,7 @@ func (r *Service) readAndStoreInputs( "error", err, ) } - scanErr = errScanIncomplete + scanIncomplete = true } else { r.Logger.Debug("Updated LastInputCheckBlock for applications without inputs", "app_ids", appsToUpdate, @@ -495,7 +510,10 @@ func (r *Service) readAndStoreInputs( } } - return scanErr + if scanIncomplete { + return errScanIncomplete + } + return nil } // readInputsFromBlockchain fetches inputs for each application independently. diff --git a/internal/evmreader/readiness_test.go b/internal/evmreader/readiness_test.go index 1aec31f14..2bff483ba 100644 --- a/internal/evmreader/readiness_test.go +++ b/internal/evmreader/readiness_test.go @@ -6,6 +6,8 @@ package evmreader import ( "context" "errors" + "fmt" + "github.com/cartesi/rollups-node/internal/repository" "log/slog" "math/big" "testing" @@ -115,3 +117,96 @@ func TestReadinessTracksScanFailures(t *testing.T) { }) } } + +func TestReadinessSeparatesInputCorruptionFromSharedFailures(t *testing.T) { + for _, failure := range []string{"recorded conflict", "status write", "RPC", "store", "epoch query", "zero epoch", "zero epoch status write", "non-open epoch", "non-open epoch status write"} { + t.Run(failure, func(t *testing.T) { + bad := &Application{ID: 1, Name: "degraded", Enabled: true, + IApplicationAddress: app1Addr, IInputBoxAddress: inputBoxAddr, + DataAvailability: DataAvailability_InputBox[:], EpochLength: 10, + Status: ApplicationStatus_OK, LastInputCheckBlock: 100, + LastOutputCheckBlock: 110, LastForecloseCheckBlock: 110} + good := *bad + good.ID, good.Name, good.IApplicationAddress = 2, "healthy", common.HexToAddress("0x2222") + // Terminal execution status must not prevent healthy observation. + good.Status = ApplicationStatus_Corrupted + repo := newMockRepository() + list := repo.On("ListApplications", mock.Anything, mock.Anything, mock.Anything, false). + Return([]*Application{bad, &good}, uint64(2), nil) + input := newMockInputBox() + input.On("GetNumberOfInputs", mock.Anything, good.IApplicationAddress).Return(big.NewInt(0), nil) + count := input.On("GetNumberOfInputs", mock.Anything, bad.IApplicationAddress).Return(big.NewInt(1), nil) + input.On("RetrieveInputs", mock.Anything, mock.Anything, mock.Anything). + Return([]iinputbox.IInputBoxInputAdded{makeInputEvent(bad.IApplicationAddress, 0, 101)}, nil) + repo.On("GetNumberOfInputs", mock.Anything, mock.Anything).Return(uint64(0), nil) + repo.On("GetEpoch", mock.Anything, good.IApplicationAddress.String(), mock.Anything).Return((*Epoch)(nil), nil) + epoch := repo.On("GetEpoch", mock.Anything, bad.IApplicationAddress.String(), mock.Anything).Return((*Epoch)(nil), nil) + store := repo.On("CreateEpochsAndInputs", mock.Anything, bad.IApplicationAddress.String(), mock.Anything, uint64(110)). + Return(fmt.Errorf("insert: %w", repository.ErrInputLogIdentityConflict)) + status := repo.On("UpdateApplicationStatus", mock.Anything, bad.ID, ApplicationStatus_Corrupted, mock.Anything).Return(nil) + repo.On("UpdateEventLastCheckBlock", mock.Anything, []int64{good.ID}, MonitoredEvent_InputAdded, uint64(110)). + Return(nil).Run(func(mock.Arguments) { good.LastInputCheckBlock = 110 }) + infraErr := errors.New("infrastructure unavailable") + switch failure { + case "status write": + status.Return(infraErr) + case "RPC": + count.Return((*big.Int)(nil), infraErr) + case "store": + store.Return(infraErr) + case "epoch query": + epoch.Return((*Epoch)(nil), infraErr) + case "non-open epoch", "non-open epoch status write": + epoch.Return(&Epoch{Index: 10, Status: EpochStatus_Closed}, nil) + if failure == "non-open epoch status write" { + status.Return(infraErr) + } + case "zero epoch": + bad.EpochLength = 0 + case "zero epoch status write": + bad.EpochLength = 0 + status.Return(infraErr) + } + client := newMockEthClient() + client.On("HeaderByNumber", mock.Anything, mock.Anything).Return(&types.Header{Number: big.NewInt(110)}, nil) + r := &Service{client: client, repository: repo, inputReaderEnabled: true, + defaultBlock: DefaultBlock_Latest, readyMaxStaleness: time.Hour} + r.Logger = slog.Default() + r.resolver = newApplicationAdapterResolver(r.Logger, + newMockAdapterFactory().SetupDefaultBehaviorSingleApp(newMockApplicationContract(), input)) + local := failure == "recorded conflict" || failure == "zero epoch" || failure == "non-open epoch" + for i := 1; i <= maxConsecutiveScanFailures+1; i++ { + _, err := r.Tick(t.Context()) + require.NoError(t, err) + require.Equal(t, local || i < maxConsecutiveScanFailures, r.Ready()) + require.EqualValues(t, 110, good.LastInputCheckBlock) + require.EqualValues(t, 100, bad.LastInputCheckBlock, "failed input observation must not advance its cursor") + } + if failure == "non-open epoch" || failure == "non-open epoch status write" { + repo.AssertNumberOfCalls(t, "CreateEpochsAndInputs", 0) + } + if failure == "recorded conflict" { + repo.AssertNumberOfCalls(t, "CreateEpochsAndInputs", maxConsecutiveScanFailures+1) + } + if local { + require.Equal(t, ApplicationStatus_Corrupted, bad.Status) + repo.AssertNumberOfCalls(t, "UpdateApplicationStatus", 1) + } + // Disabling the stalled app recovers readiness without changing its status. + previousStatus := bad.Status + bad.Enabled = false + list.Return([]*Application{&good}, uint64(1), nil) + _, err := r.Tick(t.Context()) + require.NoError(t, err) + require.True(t, r.Ready()) + require.Zero(t, r.consecutiveScanFailures.Load()) + require.Equal(t, previousStatus, bad.Status) + // Both cursors at the head also constitute a healthy idle scan. + bad.Enabled, bad.LastInputCheckBlock = true, 110 + list.Return([]*Application{bad, &good}, uint64(2), nil) + _, err = r.Tick(t.Context()) + require.NoError(t, err) + require.True(t, r.Ready()) + }) + } +} diff --git a/pkg/service/supervisor.go b/pkg/service/supervisor.go index bd8e35f5d..221bf811e 100644 --- a/pkg/service/supervisor.go +++ b/pkg/service/supervisor.go @@ -11,6 +11,7 @@ import ( "os" "os/signal" "slices" + "sync" "sync/atomic" "syscall" @@ -37,7 +38,7 @@ type Supervisor interface { NotReady() []string Serve() error Stop() bool - // Fatal records the first fatal cause and initiates shutdown. Nil uses ErrServiceStopped. + // Fatal joins all fatal causes and initiates shutdown. Fatal(error) } @@ -49,10 +50,11 @@ type supervisorImpl struct { context context.Context cancel context.CancelFunc sigShutdown chan os.Signal // SIGINT/SIGTERM to exit gracefully + fatalErr error + fatalMux sync.RWMutex serving atomic.Bool stopping atomic.Bool - fatal atomic.Pointer[error] } func NewSupervisor(ctx context.Context, c *SupervisorConfigs) (Supervisor, error) { @@ -167,8 +169,11 @@ func (s *supervisorImpl) Serve() (err error) { defer func() { s.Stop() // make sure context is canceled - if fatal := s.fatal.Load(); fatal != nil { - err = errors.Join(err, *fatal) + + s.fatalMux.RLock() + defer s.fatalMux.RUnlock() + if s.fatalErr != nil { + err = errors.Join(err, s.fatalErr) } }() @@ -179,7 +184,7 @@ func (s *supervisorImpl) Serve() (err error) { s.logger.Info("Supervised services started") - stopSvcCh := make(chan struct{}, len(s.services)) + svcErrCh := make(chan error, len(s.services)) for _, svc := range s.services { go func() { s.logger.Info("Starting subservice", "subservice", svc) @@ -188,33 +193,62 @@ func (s *supervisorImpl) Serve() (err error) { switch { case unexpected: s.logger.Error("Subservice stopped unexpectedly, shutting down", - "service", svc, + "subservice", svc, "err", svcErr, ) - // Only the Stop winner writes err; stopSvcCh joins that write. - err = ErrServiceStopped - case svcErr == nil || errors.Is(svcErr, context.Canceled): + if svcErr == nil { + svcErr = ErrServiceStopped + } + case svcErr == nil || isCancellationOnly(svcErr): s.logger.Info("Subservice stopped", "subservice", svc, ) + svcErr = nil default: + // Non-cancellation drain failures, including deadlines, fail shutdown. s.logger.Warn("Subservice failed during shutting down", "subservice", svc, "err", svcErr, ) } - stopSvcCh <- struct{}{} + svcErrCh <- svcErr }() } - // wait for all services to terminate + // Join every service and aggregate errors only in the supervisor goroutine. + errs := make([]error, 0, len(s.services)) for range s.services { - <-stopSvcCh + if err := <-svcErrCh; err != nil { + errs = append(errs, err) + } } s.logger.Info("Supervisor terminated") - return err + return errors.Join(errs...) +} + +// isCancellationOnly checks every leaf before suppressing a shutdown error. +// Matching the whole tree with errors.Is would hide failures joined with cancellation. +func isCancellationOnly(err error) bool { + switch wrapped := err.(type) { + case interface{ Unwrap() []error }: + children := wrapped.Unwrap() + if len(children) == 0 { + return errors.Is(err, context.Canceled) + } + for _, child := range children { + if !isCancellationOnly(child) { + return false + } + } + return true + case interface{ Unwrap() error }: + if child := wrapped.Unwrap(); child != nil { + return isCancellationOnly(child) + } + } + return errors.Is(err, context.Canceled) } // Fatal publishes the failure before cancellation so Serve observes it even if @@ -223,7 +257,11 @@ func (s *supervisorImpl) Fatal(err error) { if err == nil { err = ErrServiceStopped } - s.fatal.CompareAndSwap(nil, &err) + s.fatalMux.Lock() + defer s.fatalMux.Unlock() + + s.fatalErr = errors.Join(s.fatalErr, err) + s.Stop() } diff --git a/pkg/service/supervisor_test.go b/pkg/service/supervisor_test.go index fd388e577..2b9415430 100644 --- a/pkg/service/supervisor_test.go +++ b/pkg/service/supervisor_test.go @@ -9,6 +9,7 @@ import ( "errors" "fmt" "log/slog" + "sync" "testing" "time" @@ -187,7 +188,8 @@ func (s *SupervisorSuite) TestItLogsServiceErrors() { select { case err := <-errCh: - require.ErrorIs(s.T(), err, ErrServiceStopped) + require.ErrorContains(s.T(), err, "oops by child-1") + require.ErrorContains(s.T(), err, "oops by child-2") logged := buf.String() require.Contains(s.T(), logged, "oops by child-1") require.Contains(s.T(), logged, "oops by child-2") @@ -215,7 +217,7 @@ func (s *SupervisorSuite) TestItStopsWhenOnServiceError() { select { case err := <-errCh: - require.ErrorIs(s.T(), err, ErrServiceStopped) + require.ErrorContains(s.T(), err, "oops by error-child") logged := buf.String() require.Contains(s.T(), logged, "oops by error-child") case <-time.After(2 * time.Second): @@ -443,7 +445,8 @@ func TestFatalShutdownPreservesCauseAndWaitsForServices(t *testing.T) { require.True(t, waitCh(child.started)) cause := errors.New("unconfirmed advance result") sup.Fatal(cause) - sup.Fatal(errors.New("later failure")) + later := errors.New("later failure") + sup.Fatal(later) require.False(t, sup.Alive()) select { case <-done: @@ -454,6 +457,7 @@ func TestFatalShutdownPreservesCauseAndWaitsForServices(t *testing.T) { select { case err := <-done: require.ErrorIs(t, err, cause) + require.ErrorIs(t, err, later) case <-time.After(2 * time.Second): t.Fatal("Serve did not finish shutdown") } @@ -484,3 +488,159 @@ func TestSupervisorRejectsEmptyFactoryList(t *testing.T) { require.ErrorContains(t, err, "at least one service factory") } } + +func TestConcurrentFatalCallsPreserveAllCauses(t *testing.T) { + child := newTestService("child") + child.serveDone = make(chan struct{}) + sup, err := NewSupervisor(t.Context(), &SupervisorConfigs{ + BaseConfigs: BaseConfigs{Logger: discardLogger()}, + Factories: []FactoryFunction{func(context.Context, Supervisor) (SupervisedService, error) { return child, nil }}, + }) + require.NoError(t, err) + t.Cleanup(func() { sup.Stop(); close(child.serveDone) }) + done := make(chan error, 1) + go func() { done <- sup.Serve() }() + require.True(t, waitCh(child.started)) + + causes := make([]error, 64) + start := make(chan struct{}) + var calls sync.WaitGroup + for i := range causes { + if i != 0 { + causes[i] = fmt.Errorf("fatal failure %d", i) + } + calls.Add(1) + go func(cause error) { + defer calls.Done() + <-start + sup.Fatal(cause) + }(causes[i]) + } + close(start) + calls.Wait() + require.False(t, sup.Alive()) + child.serveDone <- struct{}{} + select { + case err := <-done: + require.ErrorIs(t, err, ErrServiceStopped, "nil Fatal calls retain the default cause") + for _, cause := range causes[1:] { + require.ErrorIs(t, err, cause) + } + case <-time.After(2 * time.Second): + t.Fatal("Serve did not finish shutdown") + } +} + +func (s *SupervisorSuite) TestUnexpectedExitPreservesServiceAndCause() { + for _, cause := range []error{nil, errors.New("bind failed"), context.Canceled} { + s.Run(fmt.Sprintf("cause=%v", cause), func() { + child := newTestService("unexpected-child") + child.duration = 0 + child.err = cause + sup, err := s.newSupervisor(s.T(), discardLogger(), child) + s.Require().NoError(err) + err = sup.Serve() + if cause == nil { + s.Require().ErrorIs(err, ErrServiceStopped) + } else { + s.Require().ErrorIs(err, cause) + } + }) + } +} + +func (s *SupervisorSuite) TestConcurrentDrainErrorsPreserveAllCauses() { + for _, fatal := range []bool{false, true} { + s.Run(fmt.Sprintf("fatal=%v", fatal), func() { + children := make([]*testServiceImpl, 16) + release := make(chan struct{}) + for i := range children { + child := newTestService(fmt.Sprintf("draining-child-%d", i)) + child.err = fmt.Errorf("drain failure %d", i) + child.serveDone = release + children[i] = child + } + children[0].err = context.DeadlineExceeded + children[1].err = fmt.Errorf("normal shutdown: %w", context.Canceled) + sup, err := s.newSupervisor(s.T(), discardLogger(), children...) + s.Require().NoError(err) + s.T().Cleanup(func() { sup.Stop() }) + done := make(chan error, 1) + go func() { done <- sup.Serve() }() + for _, child := range children { + s.Require().True(waitCh(child.started)) + } + fatalCause := errors.New("fatal trigger") + if fatal { + sup.Fatal(fatalCause) + } else { + sup.Stop() + } + close(release) + select { + case err := <-done: + s.Require().NotErrorIs(err, context.Canceled) + if fatal { + s.Require().ErrorIs(err, fatalCause) + } + for i, child := range children { + if i == 1 { + continue + } + s.Require().ErrorIs(err, child.err) + } + case <-time.After(2 * time.Second): + s.T().Fatal("supervisor did not finish draining") + } + }) + } +} + +func (s *SupervisorSuite) TestShutdownSuppressesOnlyCancellationErrors() { + storageErr := errors.New("storage flush failed") + otherErr := errors.New("connection close failed") + wrappedCancellation := fmt.Errorf("shutdown: %w", context.Canceled) + for _, tc := range []struct { + name string + err error + want []error + }{ + {name: "nil"}, + {name: "cancellation", err: context.Canceled}, + {name: "wrapped cancellation", err: wrappedCancellation}, + {name: "joined cancellations", err: errors.Join(context.Canceled, wrappedCancellation)}, + {name: "wrapped joined cancellations", err: fmt.Errorf("drain: %w", errors.Join(context.Canceled, wrappedCancellation))}, + {name: "storage failure", err: storageErr, want: []error{storageErr}}, + {name: "deadline", err: context.DeadlineExceeded, want: []error{context.DeadlineExceeded}}, + {name: "joined storage failure", err: errors.Join(context.Canceled, storageErr), want: []error{storageErr}}, + {name: "joined deadline", err: errors.Join(context.Canceled, context.DeadlineExceeded), want: []error{context.DeadlineExceeded}}, + {name: "nested failures", err: fmt.Errorf("drain: %w", errors.Join(wrappedCancellation, errors.Join(storageErr, otherErr))), want: []error{storageErr, otherErr}}, + {name: "multiple wrapped causes", err: fmt.Errorf("drain: %w; storage: %w", context.Canceled, storageErr), want: []error{storageErr}}, + } { + s.Run(tc.name, func() { + child := newTestService("draining-child") + child.err = tc.err + sup, err := s.newSupervisor(s.T(), discardLogger(), child) + s.Require().NoError(err) + s.T().Cleanup(func() { sup.Stop() }) + done := make(chan error, 1) + go func() { done <- sup.Serve() }() + s.Require().True(waitCh(child.started)) + sup.Stop() + select { + case err := <-done: + if len(tc.want) == 0 { + s.Require().NoError(err) + } else { + for _, cause := range tc.want { + s.Require().ErrorIs(err, cause) + } + // Keep the original error and its diagnostic wrapping intact. + s.Require().ErrorIs(err, tc.err) + } + case <-time.After(2 * time.Second): + s.T().Fatal("supervisor did not finish shutdown") + } + }) + } +} diff --git a/pkg/service/tick.go b/pkg/service/tick.go index d5fcae5d9..4eaef6b95 100644 --- a/pkg/service/tick.go +++ b/pkg/service/tick.go @@ -14,6 +14,9 @@ import ( ) type TickImpl interface { + // Tick returns true when progress warrants immediate continuation, even if + // other work failed. Return false when no progress is possible so retries + // wait for the polling timer instead of spinning on an error. Tick(ctx context.Context) (bool, error) } @@ -68,9 +71,6 @@ func (s *TickServiceTemplate) tick(ctx context.Context) { "reschedule", reschedule, "error", err, ) - // Failed work must wait for the polling timer even if it requests - // an immediate retry, otherwise repeated errors can spin the CPU. - return } else { s.Logger.Debug("Tick", "duration", elapsed, diff --git a/pkg/service/tick_test.go b/pkg/service/tick_test.go index 2fa8c5ad7..e868e06cd 100644 --- a/pkg/service/tick_test.go +++ b/pkg/service/tick_test.go @@ -186,7 +186,7 @@ type tickFunc func(context.Context) (bool, error) func (f tickFunc) Tick(ctx context.Context) (bool, error) { return f(ctx) } -func TestTickStopsImmediateReschedulingOnError(t *testing.T) { +func TestTickHonorsRescheduleWithPartialProgress(t *testing.T) { for _, fail := range []bool{false, true} { t.Run(fmt.Sprintf("error=%v", fail), func(t *testing.T) { calls := 0 @@ -194,17 +194,13 @@ func TestTickStopsImmediateReschedulingOnError(t *testing.T) { s.tickImpl = tickFunc(func(context.Context) (bool, error) { calls++ if fail { - // Bound even a regressed implementation so this test cannot spin. + // Some work progresses despite another work item failing. return calls < 3, errors.New("retryable failure") } return calls < 3, nil }) s.tick(t.Context()) - if fail { - require.Equal(t, 1, calls, "an error must end the immediate reschedule loop") - } else { - require.Equal(t, 3, calls, "successful work should reschedule immediately") - } + require.Equal(t, 3, calls, "progress should reschedule immediately despite errors") }) } } @@ -221,7 +217,7 @@ func TestServeRetriesFailedTickOnPollingTimer(t *testing.T) { s.tickImpl = tickFunc(func(context.Context) (bool, error) { attempts = append(attempts, time.Now()) if len(attempts) == 1 { - return true, errors.New("retryable failure") + return false, errors.New("retryable failure") } cancel() return false, nil @@ -229,5 +225,24 @@ func TestServeRetriesFailedTickOnPollingTimer(t *testing.T) { require.ErrorIs(t, s.Serve(ctx), context.Canceled) require.Len(t, attempts, 2) require.GreaterOrEqual(t, attempts[1].Sub(attempts[0]), interval/2, - "a failed tick must fall back to the polling timer") + "a failed tick without progress must fall back to the polling timer") +} + +func TestServeReschedulesPartialProgressDespiteError(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + calls := 0 + s := &TickServiceTemplate{ + BaseTemplate: BaseTemplate{Logger: discardLogger()}, + interval: time.Hour, + } + s.tickImpl = tickFunc(func(context.Context) (bool, error) { + calls++ + if calls == 3 { + cancel() + } + return true, errors.New("another work item failed") + }) + require.ErrorIs(t, s.Serve(ctx), context.Canceled) + require.Equal(t, 3, calls, "partial progress must continue immediately and stop on cancellation") }