From 5fad4a4c0bccdada04a3a2c5083751e4381b0c20 Mon Sep 17 00:00:00 2001 From: Brandur Date: Fri, 21 Feb 2025 00:04:20 -0800 Subject: [PATCH] Job executor: Unmarshal job args late, after middlewares have run Here, modify the job executor so that instead of exclusively invoking a work unit's `UnmarshalJSON` implementation before starting to execute the stack of work middleware, invoke it much later from within the executor's `doInner` implementation. This gives middlewares the opportunity to modify args before unmarshaling, thereby enabling use cases like args encryption. This does have the downside of having to modify the worker interface's `Middleware` function from taking a generic job: type Worker[T JobArgs] interface { Middleware(job *Job[T]) []rivertype.WorkerMiddleware To taking a `JobRow` instead: type Worker[T JobArgs] interface { Middleware(job *rivertype.JobRow) []rivertype.WorkerMiddleware Because the full generic `Job[T]` is not yet available by the time we need to extract a middleware stack. --- CHANGELOG.md | 5 +++ client_test.go | 57 ++++++++++++++++++++++++---- common_test.go | 2 +- internal/execution/execution.go | 21 +++++----- internal/jobexecutor/job_executor.go | 21 ++++++---- rivertest/worker.go | 2 +- work_unit_wrapper.go | 2 +- worker.go | 4 +- 8 files changed, 85 insertions(+), 29 deletions(-) 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