From 8a8fd7c9c2572cee95b7db169c4fb226ddb87f90 Mon Sep 17 00:00:00 2001 From: Blake Gentry Date: Sun, 4 Oct 2026 17:04:57 -0500 Subject: [PATCH 1/2] add job completion concurrency capabilities Add optional `ExecutorJobCompletionConcurrency` and `PilotJobCompletionConcurrency` interfaces through which an executor and a pilot can declare how many `JobSetStateIfRunningMany` calls they can safely run at once. Implementations that don't implement them are treated as serial. The pgx and `database/sql` pool executors and the standard pilot report two. Their transaction and subtransaction executors report one, since a single transaction can't run statements concurrently. SQLite and other drivers don't opt in. A new driver test runs the advertised number of disjoint completion batches concurrently through a pool executor and checks that a transaction never claims more than one. --- riverdriver/river_driver_interface.go | 9 +++ .../river_database_sql_driver.go | 12 +++ riverdriver/riverdrivertest/executor_tx.go | 80 +++++++++++++++++++ riverdriver/riverpgxv5/river_pgx_v5_driver.go | 8 ++ rivershared/riverpilot/pilot.go | 9 +++ rivershared/riverpilot/standard_pilot.go | 4 + 6 files changed, 122 insertions(+) diff --git a/riverdriver/river_driver_interface.go b/riverdriver/river_driver_interface.go index dde894b89..2641016db 100644 --- a/riverdriver/river_driver_interface.go +++ b/riverdriver/river_driver_interface.go @@ -346,6 +346,15 @@ type Executor interface { TableTruncate(ctx context.Context, params *TableTruncateParams) error } +// ExecutorJobCompletionConcurrency is implemented by executors that can safely +// run multiple JobSetStateIfRunningMany calls concurrently. Executors without +// this capability use one completion call at a time. +// +// API is not stable. DO NOT IMPLEMENT. +type ExecutorJobCompletionConcurrency interface { + JobSetStateIfRunningManyConcurrency() int +} + // ExecutorTx is an executor which is a transaction. In addition to standard // Executor operations, it may be committed or rolled back. // diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver.go b/riverdriver/riverdatabasesql/river_database_sql_driver.go index 0671e852f..5331ee84c 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver.go @@ -779,6 +779,10 @@ func (e *Executor) JobSetStateIfRunningMany(ctx context.Context, params *riverdr return jobRowsFromInternalPartial(jobs), nil } +// JobSetStateIfRunningManyConcurrency returns the number of completion calls +// that can safely run concurrently through a database/sql pool. +func (e *Executor) JobSetStateIfRunningManyConcurrency() int { return 2 } + func (e *Executor) JobUpdate(ctx context.Context, params *riverdriver.JobUpdateParams) (*rivertype.JobRow, error) { metadata := params.Metadata if metadata == nil { @@ -1207,6 +1211,10 @@ func (t *ExecutorTx) Commit(ctx context.Context) error { return t.tx.Commit() } +// JobSetStateIfRunningManyConcurrency overrides the embedded Executor's value +// because a single transaction can't run statements concurrently. +func (t *ExecutorTx) JobSetStateIfRunningManyConcurrency() int { return 1 } + func (t *ExecutorTx) Rollback(ctx context.Context) error { // unfortunately, `database/sql` does not take a context ... return t.tx.Rollback() @@ -1251,6 +1259,10 @@ func (t *ExecutorSubTx) Commit(ctx context.Context) error { return nil } +// JobSetStateIfRunningManyConcurrency overrides the embedded Executor's value +// because a single transaction can't run statements concurrently. +func (t *ExecutorSubTx) JobSetStateIfRunningManyConcurrency() int { return 1 } + func (t *ExecutorSubTx) Rollback(ctx context.Context) error { defer t.beginOnce.Done() diff --git a/riverdriver/riverdrivertest/executor_tx.go b/riverdriver/riverdrivertest/executor_tx.go index 8d43b60df..a74461726 100644 --- a/riverdriver/riverdrivertest/executor_tx.go +++ b/riverdriver/riverdrivertest/executor_tx.go @@ -216,6 +216,86 @@ func exerciseExecutorTx[TTx any](ctx context.Context, t *testing.T, }) }) + t.Run("JobSetStateIfRunningManyConcurrency", func(t *testing.T) { + t.Parallel() + + completeManyParams := func(schema string, now time.Time, ids ...int64) *riverdriver.JobSetStateIfRunningManyParams { + params := &riverdriver.JobSetStateIfRunningManyParams{Now: &now, Schema: schema} + for _, id := range ids { + params.ID = append(params.ID, id) + params.Attempt = append(params.Attempt, nil) + params.ErrData = append(params.ErrData, nil) + params.FinalizedAt = append(params.FinalizedAt, &now) + params.MetadataDoMerge = append(params.MetadataDoMerge, false) + params.MetadataUpdates = append(params.MetadataUpdates, nil) + params.ScheduledAt = append(params.ScheduledAt, nil) + params.State = append(params.State, rivertype.JobStateCompleted) + } + return params + } + + t.Run("PoolExecutorRunsConcurrentBatches", func(t *testing.T) { + t.Parallel() + + driver, schema := driverWithSchema(ctx, t, nil) + exec := driver.GetExecutor() + + concurrency := 1 + if capability, ok := exec.(riverdriver.ExecutorJobCompletionConcurrency); ok { + concurrency = capability.JobSetStateIfRunningManyConcurrency() + } + require.GreaterOrEqual(t, concurrency, 1) + require.LessOrEqual(t, concurrency, 2) + + var ( + now = time.Now().UTC() + jobs = make([]*rivertype.JobRow, 2*concurrency) + ) + for i := range jobs { + jobs[i] = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Schema: schema, State: new(rivertype.JobStateRunning)}) + } + + // Run the advertised number of disjoint batches at once. + var ( + errs = make(chan error, concurrency) + wg sync.WaitGroup + ) + for i := range concurrency { + wg.Go(func() { + _, err := exec.JobSetStateIfRunningMany(ctx, completeManyParams(schema, now, jobs[2*i].ID, jobs[2*i+1].ID)) + errs <- err + }) + } + wg.Wait() + close(errs) + for err := range errs { + require.NoError(t, err) + } + + for _, job := range jobs { + updatedJob, err := exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ID: job.ID, Schema: schema}) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCompleted, updatedJob.State) + } + }) + + t.Run("TransactionExecutorIsSerial", func(t *testing.T) { + t.Parallel() + + exec := setup(ctx, t) + + tx, err := exec.Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback(ctx) }) + + // A transaction can't run statements concurrently, even when the + // pool executor it came from can. + if capability, ok := tx.(riverdriver.ExecutorJobCompletionConcurrency); ok { + require.Equal(t, 1, capability.JobSetStateIfRunningManyConcurrency()) + } + }) + }) + t.Run("PGAdvisoryXactLock", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver.go b/riverdriver/riverpgxv5/river_pgx_v5_driver.go index 71d56fef6..3b0b96d24 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver.go @@ -724,6 +724,10 @@ func (e *Executor) JobSetStateIfRunningMany(ctx context.Context, params *riverdr return jobRowsFromInternalPartial(jobs), nil } +// JobSetStateIfRunningManyConcurrency returns the number of completion calls +// that can safely run concurrently through a pgx pool. +func (e *Executor) JobSetStateIfRunningManyConcurrency() int { return 2 } + func (e *Executor) JobUpdate(ctx context.Context, params *riverdriver.JobUpdateParams) (*rivertype.JobRow, error) { metadata := params.Metadata if metadata == nil { @@ -1142,6 +1146,10 @@ func (t *ExecutorTx) Commit(ctx context.Context) error { return t.tx.Commit(ctx) } +// JobSetStateIfRunningManyConcurrency overrides the embedded Executor's value +// because a single transaction can't run statements concurrently. +func (t *ExecutorTx) JobSetStateIfRunningManyConcurrency() int { return 1 } + func (t *ExecutorTx) Rollback(ctx context.Context) error { return t.tx.Rollback(ctx) } diff --git a/rivershared/riverpilot/pilot.go b/rivershared/riverpilot/pilot.go index 25108d29e..9a3237101 100644 --- a/rivershared/riverpilot/pilot.go +++ b/rivershared/riverpilot/pilot.go @@ -101,6 +101,15 @@ func (p *PilotInitParams) Validate() *PilotInitParams { return p } +// PilotJobCompletionConcurrency is implemented by pilots whose completion +// logic can safely run multiple JobSetStateIfRunningMany calls concurrently. +// Pilots without this capability use one completion call at a time. +// +// API is not stable. DO NOT USE. +type PilotJobCompletionConcurrency interface { + JobSetStateIfRunningManyConcurrency() int +} + // PilotJobRescuer contains optional Pilot functionality related to rescuing // stuck jobs. Pilots that don't implement it fall back to the standard // executor-backed behavior. diff --git a/rivershared/riverpilot/standard_pilot.go b/rivershared/riverpilot/standard_pilot.go index 2e9b9aec5..744c0f3cd 100644 --- a/rivershared/riverpilot/standard_pilot.go +++ b/rivershared/riverpilot/standard_pilot.go @@ -55,6 +55,10 @@ func (p *StandardPilot) JobSetStateIfRunningMany(ctx context.Context, exec river return exec.JobSetStateIfRunningMany(ctx, params) } +// JobSetStateIfRunningManyConcurrency returns the number of completion calls +// that can safely run concurrently through the standard pilot. +func (p *StandardPilot) JobSetStateIfRunningManyConcurrency() int { return 2 } + func (p *StandardPilot) PeriodicJobKeepAliveAndReap(ctx context.Context, exec riverdriver.Executor, params *PeriodicJobKeepAliveAndReapParams) ([]*PeriodicJob, error) { return nil, nil } From adba3de0f9413a331c6926ddf17ed68e2148b4cb Mon Sep 17 00:00:00 2001 From: Blake Gentry Date: Sun, 4 Oct 2026 17:04:57 -0500 Subject: [PATCH 2/2] parallelize full completion batches The batch completer persists one `JobSetStateIfRunningMany` call at a time. Under sustained load a second full backlog builds while the first query is still in flight, so producers wait on the backlog and workers sit idle even though the database has spare capacity. When both the executor and the pilot declare completion concurrency, run up to two completion queries at once, each capped at the existing 5,000-job batch size. A second query starts only once its own backlog is full, so sparse traffic keeps the existing coalescing and connection use. SQLite, transactions, and pilots that don't opt in stay serial. Concurrent batches could otherwise hold different executions of the same job, and PostgreSQL may lock their rows out of submission order, letting an older result overwrite a newer attempt. The completer tracks in-flight and retrying job IDs so a later result for the same job waits behind the earlier batch while independent batches still persist concurrently. Shutdown waits for active queries and makes one bounded final flush attempt. Subscribers may see completion events for different jobs in a different order than before, and the completer may briefly use one more database connection. Tests block the first query and check that a sparse second batch stays queued through a flush interval and starts as soon as it reaches the batch threshold, cover capability negotiation, and force a duplicate job ID to interleave across batches. --- CHANGELOG.md | 4 + internal/jobcompleter/job_completer.go | 303 ++++++++++++++------ internal/jobcompleter/job_completer_test.go | 298 ++++++++++++++++++- 3 files changed, 510 insertions(+), 95 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c7dc90d86..fd76b9e80 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- When both the driver and pilot support it, which the PostgreSQL drivers do, the batch job completer now persists up to two full batches of job completions concurrently, improving completion throughput under heavy load. Results for the same job are still applied in order, but subscribers may receive completion events for different jobs in a different order than before, and the completer may briefly use one additional database connection. [PR #1434](https://github.com/riverqueue/river/pull/1434). + ## [0.49.0] - 2026-10-05 ⚠️ If using River Pro, make sure to upgrade it to at least River Pro v0.32.0 to get a compatible package. diff --git a/internal/jobcompleter/job_completer.go b/internal/jobcompleter/job_completer.go index 44e17a8f8..2e075d701 100644 --- a/internal/jobcompleter/job_completer.go +++ b/internal/jobcompleter/job_completer.go @@ -326,6 +326,24 @@ type batchCompleterSetState struct { Stats *jobstats.JobStatistics } +func completionConcurrency(exec riverdriver.Executor, pilot riverpilot.Pilot) int { + const maxConcurrency = 2 + + execProvider, ok := exec.(riverdriver.ExecutorJobCompletionConcurrency) + if !ok { + return 1 + } + pilotProvider, ok := pilot.(riverpilot.PilotJobCompletionConcurrency) + if !ok { + return 1 + } + return min( + maxConcurrency, + max(execProvider.JobSetStateIfRunningManyConcurrency(), 1), + max(pilotProvider.JobSetStateIfRunningManyConcurrency(), 1), + ) +} + // BatchCompleter accumulates incoming completions, and instead of completing // them immediately, every so often complete many of them as a single efficient // batch. To minimize the amount of driver surface area we need, the batching is @@ -336,19 +354,22 @@ type BatchCompleter struct { baseservice.BaseService startstop.BaseStartStop - backlogWaitThreshold int // configurable for testing purposes; backlog at which completions start waiting for the completer to catch up - batchReadyChan chan struct{} - completionMaxSize int // configurable for testing purposes; max jobs to complete in single database operation - disableSleep bool // disable sleep in testing - maxBacklog int // configurable for testing purposes; emergency backlog threshold where a warning is logged - exec riverdriver.Executor - pilot riverpilot.Pilot - schema string - setStateParams map[int64]batchCompleterSetState - setStateParamsMu sync.RWMutex - subscribeCh SubscribeChan - waitOnBacklogChan chan struct{} - waitOnBacklogWaiting bool + backlogWaitThreshold int // configurable for testing purposes; backlog at which completions start waiting for the completer to catch up + batchReadyChan chan struct{} + completionConcurrency int // configurable for testing purposes; max concurrent database completion batches + completionMaxSize int // configurable for testing purposes; max jobs to complete in single database operation + deferredSetStateParams map[int64]batchCompleterSetState + disableSleep bool // disable sleep in testing + maxBacklog int // configurable for testing purposes; emergency backlog threshold where a warning is logged + exec riverdriver.Executor + inFlightIDs map[int64]struct{} + pilot riverpilot.Pilot + schema string + setStateParams map[int64]batchCompleterSetState + setStateParamsMu sync.RWMutex + subscribeCh SubscribeChan + waitOnBacklogChan chan struct{} + waitOnBacklogWaiting bool } func NewBatchCompleter(archetype *baseservice.Archetype, schema string, exec riverdriver.Executor, pilot riverpilot.Pilot, subscribeCh SubscribeChan) *BatchCompleter { @@ -359,15 +380,18 @@ func NewBatchCompleter(archetype *baseservice.Archetype, schema string, exec riv ) return baseservice.Init(archetype, &BatchCompleter{ - backlogWaitThreshold: backlogWaitThreshold, - batchReadyChan: make(chan struct{}, 1), - completionMaxSize: completionMaxSize, - exec: exec, - maxBacklog: maxBacklog, - pilot: pilot, - schema: schema, - setStateParams: make(map[int64]batchCompleterSetState), - subscribeCh: subscribeCh, + backlogWaitThreshold: backlogWaitThreshold, + batchReadyChan: make(chan struct{}, 1), + completionConcurrency: completionConcurrency(exec, pilot), + completionMaxSize: completionMaxSize, + deferredSetStateParams: make(map[int64]batchCompleterSetState), + exec: exec, + inFlightIDs: make(map[int64]struct{}), + maxBacklog: maxBacklog, + pilot: pilot, + schema: schema, + setStateParams: make(map[int64]batchCompleterSetState), + subscribeCh: subscribeCh, }) } @@ -395,83 +419,148 @@ func (c *BatchCompleter) Start(ctx context.Context) error { ticker := time.NewTicker(50 * time.Millisecond) defer ticker.Stop() + batchDoneChan := make(chan error, max(c.completionConcurrency, 1)) + numInFlight := 0 - backlogSize := func() int { + readyBacklogSize := func() int { c.setStateParamsMu.RLock() defer c.setStateParamsMu.RUnlock() return len(c.setStateParams) } - for numTicks := 0; ; numTicks++ { + startBatch := func() bool { + setStateBatch := c.takeBatch(c.completionMaxSize) + if len(setStateBatch) == 0 { + return false + } + + numInFlight++ + go func() { + batchDoneChan <- c.handleSetStateBatch(ctx, setStateBatch) + }() + return true + } + + const batchCompleterStartThreshold = 100 + startReadyBatches := func(force bool) { + for numInFlight < max(c.completionConcurrency, 1) { + backlogSize := readyBacklogSize() + if backlogSize == 0 { + return + } + if numInFlight > 0 { + // Preserve a full-size query for the concurrent path. Sparse + // completions continue coalescing until the active query exits. + if backlogSize < c.batchReadyThreshold() { + return + } + } else if backlogSize < min(c.backlogWaitThresholdEffective(), batchCompleterStartThreshold) && !force { + return + } + if !startBatch() { + return + } + force = false + } + } + + logBatchError := func(err error) { + if err != nil { + c.Logger.ErrorContext(ctx, c.Name+": Error completing batch", "err", err) + } + } + + numTicks := 0 + for { select { case <-stopCtx.Done(): - // Try to insert last batch before leaving. Note we use the - // original context so operations aren't immediately cancelled. - if err := c.handleBatch(ctx); err != nil { - c.Logger.ErrorContext(ctx, c.Name+": Error completing batch", "err", err) + // Finish active queries, then flush any deferred per-job results. + // Keep using the original context so operations aren't immediately + // cancelled by the service stop context. Stop on the first error so + // a requeued batch can't make shutdown retry indefinitely. + for numInFlight > 0 { + logBatchError(<-batchDoneChan) + numInFlight-- + } + for { + err := c.handleBatch(ctx) + logBatchError(err) + if err != nil || readyBacklogSize() == 0 { + break + } } return case <-c.batchReadyChan: + startReadyBatches(false) case <-ticker.C: + // The ticker fires quite often to make sure that given a huge + // glut of jobs, we don't accidentally build up too much of a + // backlog by waiting too long. However, don't start a complete + // operation until we reach a minimum threshold unless this is a + // periodic flush tick. Sparse second batches are never forced to + // run concurrently with a first batch. + force := numTicks == 0 || numTicks%5 == 0 + startReadyBatches(force) + numTicks++ + case err := <-batchDoneChan: + numInFlight-- + logBatchError(err) + startReadyBatches(false) } + } + }() - // The ticker fires quite often to make sure that given a huge glut - // of jobs, we don't accidentally build up too much of a backlog by - // waiting too long. However, don't start a complete operation until - // we reach a minimum threshold unless we're on a tick that's a - // multiple of 5. So, jobs will be completed every 250ms even if the - // threshold hasn't been met. - const batchCompleterStartThreshold = 100 - if backlogSize() < min(c.backlogWaitThresholdEffective(), batchCompleterStartThreshold) && numTicks != 0 && numTicks%5 != 0 { - continue - } + return nil +} - for { - if err := c.handleBatch(ctx); err != nil { - c.Logger.ErrorContext(ctx, c.Name+": Error completing batch", "err", err) - } +func (c *BatchCompleter) takeBatch(maxSize int) map[int64]batchCompleterSetState { + c.setStateParamsMu.Lock() + defer c.setStateParamsMu.Unlock() - // New jobs to complete may have come in while working the batch - // above. If enough have to bring us above the minimum complete - // threshold, loop again and do another batch. Otherwise, break - // and listen for a new tick. - if backlogSize() < batchCompleterStartThreshold { - break - } - } + if len(c.setStateParams) == 0 { + return nil + } + if len(c.inFlightIDs) == 0 && (maxSize <= 0 || len(c.setStateParams) <= maxSize) { + setStateBatch := c.setStateParams + c.setStateParams = make(map[int64]batchCompleterSetState) + for id := range setStateBatch { + c.inFlightIDs[id] = struct{}{} } - }() + return setStateBatch + } - return nil + batchCapacity := len(c.setStateParams) + if maxSize > 0 { + batchCapacity = min(batchCapacity, maxSize) + } + setStateBatch := make(map[int64]batchCompleterSetState, batchCapacity) + for id, setState := range c.setStateParams { + if _, inFlight := c.inFlightIDs[id]; inFlight { + continue + } + setStateBatch[id] = setState + delete(c.setStateParams, id) + c.inFlightIDs[id] = struct{}{} + if maxSize > 0 && len(setStateBatch) == maxSize { + break + } + } + return setStateBatch } func (c *BatchCompleter) handleBatch(ctx context.Context) error { - var setStateBatch map[int64]batchCompleterSetState - func() { - c.setStateParamsMu.Lock() - defer c.setStateParamsMu.Unlock() - - setStateBatch = c.setStateParams - - // Don't bother resetting the map if there's nothing to process, - // allowing the completer to idle efficiently. - if len(setStateBatch) > 0 { - c.setStateParams = make(map[int64]batchCompleterSetState) - } else { - // Set nil to avoid a data race below in case the map is set as a - // new job comes in. - setStateBatch = nil - } - }() + return c.handleSetStateBatch(ctx, c.takeBatch(0)) +} +func (c *BatchCompleter) handleSetStateBatch(ctx context.Context, setStateBatch map[int64]batchCompleterSetState) error { if len(setStateBatch) < 1 { return nil } handleBatchError := func(err error) error { if isNonRetryableCompleterError(err) { - c.releaseBacklogWaitIfReady(ctx) + c.finishBatch(ctx, setStateBatch) return err } @@ -580,26 +669,32 @@ func (c *BatchCompleter) handleBatch(ctx context.Context) error { if len(events) > 0 { c.subscribeCh <- events } - - func() { - c.setStateParamsMu.Lock() - defer c.setStateParamsMu.Unlock() - - if c.waitOnBacklogWaiting && len(c.setStateParams) < c.backlogResumeThreshold() { - c.Logger.DebugContext(ctx, c.Name+": Disabling waitOnBacklog; ready to complete more jobs") - close(c.waitOnBacklogChan) - c.waitOnBacklogWaiting = false - } - }() + c.finishBatch(ctx, setStateBatch) return nil } -func (c *BatchCompleter) releaseBacklogWaitIfReady(ctx context.Context) { +func (c *BatchCompleter) finishBatch(ctx context.Context, setStateBatch map[int64]batchCompleterSetState) { c.setStateParamsMu.Lock() - defer c.setStateParamsMu.Unlock() + for id := range setStateBatch { + delete(c.inFlightIDs, id) + if deferred, exists := c.deferredSetStateParams[id]; exists { + c.setStateParams[id] = deferred + delete(c.deferredSetStateParams, id) + } + } + backlogSize := c.backlogSizeLocked() + c.releaseBacklogWaitIfReadyLocked(ctx, backlogSize) + readyBacklogSize := len(c.setStateParams) + c.setStateParamsMu.Unlock() + + if readyBacklogSize >= c.batchReadyThreshold() { + c.signalBatchReady() + } +} - if c.waitOnBacklogWaiting && len(c.setStateParams) < c.backlogResumeThreshold() { +func (c *BatchCompleter) releaseBacklogWaitIfReadyLocked(ctx context.Context, backlogSize int) { + if c.waitOnBacklogWaiting && backlogSize < c.backlogResumeThreshold() { c.Logger.DebugContext(ctx, c.Name+": Disabling waitOnBacklog; ready to complete more jobs") close(c.waitOnBacklogChan) c.waitOnBacklogWaiting = false @@ -609,20 +704,27 @@ func (c *BatchCompleter) releaseBacklogWaitIfReady(ctx context.Context) { func (c *BatchCompleter) requeueBatch(ctx context.Context, setStateBatch map[int64]batchCompleterSetState) { c.setStateParamsMu.Lock() for id, setState := range setStateBatch { - if _, exists := c.setStateParams[id]; exists { + delete(c.inFlightIDs, id) + + // A result that arrived for the same job while this batch was in + // flight comes from a later execution, so it supersedes the failed + // result rather than queueing behind it. Retrying the stale result + // first could overwrite the later attempt's running row and turn the + // newer result into a no-op. + if deferred, exists := c.deferredSetStateParams[id]; exists { + c.setStateParams[id] = deferred + delete(c.deferredSetStateParams, id) continue } + c.setStateParams[id] = setState } - backlogSize := len(c.setStateParams) - if c.waitOnBacklogWaiting && backlogSize < c.backlogResumeThreshold() { - c.Logger.DebugContext(ctx, c.Name+": Disabling waitOnBacklog; ready to complete more jobs") - close(c.waitOnBacklogChan) - c.waitOnBacklogWaiting = false - } + backlogSize := c.backlogSizeLocked() + c.releaseBacklogWaitIfReadyLocked(ctx, backlogSize) + readyBacklogSize := len(c.setStateParams) c.setStateParamsMu.Unlock() - if backlogSize >= c.batchReadyThreshold() { + if readyBacklogSize >= c.batchReadyThreshold() { c.signalBatchReady() } @@ -662,7 +764,7 @@ func (c *BatchCompleter) tryEnqueueSetState(ctx context.Context, now time.Time, } var ( - backlogSize = len(c.setStateParams) + backlogSize = c.backlogSizeLocked() waitAt = c.backlogWaitThresholdEffective() ) if backlogSize >= waitAt { @@ -670,11 +772,24 @@ func (c *BatchCompleter) tryEnqueueSetState(ctx context.Context, now time.Time, } statsSnapshot := *stats - c.setStateParams[params.ID] = batchCompleterSetState{Params: params, StartTime: now, Stats: &statsSnapshot} + setState := batchCompleterSetState{Params: params, StartTime: now, Stats: &statsSnapshot} + // A result for a job whose earlier result is still being persisted waits + // behind it so that concurrent batches can't apply the two out of order. + // A queued result that isn't in flight yet, including one requeued after + // a failed batch, is simply replaced by the newer one. + if _, inFlight := c.inFlightIDs[params.ID]; inFlight { + c.deferredSetStateParams[params.ID] = setState + } else { + c.setStateParams[params.ID] = setState + } return len(c.setStateParams), nil } +func (c *BatchCompleter) backlogSizeLocked() int { + return len(c.setStateParams) + len(c.deferredSetStateParams) +} + // backlogResumeThreshold returns the low-water mark below which waiting // completers are released. Keeping this below the wait threshold avoids rapidly // cycling between waiting and not waiting when the completer is near capacity. diff --git a/internal/jobcompleter/job_completer_test.go b/internal/jobcompleter/job_completer_test.go index c88dcd401..0890ec3a2 100644 --- a/internal/jobcompleter/job_completer_test.go +++ b/internal/jobcompleter/job_completer_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "slices" "sync" "sync/atomic" "testing" @@ -23,12 +24,14 @@ import ( "github.com/riverqueue/river/rivershared/riversharedtest" "github.com/riverqueue/river/rivershared/startstop" "github.com/riverqueue/river/rivershared/testfactory" + "github.com/riverqueue/river/rivershared/testsignal" "github.com/riverqueue/river/rivertype" ) type partialExecutorMock struct { riverdriver.Executor + CompletionConcurrency int JobSetStateIfRunningManyCalled bool JobSetStateIfRunningManyFunc func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) mu sync.Mutex @@ -37,10 +40,14 @@ type partialExecutorMock struct { // NewPartialExecutorMock returns a new mock with all mock functions set to call // down into the given real executor. func NewPartialExecutorMock(exec riverdriver.Executor) *partialExecutorMock { - return &partialExecutorMock{ + mock := &partialExecutorMock{ Executor: exec, JobSetStateIfRunningManyFunc: exec.JobSetStateIfRunningMany, } + if provider, ok := exec.(riverdriver.ExecutorJobCompletionConcurrency); ok { + mock.CompletionConcurrency = provider.JobSetStateIfRunningManyConcurrency() + } + return mock } func (m *partialExecutorMock) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { @@ -56,6 +63,10 @@ func (m *partialExecutorMock) JobSetStateIfRunningMany(ctx context.Context, para return m.JobSetStateIfRunningManyFunc(ctx, params) } +func (m *partialExecutorMock) JobSetStateIfRunningManyConcurrency() int { + return m.CompletionConcurrency +} + func (m *partialExecutorMock) setCalled(setCalledFunc func()) { m.mu.Lock() defer m.mu.Unlock() @@ -68,6 +79,14 @@ type partialExecutorTxMock struct { partial *partialExecutorMock } +type unspecifiedCompletionConcurrencyExecutor struct { + riverdriver.Executor +} + +type unspecifiedCompletionConcurrencyPilot struct { + riverpilot.Pilot +} + func (m *partialExecutorTxMock) JobSetStateIfRunningMany(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { return m.partial.JobSetStateIfRunningMany(ctx, params) } @@ -724,6 +743,184 @@ func TestBatchCompleter_BackpressureRequeuesBatchAfterCompletionFailure(t *testi require.NoError(t, riversharedtest.WaitOrTimeout(t, errCh)) } +func TestBatchCompleter_CompletionConcurrency(t *testing.T) { + t.Parallel() + + standardPilot := &riverpilot.StandardPilot{} + + t.Run("BothOptIn", func(t *testing.T) { + t.Parallel() + + exec := &partialExecutorMock{CompletionConcurrency: 2} + + require.Equal(t, 2, completionConcurrency(exec, standardPilot)) + }) + + t.Run("ExecutorDoesNotOptIn", func(t *testing.T) { + t.Parallel() + + exec := &unspecifiedCompletionConcurrencyExecutor{} + + require.Equal(t, 1, completionConcurrency(exec, standardPilot)) + }) + + t.Run("PilotDoesNotOptIn", func(t *testing.T) { + t.Parallel() + + exec := &partialExecutorMock{CompletionConcurrency: 2} + pilot := &unspecifiedCompletionConcurrencyPilot{Pilot: standardPilot} + + require.Equal(t, 1, completionConcurrency(exec, pilot)) + }) +} + +func TestBatchCompleter_ConcurrentBatchRequiresFullBacklog(t *testing.T) { + t.Parallel() + + ctx := context.Background() + var calls testsignal.TestSignal[int] + calls.Init(t) + release := make(chan struct{}, 2) + execMock := &partialExecutorMock{CompletionConcurrency: 2} + execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + calls.Signal(len(params.ID)) + <-release + + rows := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + rows[i] = &rivertype.JobRow{ID: id, State: params.State[i]} + } + return rows, nil + } + + subscribeCh := make(chan []CompleterJobUpdated, 2) + completer := NewBatchCompleter( + riversharedtest.BaseServiceArchetype(t), + "", + execMock, + &riverpilot.StandardPilot{}, + subscribeCh, + ) + completer.completionMaxSize = 2 + require.NoError(t, completer.Start(ctx)) + t.Cleanup(func() { + for range 2 { + select { + case release <- struct{}{}: + default: + } + } + completer.Stop() + }) + riversharedtest.WaitOrTimeout(t, completer.Started()) + + enqueue := func(id int64) { + t.Helper() + + require.NoError(t, completer.JobSetStateIfRunning( + ctx, + &jobstats.JobStatistics{}, + riverdriver.JobSetStateCompleted(id, time.Now(), nil), + )) + } + + enqueue(1) + enqueue(2) + require.Equal(t, 2, calls.WaitOrTimeout()) + + enqueue(3) + select { + case count := <-calls.WaitC(): + require.Failf(t, "unexpected sparse concurrent batch", "got batch size %d", count) + case <-time.After(300 * time.Millisecond): + } + + enqueue(4) + require.Equal(t, 2, calls.WaitOrTimeout()) + release <- struct{}{} + release <- struct{}{} + completer.Stop() +} + +func TestBatchCompleter_ConcurrentBatchesSerializeDuplicateJob(t *testing.T) { + t.Parallel() + + type call struct { + ids []int64 + release chan struct{} + } + + ctx := context.Background() + calls := make(chan call, 3) + releaseAll := make(chan struct{}) + execMock := &partialExecutorMock{CompletionConcurrency: 2} + execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + currentCall := call{ids: slices.Clone(params.ID), release: make(chan struct{})} + calls <- currentCall + select { + case <-currentCall.release: + case <-releaseAll: + } + + rows := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + rows[i] = &rivertype.JobRow{ID: id, State: params.State[i]} + } + return rows, nil + } + + subscribeCh := make(chan []CompleterJobUpdated, 3) + completer := NewBatchCompleter( + riversharedtest.BaseServiceArchetype(t), + "", + execMock, + &riverpilot.StandardPilot{}, + subscribeCh, + ) + completer.completionMaxSize = 2 + require.NoError(t, completer.Start(ctx)) + t.Cleanup(func() { + close(releaseAll) + completer.Stop() + }) + riversharedtest.WaitOrTimeout(t, completer.Started()) + + enqueue := func(id int64) { + t.Helper() + + require.NoError(t, completer.JobSetStateIfRunning( + ctx, + &jobstats.JobStatistics{}, + riverdriver.JobSetStateCompleted(id, time.Now(), nil), + )) + } + + enqueue(1) + enqueue(2) + firstCall := riversharedtest.WaitOrTimeout(t, calls) + require.ElementsMatch(t, []int64{1, 2}, firstCall.ids) + + enqueue(1) + enqueue(3) + enqueue(4) + secondCall := riversharedtest.WaitOrTimeout(t, calls) + require.ElementsMatch(t, []int64{3, 4}, secondCall.ids) + close(secondCall.release) + + select { + case unexpectedCall := <-calls: + close(unexpectedCall.release) + require.Failf(t, "duplicate job ran concurrently", "got IDs %v", unexpectedCall.ids) + case <-time.After(100 * time.Millisecond): + } + + close(firstCall.release) + thirdCall := riversharedtest.WaitOrTimeout(t, calls) + require.Equal(t, []int64{1}, thirdCall.ids) + close(thirdCall.release) + completer.Stop() +} + func TestBatchCompleter_NonRetryableCompletionFailureDoesNotRequeueBatch(t *testing.T) { t.Parallel() @@ -764,6 +961,105 @@ func TestBatchCompleter_NonRetryableCompletionFailureDoesNotRequeueBatch(t *test } } +func TestBatchCompleter_NewerResultSupersedesRequeuedResult(t *testing.T) { + t.Parallel() + + ctx := context.Background() + + type testBundle struct { + completer *BatchCompleter + execMock *partialExecutorMock + + // persisted records the state of each job ID in every set-state call + // that succeeded. + persisted []map[int64]rivertype.JobState + persistedMu sync.Mutex + } + + setup := func(t *testing.T, failCalls int, beforeFailedCall func(bundle *testBundle, callNum int)) *testBundle { + t.Helper() + + bundle := &testBundle{execMock: &partialExecutorMock{}} + + var numCalls int + bundle.execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + numCalls++ + if numCalls <= failCalls { + if beforeFailedCall != nil { + beforeFailedCall(bundle, numCalls) + } + return nil, errors.New("error from batch completion") + } + + states := make(map[int64]rivertype.JobState, len(params.ID)) + rows := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + states[id] = params.State[i] + rows[i] = &rivertype.JobRow{ID: id, State: params.State[i]} + } + + bundle.persistedMu.Lock() + bundle.persisted = append(bundle.persisted, states) + bundle.persistedMu.Unlock() + + return rows, nil + } + + bundle.completer = NewBatchCompleter(riversharedtest.BaseServiceArchetype(t), "", bundle.execMock, &riverpilot.StandardPilot{}, make(chan []CompleterJobUpdated, 10)) + bundle.completer.disableSleep = true + + return bundle + } + + setState := func(t *testing.T, completer *BatchCompleter, params *riverdriver.JobSetStateIfRunningParams) { + t.Helper() + + require.NoError(t, completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, params)) + } + + t.Run("NewerResultArrivesAfterRequeue", func(t *testing.T) { + t.Parallel() + + bundle := setup(t, numRetries, nil) + + setState(t, bundle.completer, riverdriver.JobSetStateCompleted(1, time.Now(), nil)) + require.Error(t, bundle.completer.handleBatch(ctx)) + + setState(t, bundle.completer, riverdriver.JobSetStateErrorRetryable(1, time.Now(), []byte(`{}`), nil)) + require.NoError(t, bundle.completer.handleBatch(ctx)) + + require.Equal(t, []map[int64]rivertype.JobState{{1: rivertype.JobStateRetryable}}, bundle.persisted) + }) + + t.Run("NewerResultArrivesWhileInFlight", func(t *testing.T) { + t.Parallel() + + // The newer result for job 1 arrives while the batch carrying the + // older result is failing, so it's deferred behind the in-flight batch. + bundle := setup(t, numRetries, func(bundle *testBundle, callNum int) { + if callNum == 1 { + setState(t, bundle.completer, riverdriver.JobSetStateErrorRetryable(1, time.Now(), []byte(`{}`), nil)) + } + }) + + setState(t, bundle.completer, riverdriver.JobSetStateCompleted(1, time.Now(), nil)) + setState(t, bundle.completer, riverdriver.JobSetStateCompleted(2, time.Now(), nil)) + require.Error(t, bundle.completer.handleBatch(ctx)) + + bundle.completer.setStateParamsMu.RLock() + require.Empty(t, bundle.completer.deferredSetStateParams) + require.Len(t, bundle.completer.setStateParams, 2) + bundle.completer.setStateParamsMu.RUnlock() + + require.NoError(t, bundle.completer.handleBatch(ctx)) + + require.Equal(t, []map[int64]rivertype.JobState{{ + 1: rivertype.JobStateRetryable, + 2: rivertype.JobStateCompleted, + }}, bundle.persisted) + }) +} + func TestBatchCompleter_JobStatsSnapshotsPerUpdate(t *testing.T) { t.Parallel()