diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c28b670..c042b6e4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,7 +9,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Changed +⚠️ Version 0.19.0 has minor breaking changes for the `Worker.Middleware`, introduced fairly recently in 0.17.0. We tried not to make this change, but found the existing middleware interface insufficient to provide the necessary range of functionality we wanted, and this is a secondary middleware facility that won't be in use for many users, so it seemed worthwhile. + +### Changed + - The `river.RecordOutput` function now returns an error if the output is too large. The output is limited to 32MB in size. [PR #782](https://github.com/riverqueue/river/pull/782). +- **Breaking change:** The `Worker` interface's `Middleware` function now takes a `JobRow` parameter instead of a generic `Job[T]`. This was necessary to expand the potential of what middleware can do: by letting the executor extract a middleware stack from a worker before a job is fully unmarshaled, the middleware can also participate in the unmarshaling process. [PR #783](https://github.com/riverqueue/river/pull/783). ## [0.18.0] - 2025-02-20 diff --git a/client_test.go b/client_test.go index f0d8f1f4..9d2af82b 100644 --- a/client_test.go +++ b/client_test.go @@ -661,8 +661,8 @@ func Test_Client(t *testing.T) { require.Equal(t, "called", ctx.Value(privateKey("middleware"))) return nil }, - middlewareFunc: func(job *Job[callbackArgs]) []rivertype.WorkerMiddleware { - require.Equal(t, "middleware_test", job.Args.Name, "JSON should be decoded before middleware is called") + middlewareFunc: func(job *rivertype.JobRow) []rivertype.WorkerMiddleware { + require.Equal(t, "callback", job.Kind) return []rivertype.WorkerMiddleware{ &overridableJobMiddleware{ @@ -694,6 +694,49 @@ func Test_Client(t *testing.T) { require.True(t, middlewareCalled) }) + t.Run("MiddlewareModifiesEncodedArgs", func(t *testing.T) { + t.Parallel() + + _, bundle := setup(t) + middlewareCalled := false + + worker := &workerWithMiddleware[callbackArgs]{ + workFunc: func(ctx context.Context, job *Job[callbackArgs]) error { + require.Equal(t, "middleware name", job.Args.Name) + return nil + }, + middlewareFunc: func(job *rivertype.JobRow) []rivertype.WorkerMiddleware { + return []rivertype.WorkerMiddleware{ + &overridableJobMiddleware{ + workFunc: func(ctx context.Context, job *rivertype.JobRow, doInner func(ctx context.Context) error) error { + middlewareCalled = true + require.Equal(t, `{"name": "inserted name"}`, string(job.EncodedArgs)) + job.EncodedArgs = []byte(`{"name": "middleware name"}`) + return doInner(ctx) + }, + }, + } + }, + } + + AddWorker(bundle.config.Workers, worker) + + driver := riverpgxv5.New(bundle.dbPool) + client, err := NewClient(driver, bundle.config) + require.NoError(t, err) + + subscribeChan := subscribe(t, client) + startClient(ctx, t, client) + + result, err := client.Insert(ctx, callbackArgs{Name: "inserted name"}, nil) + require.NoError(t, err) + + event := riversharedtest.WaitOrTimeout(t, subscribeChan) + require.Equal(t, EventKindJobCompleted, event.Kind) + require.Equal(t, result.Job.ID, event.Job.ID) + require.True(t, middlewareCalled) + }) + t.Run("PauseAndResumeSingleQueue", func(t *testing.T) { t.Parallel() @@ -943,15 +986,15 @@ func Test_Client(t *testing.T) { type workerWithMiddleware[T JobArgs] struct { WorkerDefaults[T] workFunc func(context.Context, *Job[T]) error - middlewareFunc func(*Job[T]) []rivertype.WorkerMiddleware + middlewareFunc func(*rivertype.JobRow) []rivertype.WorkerMiddleware } -func (w *workerWithMiddleware[T]) Work(ctx context.Context, job *Job[T]) error { - return w.workFunc(ctx, job) +func (w *workerWithMiddleware[T]) Middleware(job *rivertype.JobRow) []rivertype.WorkerMiddleware { + return w.middlewareFunc(job) } -func (w *workerWithMiddleware[T]) Middleware(job *Job[T]) []rivertype.WorkerMiddleware { - return w.middlewareFunc(job) +func (w *workerWithMiddleware[T]) Work(ctx context.Context, job *Job[T]) error { + return w.workFunc(ctx, job) } func Test_Client_Stop(t *testing.T) { diff --git a/common_test.go b/common_test.go index 720c64d1..30b374f6 100644 --- a/common_test.go +++ b/common_test.go @@ -32,7 +32,7 @@ func waitForNJobs(subscribeChan <-chan *river.Event, numJobs int) { } case <-time.After(time.Until(deadline)): - panic(fmt.Sprintf("WaitOrTimeout timed out after waiting %s (received %d job(s), wanted %d)", + panic(fmt.Sprintf("waitForNJobs timed out after waiting %s (received %d job(s), wanted %d)", timeout, len(events), numJobs)) } } diff --git a/internal/execution/execution.go b/internal/execution/execution.go index bcf9917e..9da462e4 100644 --- a/internal/execution/execution.go +++ b/internal/execution/execution.go @@ -29,19 +29,22 @@ func MaybeApplyTimeout(ctx context.Context, timeout time.Duration) (context.Cont // MiddlewareChain chains together the given middleware functions, returning a // single function that applies them all in reverse order. func MiddlewareChain(global, worker []rivertype.WorkerMiddleware, doInner Func, jobRow *rivertype.JobRow) Func { + // Quick return for no middleware, which will often be the case. + if len(global) < 1 && len(worker) < 1 { + return doInner + } + allMiddleware := make([]rivertype.WorkerMiddleware, 0, len(global)+len(worker)) allMiddleware = append(allMiddleware, global...) allMiddleware = append(allMiddleware, worker...) - if len(allMiddleware) > 0 { - // Wrap middlewares in reverse order so the one defined first is wrapped - // as the outermost function and is first to receive the operation. - for i := len(allMiddleware) - 1; i >= 0; i-- { - middlewareItem := allMiddleware[i] // capture the current middleware item - previousDoInner := doInner // capture the current doInner function - doInner = func(ctx context.Context) error { - return middlewareItem.Work(ctx, jobRow, previousDoInner) - } + // Wrap middlewares in reverse order so the one defined first is wrapped + // as the outermost function and is first to receive the operation. + for i := len(allMiddleware) - 1; i >= 0; i-- { + middlewareItem := allMiddleware[i] // capture the current middleware item + previousDoInner := doInner // capture the current doInner function + doInner = func(ctx context.Context) error { + return middlewareItem.Work(ctx, jobRow, previousDoInner) } } diff --git a/internal/jobexecutor/job_executor.go b/internal/jobexecutor/job_executor.go index 33da7f1e..980fd065 100644 --- a/internal/jobexecutor/job_executor.go +++ b/internal/jobexecutor/job_executor.go @@ -183,16 +183,21 @@ func (e *JobExecutor) execute(ctx context.Context) (res *jobExecutorResult) { return &jobExecutorResult{Err: &rivertype.UnknownJobKindError{Kind: e.JobRow.Kind}, MetadataUpdates: metadataUpdates} } - if err := e.WorkUnit.UnmarshalJob(); err != nil { - return &jobExecutorResult{Err: err, MetadataUpdates: metadataUpdates} - } + doInner := execution.Func(func(ctx context.Context) error { + if err := e.WorkUnit.UnmarshalJob(); err != nil { + return err + } + + jobTimeout := valutil.FirstNonZero(e.WorkUnit.Timeout(), e.ClientJobTimeout) + ctx, cancel := execution.MaybeApplyTimeout(ctx, jobTimeout) + defer cancel() + + return e.WorkUnit.Work(ctx) + }) - doInner := execution.MiddlewareChain(e.GlobalMiddleware, e.WorkUnit.Middleware(), e.WorkUnit.Work, e.JobRow) - jobTimeout := valutil.FirstNonZero(e.WorkUnit.Timeout(), e.ClientJobTimeout) - ctx, cancel := execution.MaybeApplyTimeout(ctx, jobTimeout) - defer cancel() + executeFunc := execution.MiddlewareChain(e.GlobalMiddleware, e.WorkUnit.Middleware(), doInner, e.JobRow) - return &jobExecutorResult{Err: doInner(ctx), MetadataUpdates: metadataUpdates} + return &jobExecutorResult{Err: executeFunc(ctx), MetadataUpdates: metadataUpdates} } func (e *JobExecutor) invokeErrorHandler(ctx context.Context, res *jobExecutorResult) bool { diff --git a/rivertest/worker.go b/rivertest/worker.go index 1198c946..1b18169f 100644 --- a/rivertest/worker.go +++ b/rivertest/worker.go @@ -300,7 +300,7 @@ func (w *wrapperWorkUnit[T]) Timeout() time.Duration { return w.worker.T func (w *wrapperWorkUnit[T]) Work(ctx context.Context) error { return w.worker.Work(ctx, w.job) } func (w *wrapperWorkUnit[T]) Middleware() []rivertype.WorkerMiddleware { - return w.worker.Middleware(w.job) + return w.worker.Middleware(w.jobRow) } func (w *wrapperWorkUnit[T]) UnmarshalJob() error { diff --git a/work_unit_wrapper.go b/work_unit_wrapper.go index 9b787ad1..2c8e3915 100644 --- a/work_unit_wrapper.go +++ b/work_unit_wrapper.go @@ -30,7 +30,7 @@ func (w *wrapperWorkUnit[T]) Timeout() time.Duration { return w.worker.T func (w *wrapperWorkUnit[T]) Work(ctx context.Context) error { return w.worker.Work(ctx, w.job) } func (w *wrapperWorkUnit[T]) Middleware() []rivertype.WorkerMiddleware { - return w.worker.Middleware(w.job) + return w.worker.Middleware(w.jobRow) } func (w *wrapperWorkUnit[T]) UnmarshalJob() error { diff --git a/worker.go b/worker.go index bf188ab9..54625d50 100644 --- a/worker.go +++ b/worker.go @@ -38,7 +38,7 @@ import ( // with the client using the AddWorker function. type Worker[T JobArgs] interface { // Middleware returns the type-specific middleware for this job. - Middleware(job *Job[T]) []rivertype.WorkerMiddleware + Middleware(job *rivertype.JobRow) []rivertype.WorkerMiddleware // NextRetry calculates when the next retry for a failed job should take // place given when it was last attempted and its number of attempts, or any @@ -74,7 +74,7 @@ type Worker[T JobArgs] interface { // struct to make it fulfill the Worker interface with default values. type WorkerDefaults[T JobArgs] struct{} -func (w WorkerDefaults[T]) Middleware(*Job[T]) []rivertype.WorkerMiddleware { return nil } +func (w WorkerDefaults[T]) Middleware(*rivertype.JobRow) []rivertype.WorkerMiddleware { return nil } // NextRetry returns an empty time.Time{} to avoid setting any job or // Worker-specific overrides on the next retry time. This means that the