From b60980b716e1e1c6970148f8f120e5ebe744246c Mon Sep 17 00:00:00 2001 From: Brandur Date: Fri, 25 Sep 2026 22:08:10 -0500 Subject: [PATCH] Add `Client.JobWaitFinalized` Adds a helper `Client.JobWaitFinalized` that polls waiting for a job to be finalized (completed, discarded, or cancelled) and returns it. The existing subscription functionality isn't appropriate for this because it only returns changes that occur within the existing client, so a job worked elsewhere wouldn't be returned. Implementation details: * Multiple calls to `JobWaitFinalized` share a single poll loop. * Polls every 250 ms. * A database and/or client change could add a notify signal that sends when a job is finalized, but that'd involve a more elaborate change, and would have downsides in that we'd have to be doing a lot of signaling even in case where no one's listening. --- CHANGELOG.md | 1 + client.go | 36 ++++ client_test.go | 226 ++++++++++++++++++++++++ job_waiter.go | 286 ++++++++++++++++++++++++++++++ job_waiter_test.go | 421 +++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 970 insertions(+) create mode 100644 job_waiter.go create mode 100644 job_waiter_test.go diff --git a/CHANGELOG.md b/CHANGELOG.md index 2e095285..0f20c00c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Added `Config.FetchOnlyKnownKinds` to restrict job fetching to registered worker kinds, including aliases. Clients with different workers can share a queue while leaving unknown jobs available without consuming attempts. Disabled by default; leader election and stuck-job rescue behavior are unchanged. [PR #1396](https://github.com/riverqueue/river/pull/1396). - Added `Config.LeaderElectionDisabled` to let a client work jobs without participating in leader election or running maintenance services. Other eligible clients in the same database and schema continue handling scheduling, retries, periodic enqueueing, rescue, and cleanup. [PR #1382](https://github.com/riverqueue/river/pull/1382). - Added support for YugabyteDB. When `LISTEN/NOTIFY` is unavailable or disabled, clients automatically poll for running job cancellations and queue pause, resume, and metadata changes, and skip unsupported notification broadcasts. This works with the default `PollOnly: false`. Native notifications require YugabyteDB 2025.2.3 or later with `ysql_yb_enable_listen_notify=true` on both Masters and TServers. [PR #1347](https://github.com/riverqueue/river/pull/1347). +- Added `Client.JobWaitFinalized` to wait for a job to reach a cancelled, completed, or discarded state using shared, batched polling. It works across clients and with clients that don't run workers, and supports context cancellation and a configurable poll interval through `JobWaitFinalizedOpts`. [PR #1393](https://github.com/riverqueue/river/pull/1393). ### Changed diff --git a/client.go b/client.go index 3387212e..34fda851 100644 --- a/client.go +++ b/client.go @@ -749,6 +749,7 @@ type Client[TTx any] struct { pluginLookupByJob *pluginlookup.JobPluginLookup pluginLookupGlobal *pluginlookup.PluginLookup insertNotifyLimiter *notifylimiter.Limiter + jobWaiter *jobWaiter notifier *notifier.Notifier // may be nil in poll-only mode periodicJobs *PeriodicJobBundle pilot riverpilot.Pilot @@ -881,6 +882,7 @@ func NewClient[TTx any](driver riverdriver.Driver[TTx], config *Config) (*Client }, config: config, driver: driver, + jobWaiter: newJobWaiter(driver.GetExecutor, config.Schema), pluginLookupByJob: pluginLookupByJob, pluginLookupGlobal: pluginLookupGlobal, producersByQueueName: make(map[string]*producer), @@ -1754,6 +1756,40 @@ func (c *Client[TTx]) jobUpdate(ctx context.Context, exec riverdriver.Executor, }) } +// JobWaitFinalized waits until the job with the given ID is observed in a +// finalized state (cancelled, completed, or discarded), and returns its persisted +// row. A cancelled or discarded job is returned with a nil error; callers should +// inspect State to determine the job's outcome. Retried and snoozed jobs continue +// waiting until they finalize. +// +// The job is checked immediately, then polled approximately every 250 milliseconds +// by default. Pass opts to customize the poll interval, or nil to use the defaults. +// Concurrent calls on the same client share a single poll loop that batches +// database queries for all jobs being waited on, using the shortest poll interval +// requested by an active call. The client does not need to be started or configured +// with workers, and the job may be executed by another client. Stop and StopAndCancel +// do not cancel waits. +// +// Cancelling ctx stops only this wait, without cancelling the job. Database +// errors are returned to the caller. ErrNotFound is returned if the job does not +// exist or is deleted before its finalized state is observed. Commit any +// transaction that inserts the job before calling JobWaitFinalized. +// +// This method observes current state, not completion history. If a finalized job +// is manually retried before being observed, the wait continues for that execution. +func (c *Client[TTx]) JobWaitFinalized(ctx context.Context, id int64, opts *JobWaitFinalizedOpts) (*rivertype.JobRow, error) { + if !c.driver.PoolIsSet() { + return nil, errNoDriverDBPool + } + if opts == nil { + opts = &JobWaitFinalizedOpts{} + } + if opts.PollInterval < 0 { + return nil, errors.New("PollInterval cannot be less than zero") + } + return c.jobWaiter.wait(ctx, id, cmp.Or(opts.PollInterval, jobWaitPollIntervalDefault)) +} + // ID returns the unique ID of this client as set in its config or // auto-generated if not specified. func (c *Client[TTx]) ID() string { diff --git a/client_test.go b/client_test.go index 4a9ae161..ee1bdd62 100644 --- a/client_test.go +++ b/client_test.go @@ -14,6 +14,7 @@ import ( "sync" "sync/atomic" "testing" + "testing/synctest" "time" "github.com/jackc/pgerrcode" @@ -5955,6 +5956,231 @@ func Test_Client_JobUpdateTx(t *testing.T) { }) } +func Test_Client_JobWaitFinalized(t *testing.T) { + t.Parallel() + + ctx := context.Background() + + type testBundle struct { + client *Client[pgx.Tx] + driver *riverpgxv5.Driver + exec riverdriver.Executor + schema string + } + + setup := func(t *testing.T) *testBundle { + t.Helper() + + var ( + dbPool = riversharedtest.DBPool(ctx, t) + driver = riverpgxv5.New(dbPool) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + ) + client, err := NewClient(driver, &Config{Logger: riversharedtest.Logger(t), Schema: schema}) + require.NoError(t, err) + + return &testBundle{ + client: client, + driver: driver, + exec: driver.GetExecutor(), + schema: schema, + } + } + + t.Run("AlreadyFinalized", func(t *testing.T) { + t.Parallel() + + for _, state := range []rivertype.JobState{rivertype.JobStateCancelled, rivertype.JobStateCompleted, rivertype.JobStateDiscarded} { + t.Run(string(state), func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{ + Metadata: []byte(`{"output":{"answer":42}}`), + Schema: bundle.schema, + State: &state, + }) + + waitCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + finalized, err := bundle.client.JobWaitFinalized(waitCtx, job.ID, nil) + require.NoError(t, err) + require.Equal(t, job.ID, finalized.ID) + require.Equal(t, state, finalized.State) + require.NotNil(t, finalized.FinalizedAt) + require.JSONEq(t, `{"answer":42}`, string(finalized.Output())) + }) + } + }) + + t.Run("CancellationDoesNotCancelJob", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema}) + + waitCtx, cancel := context.WithTimeout(ctx, 50*time.Millisecond) + defer cancel() + finalized, err := bundle.client.JobWaitFinalized(waitCtx, job.ID, nil) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Nil(t, finalized) + + job, err = bundle.client.JobGet(ctx, job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateAvailable, job.State) + }) + + t.Run("CompletedByAnotherClient", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + config := newTestConfig(t, bundle.schema) + config.PollOnly = true + workerClient, err := NewClient(bundle.driver, config) + require.NoError(t, err) + startClient(ctx, t, workerClient) + + insertRes, err := bundle.client.Insert(ctx, noOpArgs{}, nil) + require.NoError(t, err) + waitCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + finalized, err := bundle.client.JobWaitFinalized(waitCtx, insertRes.Job.ID, nil) + require.NoError(t, err) + require.Equal(t, insertRes.Job.ID, finalized.ID) + require.Equal(t, rivertype.JobStateCompleted, finalized.State) + require.Equal(t, 1, finalized.Attempt) + }) + + t.Run("FinalizationMustCommit", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, State: new(rivertype.JobStateRunning)}) + + finalizeInTxFunc := func() riverdriver.ExecutorTx { + execTx, err := bundle.exec.Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = execTx.Rollback(ctx) }) + _, err = execTx.JobSetStateIfRunningMany(ctx, &riverdriver.JobSetStateIfRunningManyParams{ + ID: []int64{job.ID}, + Attempt: []*int{nil}, + ErrData: [][]byte{nil}, + FinalizedAt: []*time.Time{new(time.Now())}, + MetadataDoMerge: []bool{false}, + MetadataUpdates: [][]byte{nil}, + ScheduledAt: []*time.Time{nil}, + Schema: bundle.schema, + State: []rivertype.JobState{rivertype.JobStateCompleted}, + }) + require.NoError(t, err) + return execTx + } + + execTx := finalizeInTxFunc() + waitCtx, cancel := context.WithTimeout(ctx, 50*time.Millisecond) + defer cancel() + finalized, err := bundle.client.JobWaitFinalized(waitCtx, job.ID, nil) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Nil(t, finalized) + require.NoError(t, execTx.Rollback(ctx)) + + persisted, err := bundle.client.JobGet(ctx, job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateRunning, persisted.State) + + execTx = finalizeInTxFunc() + require.NoError(t, execTx.Commit(ctx)) + waitCtx, cancel = context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + finalized, err = bundle.client.JobWaitFinalized(waitCtx, job.ID, nil) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCompleted, finalized.State) + }) + + t.Run("NegativePollInterval", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + job, err := bundle.client.JobWaitFinalized(ctx, 0, &JobWaitFinalizedOpts{PollInterval: -time.Second}) + require.EqualError(t, err, "PollInterval cannot be less than zero") + require.Nil(t, job) + }) + + t.Run("NotFound", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + + waitCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + job, err := bundle.client.JobWaitFinalized(waitCtx, 0, nil) + require.ErrorIs(t, err, rivertype.ErrNotFound) + require.Nil(t, job) + }) + + t.Run("PollInterval", func(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + name string + interval time.Duration + opts *JobWaitFinalizedOpts + }{ + {name: "CustomLonger", interval: time.Second, opts: &JobWaitFinalizedOpts{PollInterval: time.Second}}, + {name: "CustomShorter", interval: 10 * time.Millisecond, opts: &JobWaitFinalizedOpts{PollInterval: 10 * time.Millisecond}}, + {name: "Nil", interval: 250 * time.Millisecond}, + {name: "Zero", interval: 250 * time.Millisecond, opts: &JobWaitFinalizedOpts{}}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + // Only PoolIsSet is used; the executor stub keeps all work in + // the bubble so polling intervals use deterministic fake time. + client, err := NewClient(riverpgxv5.New(&pgxpool.Pool{}), &Config{}) + require.NoError(t, err) + exec := &jobWaiterExecutorStub{} + exec.testSignals.Init(t) + client.jobWaiter.getExecutorFunc = func() riverdriver.Executor { return exec } + + var numCalls atomic.Int32 + exec.getByIDManyFunc = func(context.Context, *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + state := rivertype.JobStateAvailable + if numCalls.Add(1) == 2 { + state = rivertype.JobStateCompleted + } + return []*rivertype.JobRow{{ID: 1, State: state}}, nil + } + + waitCtx, cancel := context.WithCancel(t.Context()) + defer cancel() + go func() { + job, err := client.JobWaitFinalized(waitCtx, 1, testCase.opts) + exec.testSignals.WaitFinished.Signal(jobWaiterResult{err: err, job: job}) + }() + synctest.Wait() + require.EqualValues(t, 1, numCalls.Load()) + + time.Sleep(testCase.interval - time.Nanosecond) + synctest.Wait() + require.EqualValues(t, 1, numCalls.Load()) + time.Sleep(time.Nanosecond) + synctest.Wait() + require.EqualValues(t, 2, numCalls.Load()) + result := exec.testSignals.WaitFinished.WaitOrTimeout() + require.NoError(t, result.err) + require.Equal(t, rivertype.JobStateCompleted, result.job.State) + }) + }) + } + }) +} + func Test_Client_ErrorHandler(t *testing.T) { t.Parallel() diff --git a/job_waiter.go b/job_waiter.go new file mode 100644 index 00000000..6c6a7c61 --- /dev/null +++ b/job_waiter.go @@ -0,0 +1,286 @@ +package river + +import ( + "context" + "fmt" + "slices" + "sync" + "time" + + "github.com/riverqueue/river/internal/rivercommon" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/util/timeoututil" + "github.com/riverqueue/river/rivertype" +) + +const ( + jobWaitBatchSize = 1_000 + jobWaitPollIntervalDefault = 250 * time.Millisecond +) + +// JobWaitFinalizedOpts are optional settings for waiting for a job to finalize. +type JobWaitFinalizedOpts struct { + // PollInterval is the interval between database polls after the initial, + // immediate check. Defaults to 250 milliseconds when zero. Must not be negative. + // + // Concurrent waits on the same client share one poll loop using the shortest + // interval requested by an active wait, so a job may be checked more often + // than this interval. The loop adjusts as waits finish or are cancelled. + PollInterval time.Duration +} + +// jobWaiter polls only while callers are waiting. Its lifetime is independent of +// worker services so it also works on clients that only insert and inspect jobs. +type jobWaiter struct { + getExecutorFunc func() riverdriver.Executor + schema string + + mu sync.Mutex + activeRun *jobWaiterRun +} + +func newJobWaiter(getExecutorFunc func() riverdriver.Executor, schema string) *jobWaiter { + return &jobWaiter{ + getExecutorFunc: getExecutorFunc, + schema: schema, + } +} + +func (w *jobWaiter) pollInterval(run *jobWaiterRun) time.Duration { + w.mu.Lock() + defer w.mu.Unlock() + + var interval time.Duration + for _, requests := range run.requests { + for request := range requests { + if interval == 0 || request.pollInterval < interval { + interval = request.pollInterval + } + } + } + return interval +} + +func (w *jobWaiter) pollOnce(ctx context.Context, run *jobWaiterRun, pendingOnly bool) { + batch := func() []jobWaiterBatchEntry { + w.mu.Lock() + defer w.mu.Unlock() + + batch := make([]jobWaiterBatchEntry, 0, len(run.requests)) + for id, requests := range run.requests { + if _, pending := run.pending[id]; pendingOnly && !pending { + continue + } + delete(run.pending, id) + entry := jobWaiterBatchEntry{id: id, requests: make([]*jobWaiterRequest, 0, len(requests))} + for request := range requests { + entry.requests = append(entry.requests, request) + } + batch = append(batch, entry) + } + return batch + }() + + for chunk := range slices.Chunk(batch, jobWaitBatchSize) { + if ctx.Err() != nil { + return + } + ids := make([]int64, len(chunk)) + for i, entry := range chunk { + ids[i] = entry.id + } + jobs, err := timeoututil.WithTimeoutV(ctx, rivercommon.HotOperationTimeout, "JobWaitFinalized", func(ctx context.Context) ([]*rivertype.JobRow, error) { + return w.getExecutorFunc().JobGetByIDMany(ctx, &riverdriver.JobGetByIDManyParams{ID: ids, Schema: w.schema}) + }) + if err != nil { + err = fmt.Errorf("error waiting for jobs to finalize: %w", err) + } + jobsByID := make(map[int64]*rivertype.JobRow, len(jobs)) + for _, job := range jobs { + jobsByID[job.ID] = job + } + + func() { + w.mu.Lock() + defer w.mu.Unlock() + + for _, entry := range chunk { + result := jobWaiterResult{err: err, job: jobsByID[entry.id]} + if err == nil { + if result.job == nil { + result.err = rivertype.ErrNotFound + } else if !slices.Contains([]rivertype.JobState{ + rivertype.JobStateCancelled, rivertype.JobStateCompleted, rivertype.JobStateDiscarded, + }, result.job.State) { + continue + } + } + for _, request := range entry.requests { + // Only notify callers in the snapshot that are still registered. + // A new caller must get a read initiated after it registered. + if _, registered := run.requests[entry.id][request]; registered { + request.resultChan <- result + w.removeLocked(run, request) + } + } + } + }() + } +} + +func (w *jobWaiter) register(id int64, pollInterval time.Duration) (*jobWaiterRun, *jobWaiterRequest) { + w.mu.Lock() + defer w.mu.Unlock() + + if w.activeRun == nil { + // No individual caller owns this context: cancelling one wait must not + // interrupt a shared query that other callers still need. + ctx, cancel := context.WithCancel(context.Background()) + w.activeRun = &jobWaiterRun{ + cancelFunc: cancel, + pending: make(map[int64]struct{}), + requests: make(map[int64]map[*jobWaiterRequest]struct{}), + wakeChan: make(chan struct{}, 1), + } + go w.run(ctx, w.activeRun) + } + + run := w.activeRun + request := &jobWaiterRequest{id: id, pollInterval: pollInterval, resultChan: make(chan jobWaiterResult, 1)} + if run.requests[id] == nil { + run.requests[id] = make(map[*jobWaiterRequest]struct{}) + } + run.requests[id][request] = struct{}{} + run.pending[id] = struct{}{} + select { + case run.wakeChan <- struct{}{}: + default: + } + return run, request +} + +func (w *jobWaiter) remove(run *jobWaiterRun, request *jobWaiterRequest) { + w.mu.Lock() + defer w.mu.Unlock() + + w.removeLocked(run, request) +} + +func (w *jobWaiter) removeLocked(run *jobWaiterRun, request *jobWaiterRequest) { + if _, registered := run.requests[request.id][request]; !registered { + return + } + delete(run.requests[request.id], request) + if len(run.requests[request.id]) == 0 { + delete(run.requests, request.id) + delete(run.pending, request.id) + } + if len(run.requests) == 0 { + run.cancelFunc() + if w.activeRun == run { + w.activeRun = nil + } + } else { + // Removing the fastest wait may allow the shared loop to poll less often. + select { + case run.wakeChan <- struct{}{}: + default: + } + } +} + +func (w *jobWaiter) run(ctx context.Context, run *jobWaiterRun) { + pollInterval := w.pollInterval(run) + if pollInterval == 0 { + return + } + ticker := time.NewTicker(pollInterval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-run.wakeChan: + // New waits are checked promptly without polling every existing ID + // again each time a caller registers. + w.pollOnce(ctx, run, true) + case <-ticker.C: + w.pollOnce(ctx, run, false) + } + + interval := w.pollInterval(run) + if interval == 0 { + return + } + // Preserve the existing schedule when new waits use the same interval. + if interval != pollInterval { + ticker.Reset(interval) + pollInterval = interval + } + } +} + +func (w *jobWaiter) wait(ctx context.Context, id int64, pollInterval time.Duration) (*rivertype.JobRow, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + run, request := w.register(id, pollInterval) //nolint:contextcheck // Shared polling has its own lifetime, independent of any caller's context. + defer w.remove(run, request) + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-request.resultChan: + if err := ctx.Err(); err != nil { + return nil, err + } + if result.err != nil { + return nil, result.err + } + return result.jobCopy(), nil + } +} + +type jobWaiterBatchEntry struct { + id int64 + requests []*jobWaiterRequest +} + +type jobWaiterRequest struct { + id int64 + pollInterval time.Duration + resultChan chan jobWaiterResult +} + +type jobWaiterResult struct { + err error + job *rivertype.JobRow +} + +// Each caller owns its result, including slices and optional timestamps. The +// database row may otherwise be shared by many concurrent waits for the same ID. +func (r jobWaiterResult) jobCopy() *rivertype.JobRow { + job := *r.job + if job.AttemptedAt != nil { + job.AttemptedAt = new(*job.AttemptedAt) + } + job.AttemptedBy = slices.Clone(job.AttemptedBy) + job.EncodedArgs = slices.Clone(job.EncodedArgs) + job.Errors = slices.Clone(job.Errors) + if job.FinalizedAt != nil { + job.FinalizedAt = new(*job.FinalizedAt) + } + job.Metadata = slices.Clone(job.Metadata) + job.Tags = slices.Clone(job.Tags) + job.UniqueKey = slices.Clone(job.UniqueKey) + job.UniqueStates = slices.Clone(job.UniqueStates) + return &job +} + +type jobWaiterRun struct { + cancelFunc context.CancelFunc + pending map[int64]struct{} // guarded by jobWaiter.mu + requests map[int64]map[*jobWaiterRequest]struct{} // guarded by jobWaiter.mu + wakeChan chan struct{} +} diff --git a/job_waiter_test.go b/job_waiter_test.go new file mode 100644 index 00000000..e74ea929 --- /dev/null +++ b/job_waiter_test.go @@ -0,0 +1,421 @@ +package river + +import ( + "context" + "errors" + "slices" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/testsignal" + "github.com/riverqueue/river/rivertype" +) + +func TestJobWaiter(t *testing.T) { + t.Parallel() + + type testBundle struct { + exec *jobWaiterExecutorStub + waiter *jobWaiter + } + + setup := func(t *testing.T) *testBundle { + t.Helper() + + exec := &jobWaiterExecutorStub{} + waiter := newJobWaiter(func() riverdriver.Executor { return exec }, "custom_schema") + t.Cleanup(func() { + require.Eventually(t, func() bool { + waiter.mu.Lock() + defer waiter.mu.Unlock() + return waiter.activeRun == nil + }, riversharedtest.WaitTimeout(), time.Millisecond) + }) + return &testBundle{exec: exec, waiter: waiter} + } + + t.Run("BatchesAndDeduplicates", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + bundle.exec.testSignals.Init(t) + var numCalls atomic.Int32 + bundle.exec.getByIDManyFunc = func(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + require.Equal(t, "custom_schema", params.Schema) + require.LessOrEqual(t, len(params.ID), jobWaitBatchSize) + if numCalls.Add(1) == 1 { + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + select { + case <-bundle.exec.testSignals.LookupContinue.WaitC(): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + jobs := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + jobs[i] = &rivertype.JobRow{ID: id, State: rivertype.JobStateCompleted} + } + return jobs, nil + } + + // Hold the first query so all remaining registrations share one poll. + run, first := bundle.waiter.register(0, time.Hour) + t.Cleanup(func() { bundle.waiter.remove(run, first) }) + bundle.exec.testSignals.LookupStarted.WaitOrTimeout() + requests := make([]*jobWaiterRequest, 0, 1+2*(jobWaitBatchSize*2+1)) + requests = append(requests, first) + for id := range int64(jobWaitBatchSize*2 + 1) { + for range 2 { + _, request := bundle.waiter.register(id+1, time.Hour) + t.Cleanup(func() { bundle.waiter.remove(run, request) }) + requests = append(requests, request) + } + } + bundle.exec.testSignals.LookupContinue.Signal(struct{}{}) + for _, request := range requests { + result := riversharedtest.WaitOrTimeout(t, request.resultChan) + require.NoError(t, result.err) + require.Equal(t, request.id, result.job.ID) + } + require.EqualValues(t, 4, numCalls.Load()) + }) + + t.Run("CancelledContextDoesNotQuery", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + job, err := bundle.waiter.wait(ctx, 1, 10*time.Millisecond) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, job) + }) + + t.Run("CancelledWaitDoesNotCancelOtherWaits", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + bundle.exec.testSignals.Init(t) + var finalized atomic.Bool + bundle.exec.getByIDManyFunc = func(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + state := rivertype.JobStateRunning + if finalized.Load() { + state = rivertype.JobStateCompleted + } + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + return []*rivertype.JobRow{{ID: 1, State: state}}, nil + } + + run, other := bundle.waiter.register(1, 10*time.Millisecond) + t.Cleanup(func() { bundle.waiter.remove(run, other) }) + bundle.exec.testSignals.LookupStarted.WaitOrTimeout() + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + resultChan := make(chan jobWaiterResult, 1) + go func() { + job, err := bundle.waiter.wait(ctx, 1, 10*time.Millisecond) + resultChan <- jobWaiterResult{err: err, job: job} + }() + require.Eventually(t, func() bool { + bundle.waiter.mu.Lock() + defer bundle.waiter.mu.Unlock() + return len(run.requests[1]) == 2 + }, riversharedtest.WaitTimeout(), time.Millisecond) + cancel() + require.ErrorIs(t, riversharedtest.WaitOrTimeout(t, resultChan).err, context.Canceled) + + finalized.Store(true) + result := riversharedtest.WaitOrTimeout(t, other.resultChan) + require.NoError(t, result.err) + require.Equal(t, rivertype.JobStateCompleted, result.job.State) + }) + + t.Run("DatabaseError", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + expectedErr := errors.New("database unavailable") + bundle.exec.getByIDManyFunc = func(context.Context, *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + return nil, expectedErr + } + + job, err := bundle.waiter.wait(t.Context(), 1, 10*time.Millisecond) + require.ErrorIs(t, err, expectedErr) + require.Nil(t, job) + }) + + t.Run("DeadlineCancelsLastQueryAndCanRestart", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + bundle.exec.testSignals.Init(t) + var numCalls atomic.Int32 + bundle.exec.getByIDManyFunc = func(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + if numCalls.Add(1) == 1 { + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + <-ctx.Done() + bundle.exec.testSignals.QueryCancelled.Signal(struct{}{}) + return nil, ctx.Err() + } + return []*rivertype.JobRow{{ID: 1, State: rivertype.JobStateCompleted}}, nil + } + + ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) + defer cancel() + job, err := bundle.waiter.wait(ctx, 1, 10*time.Millisecond) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Nil(t, job) + bundle.exec.testSignals.QueryCancelled.WaitOrTimeout() + + job, err = bundle.waiter.wait(t.Context(), 1, 10*time.Millisecond) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCompleted, job.State) + }) + + t.Run("NewCallerDoesNotReceiveAnEarlierSnapshot", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + bundle.exec.testSignals.Init(t) + var numCalls atomic.Int32 + bundle.exec.getByIDManyFunc = func(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + attempt := numCalls.Add(1) + if attempt == 1 { + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + select { + case <-bundle.exec.testSignals.LookupContinue.WaitC(): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + return []*rivertype.JobRow{{ID: 1, Attempt: int(attempt), State: rivertype.JobStateCompleted}}, nil + } + + run, first := bundle.waiter.register(1, time.Hour) + t.Cleanup(func() { bundle.waiter.remove(run, first) }) + bundle.exec.testSignals.LookupStarted.WaitOrTimeout() + _, second := bundle.waiter.register(1, time.Hour) + t.Cleanup(func() { bundle.waiter.remove(run, second) }) + bundle.exec.testSignals.LookupContinue.Signal(struct{}{}) + + require.Equal(t, 1, riversharedtest.WaitOrTimeout(t, first.resultChan).job.Attempt) + require.Equal(t, 2, riversharedtest.WaitOrTimeout(t, second.resultChan).job.Attempt) + }) + + t.Run("NewWaitDoesNotDelayExistingPoll", func(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + bundle := setup(t) + bundle.exec.testSignals.Init(t) + + bundle.exec.getByIDManyFunc = func(_ context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + jobs := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + jobs[i] = &rivertype.JobRow{ID: id, State: rivertype.JobStateRunning} + } + return jobs, nil + } + + for i, interval := range []time.Duration{250 * time.Millisecond, 250 * time.Millisecond, time.Hour} { + if i > 0 { + time.Sleep(100 * time.Millisecond) + } + run, request := bundle.waiter.register(int64(i+1), interval) + t.Cleanup(func() { bundle.waiter.remove(run, request) }) + synctest.Wait() + require.Equal(t, []int64{int64(i + 1)}, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + } + + time.Sleep(49 * time.Millisecond) + synctest.Wait() + bundle.exec.testSignals.LookupStarted.RequireEmpty() + time.Sleep(time.Millisecond) + synctest.Wait() + require.ElementsMatch(t, []int64{1, 2, 3}, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + }) + }) + + t.Run("NonFinalizedStatesKeepWaiting", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + states := []rivertype.JobState{ + rivertype.JobStateAvailable, rivertype.JobStatePending, rivertype.JobStateRunning, + rivertype.JobStateRetryable, rivertype.JobStateScheduled, rivertype.JobStateCompleted, + } + var numCalls atomic.Int32 + bundle.exec.getByIDManyFunc = func(context.Context, *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + state := states[numCalls.Add(1)-1] + return []*rivertype.JobRow{{ID: 1, State: state}}, nil + } + + job, err := bundle.waiter.wait(t.Context(), 1, 10*time.Millisecond) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCompleted, job.State) + require.EqualValues(t, len(states), numCalls.Load()) + }) + + t.Run("NotFoundAfterDeletion", func(t *testing.T) { + t.Parallel() + + bundle := setup(t) + var numCalls atomic.Int32 + bundle.exec.getByIDManyFunc = func(context.Context, *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + if numCalls.Add(1) == 1 { + return []*rivertype.JobRow{{ID: 1, State: rivertype.JobStateAvailable}}, nil + } + return nil, nil + } + + job, err := bundle.waiter.wait(t.Context(), 1, 10*time.Millisecond) + require.ErrorIs(t, err, rivertype.ErrNotFound) + require.Nil(t, job) + }) + + t.Run("PollIntervalChangesWithActiveWaits", func(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + name string + complete bool + fastID int64 + }{ + {name: "CancelledDifferentJob", fastID: 2}, + {name: "CancelledSameJob", fastID: 1}, + {name: "Completed", complete: true, fastID: 2}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + bundle := setup(t) + bundle.exec.testSignals.Init(t) + + var finalized atomic.Bool + bundle.exec.getByIDManyFunc = func(_ context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + bundle.exec.testSignals.LookupStarted.Signal(params.ID) + jobs := make([]*rivertype.JobRow, len(params.ID)) + for i, id := range params.ID { + state := rivertype.JobStateRunning + if id == testCase.fastID && finalized.Load() { + state = rivertype.JobStateCompleted + } + jobs[i] = &rivertype.JobRow{ID: id, State: state} + } + return jobs, nil + } + + run, slow := bundle.waiter.register(1, time.Second) + t.Cleanup(func() { bundle.waiter.remove(run, slow) }) + synctest.Wait() + require.Equal(t, []int64{1}, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + + time.Sleep(100 * time.Millisecond) + fastRun, fast := bundle.waiter.register(testCase.fastID, 100*time.Millisecond) + t.Cleanup(func() { bundle.waiter.remove(fastRun, fast) }) + synctest.Wait() + require.Same(t, run, fastRun) + require.Equal(t, []int64{testCase.fastID}, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + + // The faster wait shortens the shared interval for both jobs. + finalized.Store(testCase.complete) + time.Sleep(99 * time.Millisecond) + synctest.Wait() + bundle.exec.testSignals.LookupStarted.RequireEmpty() + time.Sleep(time.Millisecond) + synctest.Wait() + expectedIDs := []int64{1} + if testCase.fastID != 1 { + expectedIDs = append(expectedIDs, testCase.fastID) + } + require.ElementsMatch(t, expectedIDs, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + + if testCase.complete { + result := riversharedtest.WaitOrTimeout(t, fast.resultChan) + require.NoError(t, result.err) + require.Equal(t, rivertype.JobStateCompleted, result.job.State) + } else { + bundle.waiter.remove(run, fast) + } + synctest.Wait() + + // Once the faster wait leaves, the slower interval is restored. + time.Sleep(999 * time.Millisecond) + synctest.Wait() + bundle.exec.testSignals.LookupStarted.RequireEmpty() + time.Sleep(time.Millisecond) + synctest.Wait() + require.Equal(t, []int64{1}, bundle.exec.testSignals.LookupStarted.WaitOrTimeout()) + }) + }) + } + }) +} + +func TestJobWaiterResultJobCopy(t *testing.T) { + t.Parallel() + + original := &rivertype.JobRow{ + ID: 1, + AttemptedAt: new(time.Now()), + AttemptedBy: []string{"client"}, + EncodedArgs: []byte(`{}`), + Errors: []rivertype.AttemptError{{Error: "error"}}, + FinalizedAt: new(time.Now()), + Metadata: []byte(`{}`), + Tags: []string{"tag"}, + UniqueKey: []byte("key"), + UniqueStates: []rivertype.JobState{rivertype.JobStateCompleted}, + } + result := jobWaiterResult{job: original} + first, second := result.jobCopy(), result.jobCopy() + require.Equal(t, original, first) + + first.ID++ + *first.AttemptedAt = time.Time{} + first.AttemptedBy[0] = "changed" + first.EncodedArgs[0] = 'x' + first.Errors[0].Error = "changed" + *first.FinalizedAt = time.Time{} + first.Metadata[0] = 'x' + first.Tags[0] = "changed" + first.UniqueKey[0] = 'x' + first.UniqueStates[0] = rivertype.JobStateAvailable + require.Equal(t, original, second) + require.False(t, slices.Equal(first.Metadata, second.Metadata)) +} + +type jobWaiterExecutorStub struct { + riverdriver.Executor + + getByIDManyFunc func(context.Context, *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) + testSignals jobWaiterExecutorTestSignals +} + +func (e *jobWaiterExecutorStub) JobGetByIDMany(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + return e.getByIDManyFunc(ctx, params) +} + +type jobWaiterExecutorTestSignals struct { + LookupContinue testsignal.TestSignal[struct{}] + LookupStarted testsignal.TestSignal[[]int64] + QueryCancelled testsignal.TestSignal[struct{}] + WaitFinished testsignal.TestSignal[jobWaiterResult] +} + +func (s *jobWaiterExecutorTestSignals) Init(t *testing.T) { + t.Helper() + s.LookupContinue.Init(t) + s.LookupStarted.Init(t) + s.QueryCancelled.Init(t) + s.WaitFinished.Init(t) +}