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) +}