Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 38 additions & 0 deletions client_context.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,50 @@ package river
import (
"context"
"errors"
"time"

"github.com/riverqueue/river/internal/rivercommon"
"github.com/riverqueue/river/rivershared/riverpilot"
)

var errClientNotInContext = errors.New("river: client not found in context, can only be used in a Worker")

type clientContextData struct {
Pilot riverpilot.Pilot
Schema string
Time time.Time
}

type clientContextProvider interface {
clientContextData() *clientContextData
}

func clientContextDataFromContext(ctx context.Context) *clientContextData {
client, exists := ctx.Value(rivercommon.ContextKeyClient{}).(clientContextProvider)
if !exists || client == nil {
panic(errClientNotInContext)
}

data := client.clientContextData()
if data == nil {
panic(errClientNotInContext)
}

return data
}

func (c *Client[TTx]) clientContextData() *clientContextData {
if c == nil {
return nil
}

return &clientContextData{
Pilot: c.Pilot(),
Schema: c.Schema(),
Time: c.baseService.Time.Now(),
}
}

func withClient[TTx any](ctx context.Context, client *Client[TTx]) context.Context {
return context.WithValue(ctx, rivercommon.ContextKeyClient{}, client)
}
Expand Down
30 changes: 14 additions & 16 deletions job_complete_tx.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,11 @@ import (
// JobCompleteTx marks the job as completed as part of transaction tx. If tx is
// rolled back, the completion will be as well.
//
// The function needs to know the type of the River database driver, which is
// the same as the one in use by Client, but the other generic parameters can be
// inferred. An invocation should generally look like:
// The function needs to know the type of the River database driver that is
// compatible with tx. This is usually the same driver used by Client, but may
// differ when a worker and its application transactions use different database
// abstractions. The other generic parameters can be inferred. An invocation
// should generally look like:
//
// _, err := river.JobCompleteTx[*riverpgxv5.Driver](ctx, tx, job)
// if err != nil {
Expand All @@ -35,39 +37,35 @@ func JobCompleteTx[TDriver riverdriver.Driver[TTx], TTx any, TArgs JobArgs](ctx
return nil, errors.New("job must be running")
}

client := ClientFromContext[TTx](ctx)
if client == nil {
return nil, errors.New("client not found in context, can only work within a River worker")
}
clientData := clientContextDataFromContext(ctx)

driver := client.Driver()
pilot := client.Pilot()
var driver TDriver

// extract metadata updates from context
metadataUpdates, hasMetadataUpdates := jobexecutor.MetadataUpdatesFromWorkContext(ctx)
hasMetadataUpdates = hasMetadataUpdates && len(metadataUpdates) > 0
var (
metadataUpdatesBytes []byte
err error
marshalErr error
)
if hasMetadataUpdates {
metadataUpdatesBytes, err = json.Marshal(metadataUpdates)
if err != nil {
return nil, err
metadataUpdatesBytes, marshalErr = json.Marshal(metadataUpdates)
if marshalErr != nil {
return nil, marshalErr
}
}

execTx := driver.UnwrapExecutor(tx)
params := riverdriver.JobSetStateCompleted(job.ID, client.baseService.Time.Now(), nil)
rows, err := pilot.JobSetStateIfRunningMany(ctx, execTx, &riverdriver.JobSetStateIfRunningManyParams{
params := riverdriver.JobSetStateCompleted(job.ID, clientData.Time, nil)
rows, err := clientData.Pilot.JobSetStateIfRunningMany(ctx, execTx, &riverdriver.JobSetStateIfRunningManyParams{
ID: []int64{params.ID},
Attempt: []*int{params.Attempt},
ErrData: [][]byte{params.ErrData},
FinalizedAt: []*time.Time{params.FinalizedAt},
MetadataDoMerge: []bool{hasMetadataUpdates},
MetadataUpdates: [][]byte{metadataUpdatesBytes},
ScheduledAt: []*time.Time{params.ScheduledAt},
Schema: client.config.Schema,
Schema: clientData.Schema,
State: []rivertype.JobState{params.State},
})
if err != nil {
Expand Down
32 changes: 32 additions & 0 deletions job_complete_tx_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,24 @@ import (
"github.com/riverqueue/river/rivertype"
)

type wrappedPgxTx struct {
pgx.Tx
}

type wrappedPgxTxDriver struct {
*riverpgxv5.Driver
}

func (d *wrappedPgxTxDriver) UnwrapExecutor(tx *wrappedPgxTx) riverdriver.ExecutorTx {
var driver *riverpgxv5.Driver
return driver.UnwrapExecutor(tx.Tx)
}

func (d *wrappedPgxTxDriver) UnwrapTx(execTx riverdriver.ExecutorTx) *wrappedPgxTx {
var driver *riverpgxv5.Driver
return &wrappedPgxTx{Tx: driver.UnwrapTx(execTx)}
}

func TestJobCompleteTx(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -162,4 +180,18 @@ func TestJobCompleteTx(t *testing.T) {
require.NoError(t, err)
})
})

t.Run("UsesTransactionDriverInsteadOfWorkerDriver", func(t *testing.T) {
t.Parallel()

ctx, bundle := setup(ctx, t)

job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{
State: new(rivertype.JobStateRunning),
})

completedJob, err := JobCompleteTx[*wrappedPgxTxDriver](ctx, &wrappedPgxTx{Tx: bundle.tx}, &Job[JobArgs]{JobRow: job})
require.NoError(t, err)
require.Equal(t, rivertype.JobStateCompleted, completedJob.State)
})
}
17 changes: 8 additions & 9 deletions resumable_step_tx.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@ import (
// Must be called from within a ResumableStep or ResumableStepCursor callback.
// The current step name to persist is read from context.
func ResumableSetStepTx[TDriver riverdriver.Driver[TTx], TTx any, TArgs JobArgs](ctx context.Context, tx TTx, job *Job[TArgs]) (*Job[TArgs], error) {
return resumableSetStepTx(ctx, tx, job, nil)
return resumableSetStepTx[TDriver](ctx, tx, job, nil)
}

// ResumableSetStepCursorTx immediately persists the current resumable step and
Expand All @@ -48,10 +48,10 @@ func ResumableSetStepCursorTx[TDriver riverdriver.Driver[TTx], TTx any, TArgs Jo
return nil, err
}

return resumableSetStepTx(ctx, tx, job, json.RawMessage(cursorBytes))
return resumableSetStepTx[TDriver](ctx, tx, job, json.RawMessage(cursorBytes))
}

func resumableSetStepTx[TTx any, TArgs JobArgs](ctx context.Context, tx TTx, job *Job[TArgs], cursor json.RawMessage) (*Job[TArgs], error) {
func resumableSetStepTx[TDriver riverdriver.Driver[TTx], TTx any, TArgs JobArgs](ctx context.Context, tx TTx, job *Job[TArgs], cursor json.RawMessage) (*Job[TArgs], error) {
if job.State != rivertype.JobStateRunning {
return nil, errors.New("job must be running")
}
Expand All @@ -66,10 +66,9 @@ func resumableSetStepTx[TTx any, TArgs JobArgs](ctx context.Context, tx TTx, job

step := state.StepName

client := ClientFromContext[TTx](ctx)
if client == nil {
return nil, errors.New("client not found in context, can only work within a River worker")
}
clientData := clientContextDataFromContext(ctx)

var driver TDriver

metadataUpdates := map[string]any{
rivercommon.MetadataKeyResumableStep: step,
Expand Down Expand Up @@ -99,11 +98,11 @@ func resumableSetStepTx[TTx any, TArgs JobArgs](ctx context.Context, tx TTx, job
return nil, err
}

updatedJob, err := client.Driver().UnwrapExecutor(tx).JobUpdate(ctx, &riverdriver.JobUpdateParams{
updatedJob, err := driver.UnwrapExecutor(tx).JobUpdate(ctx, &riverdriver.JobUpdateParams{
ID: job.ID,
MetadataDoMerge: true,
Metadata: metadataUpdatesBytes,
Schema: client.config.Schema,
Schema: clientData.Schema,
})
if err != nil {
if errors.Is(err, rivertype.ErrNotFound) {
Expand Down
14 changes: 14 additions & 0 deletions resumable_step_tx_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -165,4 +165,18 @@ func TestResumableSetStepTx(t *testing.T) {
require.NoError(t, err)
})
})

t.Run("UsesTransactionDriverInsteadOfWorkerDriver", func(t *testing.T) {
t.Parallel()

ctx, bundle := setup(ctx, t, "step1")

job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{
State: new(rivertype.JobStateRunning),
})

updatedJob, err := ResumableSetStepTx[*wrappedPgxTxDriver](ctx, &wrappedPgxTx{Tx: bundle.tx}, &Job[JobArgs]{JobRow: job})
require.NoError(t, err)
require.Equal(t, rivertype.JobStateRunning, updatedJob.State)
})
}
Loading