diff --git a/CHANGELOG.md b/CHANGELOG.md index b20acbc1f..405fbe9a5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Fixed +- Fixed cancelled transaction starts leaving Turso connections unusable, which could prevent maintenance from recovering after a startup failure. [PR #1347](https://github.com/riverqueue/river/pull/1347). +- Fixed maintenance startup failures leaving a client renewing leadership with maintenance stopped in poll-only mode. After exhausting startup retries, clients now request local resignation without depending on database notifications. [PR #1347](https://github.com/riverqueue/river/pull/1347). +- Fixed YugabyteDB clients relying on notifications when `LISTEN/NOTIFY` is unavailable or disabled. Clients now automatically poll for running job cancellations and queue pause, resume, and metadata changes with the default `PollOnly: false`, and skip unsupported notification broadcasts. 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). - Fixed job cancellations received during a fetch being lost before the fetched jobs started. Matching jobs now receive cancellation before work begins. [PR #1397](https://github.com/riverqueue/river/pull/1397). - Fixed SQLite `JobCancel` and `JobCancelTx` notifying running workers through the shared control outbox, so their contexts are cancelled when the transaction commits. [PR #1398](https://github.com/riverqueue/river/pull/1398). - Fixed `UniqueOpts.ByArgs` skipping distinct jobs or failing inserts when JSON keys contain path syntax (like `user.id`), are empty, or come from unnamed tags like `json:",omitempty"`. Unaffected unique keys remain unchanged; affected jobs may be inserted again after upgrading or by old and new clients during a rolling upgrade. [PR #1387](https://github.com/riverqueue/river/pull/1387). diff --git a/client.go b/client.go index a0ee88d92..83bb475e7 100644 --- a/client.go +++ b/client.go @@ -1075,10 +1075,9 @@ func NewClient[TTx any](driver riverdriver.Driver[TTx], config *Config) (*Client } client.queueMaintainerLeader = maintenance.NewQueueMaintainerLeader(archetype, &maintenance.QueueMaintainerLeaderConfig{ - ClientID: config.ID, - Elector: client.elector, - QueueMaintainer: client.queueMaintainer, - RequestResignFunc: client.clientNotifyBundle.RequestResign, + ClientID: config.ID, + Elector: client.elector, + QueueMaintainer: client.queueMaintainer, }) client.services = append(client.services, client.queueMaintainerLeader) client.testSignals.queueMaintainerLeader = &client.queueMaintainerLeader.TestSignals @@ -1136,9 +1135,30 @@ func (c *Client[TTx]) Start(ctx context.Context) error { // available, the client appears to have started even though it's completely // non-functional. Here we try to make an initial assessment of health and // return quickly in case of an apparent problem. - if err := c.driver.GetExecutor().Exec(fetchCtx, "SELECT 1"); err != nil { + executor := c.driver.GetExecutor() + if err := executor.Ping(fetchCtx); err != nil { return fmt.Errorf("error making initial connection to database: %w", err) } + if err := executor.InitDriver(fetchCtx); err != nil { + return fmt.Errorf("error initializing driver: %w", err) + } + + // Database capabilities are only known after initialization. A notifier + // created by NewClient must be removed before any services start if the + // server can't deliver notifications (for example, older Yugabyte). + if c.notifier != nil && !c.driver.SupportsListener() { + c.config.Logger.InfoContext(fetchCtx, "Database does not support listener; entering poll only mode") + c.services = slices.DeleteFunc(c.services, func(service startstop.Service) bool { + return service == c.notifier + }) + c.notifier = nil + if c.elector != nil { + c.elector.SetNotifier(nil) + } + for _, producer := range c.producersByQueueName { + producer.config.Notifier = nil + } + } // Each time we start, we need a fresh completer subscribe channel to // send job completion events on, because the completer will close it @@ -1455,8 +1475,9 @@ func (c *Client[TTx]) Driver() riverdriver.Driver[TTx] { // // If the job is currently running, it is not immediately cancelled, but is // instead marked for cancellation. The client running the job will also be -// notified (via LISTEN/NOTIFY) to cancel the running job's context. Although -// the job's context will be cancelled, since Go does not provide a mechanism to +// notified to cancel the running job's context. When running without a notifier, +// clients poll for cancellation requests every two seconds. Although the job's +// context will be cancelled, since Go does not provide a mechanism to // interrupt a running goroutine the job will continue running until it returns. // As always, it is important for workers to respect context cancellation and // return promptly when the job context is done. @@ -1511,8 +1532,9 @@ func (c *Client[TTx]) JobCancel(ctx context.Context, jobID int64) (*rivertype.Jo // // If the job is currently running, it is not immediately cancelled, but is // instead marked for cancellation. The client running the job will also be -// notified (via LISTEN/NOTIFY) to cancel the running job's context. Although -// the job's context will be cancelled, since Go does not provide a mechanism to +// notified to cancel the running job's context. When running without a notifier, +// clients poll for cancellation requests every two seconds. Although the job's +// context will be cancelled, since Go does not provide a mechanism to // interrupt a running goroutine the job will continue running until it returns. // As always, it is important for workers to respect context cancellation and // return promptly when the job context is done. diff --git a/client_test.go b/client_test.go index a87e0d36e..cf60e0b0c 100644 --- a/client_test.go +++ b/client_test.go @@ -6363,6 +6363,7 @@ func Test_Client_Maintenance(t *testing.T) { // After all retries exhausted, the client should request resignation. client.queueMaintainerLeader.TestSignals.StartRetriesExhausted.WaitOrTimeout() + client.queueMaintainerLeader.TestSignals.ElectedLeader.WaitOrTimeout() }) t.Run("PeriodicJobEnqueuerWithInsertOpts", func(t *testing.T) { @@ -8411,6 +8412,26 @@ func Test_Client_Start_Error(t *testing.T) { require.Equal(t, pgerrcode.InvalidCatalogName, pgErr.Code) }) + t.Run("DatabaseErrorAfterSuccessfulStart", func(t *testing.T) { + t.Parallel() + + dbPool := riversharedtest.DBPoolClone(ctx, t) + driver := NewDriverPollOnly(dbPool) + schema := riverdbtest.TestSchema(ctx, t, driver, nil) + + client, err := NewClient(driver, newTestConfig(t, schema)) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, client.Stop(ctx)) }) + + require.NoError(t, client.Start(ctx)) + require.NoError(t, client.Stop(ctx)) + + dbPool.Close() + + err = client.Start(ctx) + require.ErrorIs(t, err, riverdriver.ErrClosedPool) + }) + t.Run("CanRestartAfterFailure", func(t *testing.T) { t.Parallel() @@ -8437,6 +8458,94 @@ func Test_Client_Start_Error(t *testing.T) { }) } +func Test_Client_YugabyteQueueControl(t *testing.T) { + t.Parallel() + + ctx := context.Background() + for _, testCase := range []struct { + enabled *bool + name string + }{ + {enabled: new(false), name: "Disabled"}, + {enabled: new(true), name: "Enabled"}, + {name: "Unavailable"}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + basePool := riversharedtest.DBPool(ctx, t) + schema := riverdbtest.TestSchema(ctx, t, riverpgxv5.New(basePool), nil) + pool := riversharedtest.DBPoolWithYugabyteVersion(ctx, t, schema, testCase.enabled) + config := newTestConfig(t, schema) + config.queuePollInterval = 20 * time.Millisecond + require.False(t, config.PollOnly) + + client, err := NewClient(riverpgxv5.New(pool), config) + require.NoError(t, err) + require.NotNil(t, client.notifier, "server capability isn't known until Start") + client.testSignals.Init(t) + + // A separate insert-only client prevents local control delivery + // from hiding a missing notification or queue poll. + controller, err := NewClient(riverpgxv5.New(pool), &Config{Schema: schema}) + require.NoError(t, err) + + for run := range 2 { + events := subscribe(t, client) + startClient(ctx, t, client) + client.queueMaintainerLeader.TestSignals.ElectedLeader.WaitOrTimeout() + if testCase.enabled == nil || !*testCase.enabled { + require.Nil(t, client.notifier) + } else { + require.NotNil(t, client.notifier) + } + + if run == 0 { + require.NoError(t, client.Queues().Add("added_after_start", QueueConfig{MaxWorkers: 1})) + } + for _, queue := range []string{QueueDefault, "added_after_start"} { + producer := client.producersByQueueName[queue] + // Only initialize the signals we consume to avoid filling + // unrelated signal buffers while the client is running. + if run == 0 { + producer.testSignals.MetadataChanged.Init(t) + } + + require.NoError(t, controller.QueuePause(ctx, queue, nil)) + event := riversharedtest.WaitOrTimeout(t, events) + require.Equal(t, EventKindQueuePaused, event.Kind) + require.Equal(t, queue, event.Queue.Name) + + inserted, err := controller.Insert(ctx, noOpArgs{}, &InsertOpts{Queue: queue}) + require.NoError(t, err) + job, err := controller.JobGet(ctx, inserted.Job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateAvailable, job.State) + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback(ctx) }) + _, err = controller.QueueUpdateTx(ctx, tx, queue, &QueueUpdateParams{ + Metadata: []byte(fmt.Sprintf(`{"revision":%d}`, run+1)), + }) + require.NoError(t, err) + require.NoError(t, tx.Commit(ctx)) + producer.testSignals.MetadataChanged.WaitOrTimeout() + + require.NoError(t, controller.QueueResume(ctx, queue, nil)) + event = riversharedtest.WaitOrTimeout(t, events) + require.Equal(t, EventKindQueueResumed, event.Kind) + require.Equal(t, queue, event.Queue.Name) + event = riversharedtest.WaitOrTimeout(t, events) + require.Equal(t, EventKindJobCompleted, event.Kind) + require.Equal(t, inserted.Job.ID, event.Job.ID) + } + require.NoError(t, client.Stop(ctx)) + } + }) + } +} + func Test_Config_WithDefaults(t *testing.T) { t.Parallel() diff --git a/docs/yugabyte.md b/docs/yugabyte.md new file mode 100644 index 000000000..f00fe0c5f --- /dev/null +++ b/docs/yugabyte.md @@ -0,0 +1,18 @@ +# YugabyteDB + +The PostgreSQL drivers automatically use polling when YugabyteDB's +`yb_enable_listen_notify` setting is absent or disabled. This includes YugabyteDB +2025.2.1, even with the default `PollOnly: false`. New jobs are picked up on the +`FetchPollInterval`, and queue pause, resume, and metadata changes are picked up +by polling queue settings (every two seconds by default). Running job cancellation +requests are also polled every two seconds, including requests from other clients +and requests made with `JobCancelTx` once committed. + +Native notifications require YugabyteDB **2025.2.3 or later**, with +`ysql_yb_enable_listen_notify=true` on **both Masters and TServers**. The feature +is disabled by default. See [Yugabyte's LISTEN/NOTIFY documentation](https://docs.yugabyte.com/stable/api/ysql/the-sql-language/statements/cmd_listen_notify/) +for the additional replication configuration requirements. `PollOnly: true` +continues to force polling even when native notifications are enabled. + +Database capabilities are cached for the lifetime of the driver. After enabling +notifications, restart the application with a new driver to detect the change. diff --git a/internal/leadership/elector.go b/internal/leadership/elector.go index c79391a7c..260e57453 100644 --- a/internal/leadership/elector.go +++ b/internal/leadership/elector.go @@ -44,7 +44,10 @@ const ( ) type Notification struct { - IsLeader bool + IsLeader bool + // Term identifies a local leadership term and increases on each election. + // Pass it to RequestResign to avoid resigning a subsequent term. + Term uint64 Timestamp time.Time } @@ -233,6 +236,7 @@ type Elector struct { isLeader bool pendingRequestResign bool subscriptions []*Subscription + term uint64 } type leadershipTerm struct { @@ -305,6 +309,29 @@ func trySendWakeup(ctx context.Context, wakeupChan chan struct{}) { } } +// RequestResign requests resignation of this elector's given leadership term. +// It is safe to call concurrently and returns without waiting for resignation. +// Requests for a past term, a follower, or a cancelled context have no effect. +// A zero term requests resignation of whichever term is current, as used for +// database notifications. Local callers should use the term from Listen. +func (e *Elector) RequestResign(ctx context.Context, term uint64) { + e.mu.Lock() + defer e.mu.Unlock() + + if ctx.Err() != nil || !e.isLeader || (term != 0 && term != e.term) { + return + } + + e.pendingRequestResign = true + trySendWakeup(ctx, e.wakeupChan) +} + +// SetNotifier changes the notifier before startup, after database capability +// detection. It must only be called while the elector is stopped. +func (e *Elector) SetNotifier(notifier *notifier.Notifier) { + e.notifier = notifier +} + func (e *Elector) Start(ctx context.Context) error { ctx, shouldStart, started, stopped := e.StartInit(ctx) if !shouldStart { @@ -448,11 +475,7 @@ func (e *Elector) handleLeadershipNotification(ctx context.Context, topic notifi switch notification.Action { case DBNotificationKindRequestResign: - if !e.markPendingRequestResign() { - return - } - - trySendWakeup(ctx, e.wakeupChan) + e.RequestResign(ctx, 0) case DBNotificationKindResigned: // If this a resignation from _this_ client, ignore the change. if notification.LeaderID == e.config.ClientID { @@ -656,6 +679,7 @@ func (e *Elector) Listen() *Subscription { initialNotification := &Notification{ IsLeader: e.isLeader, + Term: e.term, Timestamp: sub.creationTime, } sub.enqueue(initialNotification) @@ -695,30 +719,21 @@ func (e *Elector) leaderTTL() time.Duration { return e.config.ElectInterval + electIntervalTTLPaddingDefault } -func (e *Elector) markPendingRequestResign() bool { - e.mu.Lock() - defer e.mu.Unlock() - - if !e.isLeader { - return false - } - - e.pendingRequestResign = true - return true -} - func (e *Elector) publishLeadershipState(isLeader bool) { notifyTime := time.Now().UTC() e.mu.Lock() defer e.mu.Unlock() e.isLeader = isLeader - if !isLeader { + if isLeader { + e.term++ + } else { e.pendingRequestResign = false } notification := &Notification{ IsLeader: isLeader, + Term: e.term, Timestamp: notifyTime, } diff --git a/internal/leadership/elector_test.go b/internal/leadership/elector_test.go index 3b3bc3498..0dc7b79aa 100644 --- a/internal/leadership/elector_test.go +++ b/internal/leadership/elector_test.go @@ -19,6 +19,7 @@ import ( "github.com/riverqueue/river/rivershared/riversharedtest" "github.com/riverqueue/river/rivershared/startstoptest" "github.com/riverqueue/river/rivershared/testfactory" + "github.com/riverqueue/river/rivershared/testsignal" "github.com/riverqueue/river/rivershared/util/dbutil" "github.com/riverqueue/river/rivertype" ) @@ -242,6 +243,60 @@ func TestElectorHandleLeadershipNotification(t *testing.T) { }) } +func TestElectorRequestResign(t *testing.T) { + t.Parallel() + + setup := func(t *testing.T) *Elector { + t.Helper() + + elector := NewElector(riversharedtest.BaseServiceArchetype(t), nil, nil, &Config{ClientID: "test_client_id"}) + elector.wakeupChan = make(chan struct{}, 1) + elector.publishLeadershipState(true) + return elector + } + + t.Run("CoalescesRequests", func(t *testing.T) { + t.Parallel() + + elector := setup(t) + + var requested testsignal.TestSignal[struct{}] + requested.Init(t) + go func() { + for range 5 { + elector.RequestResign(context.Background(), elector.term) + } + requested.Signal(struct{}{}) + }() + requested.WaitOrTimeout() + require.True(t, elector.pendingRequestResign) + require.Len(t, elector.wakeupChan, 1) + }) + + t.Run("IgnoresCancelledContext", func(t *testing.T) { + t.Parallel() + + elector := setup(t) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + elector.RequestResign(ctx, elector.term) + require.False(t, elector.pendingRequestResign) + require.Empty(t, elector.wakeupChan) + }) + + t.Run("IgnoresFollower", func(t *testing.T) { + t.Parallel() + + elector := setup(t) + + elector.publishLeadershipState(false) + elector.RequestResign(context.Background(), elector.term) + require.False(t, elector.pendingRequestResign) + require.Empty(t, elector.wakeupChan) + }) +} + func TestElectorRunLeaderState(t *testing.T) { t.Parallel() @@ -791,6 +846,50 @@ func testElector[TElectorBundle any]( elector.testSignals.GainedLeadership.WaitOrTimeout() }) + t.Run("RequestResignLocally", func(t *testing.T) { + t.Parallel() + + elector, _ := setup(t, nil) + + startElector(ctx, t, elector) + elector.testSignals.GainedLeadership.WaitOrTimeout() + + // Late subscribers must receive the current term as well. + sub := elector.Listen() + t.Cleanup(sub.Unlisten) + first := riversharedtest.WaitOrTimeout(t, sub.C()) + require.True(t, first.IsLeader) + require.Positive(t, first.Term) + + elector.RequestResign(ctx, first.Term) + require.False(t, riversharedtest.WaitOrTimeout(t, sub.C()).IsLeader) + second := riversharedtest.WaitOrTimeout(t, sub.C()) + require.True(t, second.IsLeader) + require.Greater(t, second.Term, first.Term) + + // A failure from the old maintainer must not resign the new term. + elector.RequestResign(ctx, first.Term) + select { + case notification := <-sub.C(): + t.Fatalf("stale request changed leadership: %+v", notification) + case <-time.After(100 * time.Millisecond): + } + + elector.Stop() + require.False(t, riversharedtest.WaitOrTimeout(t, sub.C()).IsLeader) + require.NoError(t, elector.Start(ctx)) + third := riversharedtest.WaitOrTimeout(t, sub.C()) + require.True(t, third.IsLeader) + require.Greater(t, third.Term, second.Term) + + elector.RequestResign(ctx, second.Term) + select { + case notification := <-sub.C(): + t.Fatalf("stale request after restart changed leadership: %+v", notification) + case <-time.After(100 * time.Millisecond): + } + }) + t.Run("RequestResignStress", func(t *testing.T) { t.Parallel() diff --git a/internal/maintenance/queue_maintainer_leader.go b/internal/maintenance/queue_maintainer_leader.go index c196e90e9..1cc6add13 100644 --- a/internal/maintenance/queue_maintainer_leader.go +++ b/internal/maintenance/queue_maintainer_leader.go @@ -39,11 +39,6 @@ type QueueMaintainerLeaderConfig struct { // QueueMaintainer is the underlying maintainer to start/stop on leadership // changes. QueueMaintainer *QueueMaintainer - - // RequestResignFunc sends a notification requesting leader resignation. - // It's injected from the client because the notification mechanism depends - // on the driver, which the maintenance package doesn't know about. - RequestResignFunc func(ctx context.Context) error } // QueueMaintainerLeader listens for leadership changes and starts/stops the @@ -122,7 +117,7 @@ func (s *QueueMaintainerLeader) Start(ctx context.Context) error { s.mu.Unlock() startWg.Go(func() { - s.tryStart(startCtx, epoch) + s.tryStart(startCtx, epoch, notification.Term) }) default: @@ -140,7 +135,7 @@ func (s *QueueMaintainerLeader) Start(ctx context.Context) error { return nil } -func (s *QueueMaintainerLeader) tryStart(ctx context.Context, epoch int64) { +func (s *QueueMaintainerLeader) tryStart(ctx context.Context, epoch int64, term uint64) { var lastErr error for attempt := 1; attempt <= queueMaintainerMaxStartAttempts; attempt++ { if ctx.Err() != nil { @@ -183,7 +178,7 @@ func (s *QueueMaintainerLeader) tryStart(ctx context.Context, epoch int64) { s.TestSignals.StartRetriesExhausted.Signal(struct{}{}) - if err := s.config.RequestResignFunc(ctx); err != nil { - s.Logger.ErrorContext(ctx, s.Name+": Error requesting leader resignation", slog.String("err", err.Error())) - } + // Recovery must work without database notifications, including in poll-only + // mode. Target this term so a delayed failure cannot resign a newer one. + s.config.Elector.RequestResign(ctx, term) } diff --git a/internal/maintenance/queue_maintainer_leader_test.go b/internal/maintenance/queue_maintainer_leader_test.go index fc2498b8c..c25c17d0c 100644 --- a/internal/maintenance/queue_maintainer_leader_test.go +++ b/internal/maintenance/queue_maintainer_leader_test.go @@ -41,9 +41,6 @@ func TestQueueMaintainerLeader(t *testing.T) { ClientID: "test_client_id", Elector: elector, QueueMaintainer: maintainer, - RequestResignFunc: func(ctx context.Context) error { - return nil - }, }) leader.TestSignals.Init(t) @@ -59,32 +56,14 @@ func TestQueueMaintainerLeader(t *testing.T) { maintainer := NewQueueMaintainer(riversharedtest.BaseServiceArchetype(t), []startstop.Service{failingSvc}) maintainer.StaggerStartupDisable(true) - resignCalled := make(chan struct{}) - archetype := riversharedtest.BaseServiceArchetype(t) - - var ( - dbPool = riversharedtest.DBPool(ctx, t) - driver = riverpgxv5.New(dbPool) - schema = riverdbtest.TestSchema(ctx, t, driver, nil) - ) - - elector := leadership.NewElector(archetype, driver.GetExecutor(), nil, &leadership.Config{ - ClientID: "test_client_id", - Schema: schema, - }) - require.NoError(t, elector.Start(ctx)) - t.Cleanup(elector.Stop) + leader := setup(t, maintainer) + sub := leader.config.Elector.Listen() + t.Cleanup(sub.Unlisten) - leader := NewQueueMaintainerLeader(archetype, &QueueMaintainerLeaderConfig{ - ClientID: "test_client_id", - Elector: elector, - QueueMaintainer: maintainer, - RequestResignFunc: func(ctx context.Context) error { - close(resignCalled) - return nil - }, - }) - leader.TestSignals.Init(t) + // Setup starts the elector. Subscribe before the maintainer can fail so + // the leadership loss is observable even if it immediately wins again. + for !riversharedtest.WaitOrTimeout(t, sub.C()).IsLeader { + } require.NoError(t, leader.Start(ctx)) t.Cleanup(leader.Stop) @@ -97,8 +76,10 @@ func TestQueueMaintainerLeader(t *testing.T) { } leader.TestSignals.StartRetriesExhausted.WaitOrTimeout() - riversharedtest.WaitOrTimeout(t, resignCalled) - require.Equal(t, int64(queueMaintainerMaxStartAttempts), startAttempts.Load()) + require.False(t, riversharedtest.WaitOrTimeout(t, sub.C()).IsLeader) + // A new term permits maintenance startup to retry again. + require.True(t, riversharedtest.WaitOrTimeout(t, sub.C()).IsLeader) + require.GreaterOrEqual(t, startAttempts.Load(), int64(queueMaintainerMaxStartAttempts)) }) t.Run("StartsMaintainerOnLeadershipGain", func(t *testing.T) { diff --git a/internal/rivercommon/river_common.go b/internal/rivercommon/river_common.go index 4f769b391..389efdc8a 100644 --- a/internal/rivercommon/river_common.go +++ b/internal/rivercommon/river_common.go @@ -45,11 +45,6 @@ const ( // MetadataKeyRescueCount records how many times the job has been rescued. MetadataKeyRescueCount = "river:rescue_count" - - // MetadataKeyUniqueNonce is a special metadata key used by the SQLite driver to - // determine whether an upsert is was skipped or not because the `(xmax != 0)` - // trick we use in Postgres doesn't work in SQLite. - MetadataKeyUniqueNonce = "river:unique_nonce" ) type ContextKeyClient struct{} diff --git a/producer.go b/producer.go index 32d8b4a12..6376fbe82 100644 --- a/producer.go +++ b/producer.go @@ -7,7 +7,9 @@ import ( "errors" "fmt" "log/slog" + "maps" "math" + "slices" "strings" "sync" "sync/atomic" @@ -35,6 +37,7 @@ import ( ) const ( + jobCancelPollBatchSize = 1000 producerReportIntervalDefault = 30 * time.Second queuePollIntervalDefault = 2 * time.Second queueReportIntervalDefault = 10 * time.Minute @@ -44,6 +47,7 @@ const ( type producerTestSignals struct { CancelHandledDuringFetch testsignal.TestSignal[int64] // notifies when a cancellation is handled during a fetch DeletedExpiredQueueRecords testsignal.TestSignal[struct{}] // notifies when the producer deletes expired queue records + ExecutorShutdownStarted testsignal.TestSignal[struct{}] // notifies when the producer starts draining active jobs JobFetchTriggered testsignal.TestSignal[struct{}] // notifies when the producer's fetch limiter is triggered via triggerJobFetch MetadataChanged testsignal.TestSignal[struct{}] // notifies when the producer detects a metadata change Paused testsignal.TestSignal[struct{}] // notifies when the producer is paused @@ -57,6 +61,7 @@ type producerTestSignals struct { func (ts *producerTestSignals) Init(tb testutil.TestingTB) { ts.DeletedExpiredQueueRecords.Init(tb) + ts.ExecutorShutdownStarted.Init(tb) ts.JobFetchTriggered.Init(tb) ts.MetadataChanged.Init(tb) ts.Paused.Init(tb) @@ -105,8 +110,8 @@ type producerConfig struct { QueueEventCallback func(event *Event) // QueuePollInterval is the amount of time between periodic checks for - // queue setting changes. This is only used in poll-only mode (when no - // notifier is provided). + // queue setting changes and job cancellation requests. This is only used in + // poll-only mode (when no notifier is provided). QueuePollInterval time.Duration // QueueReportInterval is the amount of time between periodic reports // of the queue status. @@ -191,8 +196,10 @@ type producer struct { baseservice.BaseService startstop.BaseStartStop - // Jobs which are currently being worked. Only used by main goroutine. - activeJobs map[int64]*jobexecutor.JobExecutor + // Jobs which are currently being worked. The polling goroutine snapshots + // their IDs; only the main goroutine accesses the executors themselves. + activeJobsMu sync.Mutex + activeJobs map[int64]*jobexecutor.JobExecutor completer jobcompleter.JobCompleter config *producerConfig @@ -205,8 +212,8 @@ type producer struct { pilot riverpilot.Pilot workers *Workers - // Receives job IDs to cancel. Written by notifier goroutine, only read from - // main goroutine. + // Receives job IDs to cancel. Written by notifier and polling goroutines, + // only read from the main goroutine. cancelCh chan int64 // Set to true when the producer thinks it should trigger another fetch as @@ -426,11 +433,14 @@ func (p *producer) StartWorkContext(fetchCtx, workCtx context.Context) error { subroutineWG.Add(1) go p.pollForSettingChanges(subroutineCtx, &subroutineWG, initiallyPaused, initialMetadata) + + subroutineWG.Add(1) + go p.pollForJobCancellations(subroutineCtx, &subroutineWG) } p.fetchAndRunLoop(fetchCtx, workCtx) p.Logger.DebugContext(workCtx, p.Name+": Entering shutdown loop", slog.String("queue", p.config.Queue), slog.Int64("id", p.id.Load())) - p.executorShutdownLoop() + p.executorShutdownLoop(workCtx) p.Logger.DebugContext(workCtx, p.Name+": Shutdown loop exited, awaiting subroutines", slog.String("queue", p.config.Queue), slog.Int64("id", p.id.Load())) cancelSubroutines(fmt.Errorf("producer stopped: %w", startstop.ErrStop)) @@ -684,12 +694,25 @@ func (p *producer) innerFetchLoop(workCtx context.Context, fetchResultCh chan pr } } -func (p *producer) executorShutdownLoop() { +func (p *producer) executorShutdownLoop(ctx context.Context) { + p.testSignals.ExecutorShutdownStarted.Signal(struct{}{}) + // No more jobs will be fetched or executed. However, we must wait for all // in-progress jobs to complete. - for len(p.activeJobs) != 0 { - result := <-p.jobResultCh - p.removeActiveJob(result) + for { + p.activeJobsMu.Lock() + numActiveJobs := len(p.activeJobs) + p.activeJobsMu.Unlock() + if numActiveJobs == 0 { + return + } + + select { + case jobID := <-p.cancelCh: + p.maybeCancelJob(ctx, jobID) + case result := <-p.jobResultCh: + p.removeActiveJob(result) + } } } @@ -745,12 +768,16 @@ func (p *producer) finalizeShutdown(ctx context.Context) { func (p *producer) addActiveJob(id int64, executor *jobexecutor.JobExecutor) { p.numJobsActive.Add(1) + p.activeJobsMu.Lock() p.activeJobs[id] = executor + p.activeJobsMu.Unlock() } func (p *producer) removeActiveJob(job *rivertype.JobRow) { + p.activeJobsMu.Lock() executor := p.activeJobs[job.ID] delete(p.activeJobs, job.ID) + p.activeJobsMu.Unlock() if executor == nil || executor.TryCloseSlot() { p.numJobsActive.Add(-1) } @@ -788,7 +815,9 @@ func (p *producer) handleWorkerUnstuck() { } func (p *producer) maybeCancelJob(ctx context.Context, id int64) bool { + p.activeJobsMu.Lock() executor, ok := p.activeJobs[id] + p.activeJobsMu.Unlock() if !ok { return false } @@ -975,6 +1004,57 @@ func (p *producer) handleWorkerDone(job *rivertype.JobRow) { p.jobResultCh <- job } +// Cancellation polling is independent of fetching: saturated and paused producers +// must still cancel jobs, including while draining during graceful shutdown. +func (p *producer) pollForJobCancellations(ctx context.Context, wg *sync.WaitGroup) { + defer wg.Done() + + ticker := time.NewTicker(p.config.QueuePollInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if err := p.pollForJobCancellationsOnce(ctx); err != nil { + if !errors.Is(context.Cause(ctx), startstop.ErrStop) { + p.Logger.ErrorContext(ctx, p.Name+": Error fetching job cancellation requests", slog.String("err", err.Error())) + } + continue + } + } + } +} + +func (p *producer) pollForJobCancellationsOnce(ctx context.Context) error { + p.activeJobsMu.Lock() + jobIDs := slices.Collect(maps.Keys(p.activeJobs)) + p.activeJobsMu.Unlock() + + // Bound query size (including SQLite's parameter count) and skip database + // access entirely when there are no active jobs. + for batch := range slices.Chunk(jobIDs, jobCancelPollBatchSize) { + cancelIDs, err := timeoututil.WithTimeoutV(ctx, 10*time.Second, p.Name+".pollForJobCancellations", func(ctx context.Context) ([]int64, error) { + return p.exec.JobGetCancelRequested(ctx, &riverdriver.JobGetCancelRequestedParams{ + ID: batch, + Schema: p.config.Schema, + }) + }) + if err != nil { + return err + } + + for _, id := range cancelIDs { + select { + case <-ctx.Done(): + return ctx.Err() + case p.cancelCh <- id: + } + } + } + return nil +} + func (p *producer) pollForSettingChanges(ctx context.Context, wg *sync.WaitGroup, lastPaused bool, lastMetadata []byte) { defer wg.Done() diff --git a/producer_test.go b/producer_test.go index 8337044c5..e4f9cd0fd 100644 --- a/producer_test.go +++ b/producer_test.go @@ -480,6 +480,74 @@ func testProducer(t *testing.T, makeProducer func(ctx context.Context, t *testin } }) + t.Run("CancellationPolling", func(t *testing.T) { + t.Parallel() + + for _, state := range []string{"Paused", "Running", "Stopping"} { + t.Run(state, func(t *testing.T) { + t.Parallel() + + producer, bundle := setup(t) + if producer.config.Notifier != nil { + t.Skip("requires polling without a notifier") + } + producer.config.FetchPollInterval = time.Hour + producer.config.MaxWorkers = 1 + producer.config.QueuePollInterval = 20 * time.Millisecond + + var jobStarted testsignal.TestSignal[int64] + var workerErr testsignal.TestSignal[error] + jobStarted.Init(t) + workerErr.Init(t) + + type JobArgs struct { + testutil.JobArgsReflectKind[JobArgs] + } + AddWorker(bundle.workers, WorkFunc(func(ctx context.Context, job *Job[JobArgs]) error { + jobStarted.Signal(job.ID) + <-ctx.Done() + workerErr.Signal(context.Cause(ctx)) + return ctx.Err() + })) + + fetchCtx, fetchCancel := context.WithCancel(ctx) + defer fetchCancel() + workCtx, workCancel := context.WithCancel(ctx) + defer workCancel() + + mustInsert(ctx, t, producer, bundle, &JobArgs{}) + startProducer(t, fetchCtx, workCtx, producer) + jobID := jobStarted.WaitOrTimeout() + require.Positive(t, jobID) + + switch state { + case "Paused": + require.NoError(t, bundle.exec.QueuePause(ctx, &riverdriver.QueuePauseParams{ + Name: producer.config.Queue, + Schema: producer.config.Schema, + })) + producer.testSignals.Paused.WaitOrTimeout() + case "Stopping": + fetchCancel() + producer.testSignals.ExecutorShutdownStarted.WaitOrTimeout() + } + + // Use the executor directly so no local client shortcut can cancel + // the worker. No more jobs can be fetched in any of these states. + _, err := bundle.exec.JobCancel(ctx, &riverdriver.JobCancelParams{ + ID: jobID, + CancelAttemptedAt: time.Now().UTC(), + ControlTopic: string(notifier.NotificationTopicControl), + Schema: producer.config.Schema, + }) + require.NoError(t, err) + require.ErrorIs(t, workerErr.WaitOrTimeout(), rivertype.ErrJobCancelledRemotely) + update := riversharedtest.WaitOrTimeout(t, bundle.jobUpdates) + require.Equal(t, rivertype.JobStateCancelled, update.Job.State) + }) + } + }) + t.Run("CancelledWorkContextCancelsJob", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/postgres_capabilities.go b/riverdriver/postgres_capabilities.go new file mode 100644 index 000000000..f22ad5be3 --- /dev/null +++ b/riverdriver/postgres_capabilities.go @@ -0,0 +1,30 @@ +package riverdriver + +import "strings" + +// PostgresCapabilities describes database features detected by PostgreSQL drivers. +// Drivers cache a successful detection for their lifetime. +type PostgresCapabilities struct { + SupportsListenNotify bool + UniqueInsertMode UniqueInsertMode +} + +// NewPostgresCapabilities detects features from the server's product, version, +// and settings. +func NewPostgresCapabilities(product string, version int32, ybListenNotifyEnabled bool) *PostgresCapabilities { + // Yugabyte's native notifications require 2025.2.3 or later with + // ysql_yb_enable_listen_notify=true on both Masters and TServers. The + // yb_enable_listen_notify setting is false when absent on older versions, + // so clients automatically fall back to polling even with PollOnly false. + // Capabilities are cached per driver; a new driver is needed after changing + // the setting to detect newly enabled notification support. + return &PostgresCapabilities{ + SupportsListenNotify: !postgresProductIsYugabyte(product) || ybListenNotifyEnabled, + UniqueInsertMode: UniqueInsertModeFromProductAndVersion(product, version), + } +} + +func postgresProductIsYugabyte(product string) bool { + productLower := strings.ToLower(product) + return strings.Contains(productLower, "-yb") || strings.Contains(productLower, "yugabyte") +} diff --git a/riverdriver/postgres_capabilities_test.go b/riverdriver/postgres_capabilities_test.go new file mode 100644 index 000000000..a7369595b --- /dev/null +++ b/riverdriver/postgres_capabilities_test.go @@ -0,0 +1,35 @@ +package riverdriver + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNewPostgresCapabilities(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + expectedListenNotify bool + expectedUniqueMode UniqueInsertMode + name string + product string + version int32 + ybListenNotify bool + }{ + {expectedListenNotify: true, expectedUniqueMode: UniqueInsertModeXmax, name: "Postgres15", product: "PostgreSQL 15.12", version: 150_012}, + {expectedListenNotify: true, expectedUniqueMode: UniqueInsertModeReturningOld, name: "Postgres18", product: "PostgreSQL 18.0", version: 180_000}, + {expectedUniqueMode: UniqueInsertModeMetadataNonce, name: "YugabyteDisabled", product: "PostgreSQL 15.12-YB-2025.2.3.0-b1", version: 150_012}, + {expectedListenNotify: true, expectedUniqueMode: UniqueInsertModeMetadataNonce, name: "YugabyteEnabled", product: "PostgreSQL 15.12-YB-2025.2.3.0-b1", version: 150_012, ybListenNotify: true}, + {expectedUniqueMode: UniqueInsertModeMetadataNonce, name: "YugabyteOld", product: "PostgreSQL 15.12-YB-2025.2.1.0-b1", version: 150_012}, + {expectedUniqueMode: UniqueInsertModeMetadataNonce, name: "YugabyteProductName", product: "YugabyteDB", version: 150_012}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + capabilities := NewPostgresCapabilities(testCase.product, testCase.version, testCase.ybListenNotify) + require.Equal(t, testCase.expectedListenNotify, capabilities.SupportsListenNotify) + require.Equal(t, testCase.expectedUniqueMode, capabilities.UniqueInsertMode) + }) + } +} diff --git a/riverdriver/river_driver_interface.go b/riverdriver/river_driver_interface.go index 49c3067b6..ad18b146e 100644 --- a/riverdriver/river_driver_interface.go +++ b/riverdriver/river_driver_interface.go @@ -155,6 +155,8 @@ type Driver[TTx any] interface { // SupportsListener gets whether this driver supports a listener. Drivers // that don't support a listener support poll only mode only. + // Before InitDriver, this reports the driver's default capability; callers + // must recheck after initialization to account for the database server. // // API is not stable. DO NOT USE. SupportsListener() bool @@ -167,6 +169,8 @@ type Driver[TTx any] interface { // notification mechanism, it will still broadcast in case there are other // clients/drivers on the database that do support a listener. If // notifications can't be supported at all, no broadcast attempt is made. + // Like SupportsListener, this is refined by InitDriver. Executors also + // initialize lazily before sending notifications for clients that never start. // // API is not stable. DO NOT USE. SupportsListenNotify() bool @@ -224,6 +228,11 @@ type Executor interface { IndexReindex(ctx context.Context, params *IndexReindexParams) error IndexReindexArtifacts(ctx context.Context, params *IndexReindexArtifactsParams) ([]string, error) + // InitDriver initializes driver-specific state using information read from + // the database. Implementations must be safe to call concurrently and + // repeatedly, and should cache successfully initialized state. + InitDriver(ctx context.Context) error + JobCancel(ctx context.Context, params *JobCancelParams) (*rivertype.JobRow, error) JobCountByAllStates(ctx context.Context, params *JobCountByAllStatesParams) (map[rivertype.JobState]int, error) JobCountByQueueAndState(ctx context.Context, params *JobCountByQueueAndStateParams) ([]*JobCountByQueueAndStateResult, error) @@ -235,6 +244,11 @@ type Executor interface { JobGetByID(ctx context.Context, params *JobGetByIDParams) (*rivertype.JobRow, error) JobGetByIDMany(ctx context.Context, params *JobGetByIDManyParams) ([]*rivertype.JobRow, error) JobGetByKindMany(ctx context.Context, params *JobGetByKindManyParams) ([]*rivertype.JobRow, error) + + // JobGetCancelRequested returns IDs of running jobs with a cancellation request, + // restricted to the provided IDs. + JobGetCancelRequested(ctx context.Context, params *JobGetCancelRequestedParams) ([]int64, error) + JobGetStuck(ctx context.Context, params *JobGetStuckParams) ([]*rivertype.JobRow, error) JobInsertFastMany(ctx context.Context, params *JobInsertFastManyParams) ([]*JobInsertFastResult, error) JobInsertFastManyNoReturning(ctx context.Context, params *JobInsertFastManyParams) (int, error) @@ -289,6 +303,10 @@ type Executor interface { NotificationDeleteBefore(ctx context.Context, params *NotificationDeleteBeforeParams) (int, error) NotifyMany(ctx context.Context, params *NotifyManyParams) error + + // Ping checks that the database is reachable. + Ping(ctx context.Context) error + PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) QueueCreateOrSetUpdatedAt(ctx context.Context, params *QueueCreateOrSetUpdatedAtParams) (*rivertype.Queue, error) @@ -447,6 +465,12 @@ type JobGetByKindManyParams struct { Schema string } +// JobGetCancelRequestedParams restricts cancellation checks to specific job IDs. +type JobGetCancelRequestedParams struct { + ID []int64 + Schema string +} + type JobGetStuckParams struct { AfterID int64 Max int diff --git a/riverdriver/riverdatabasesql/internal/dbsqlc/pg_misc.sql.go b/riverdriver/riverdatabasesql/internal/dbsqlc/pg_misc.sql.go index e2542e763..463016675 100644 --- a/riverdriver/riverdatabasesql/internal/dbsqlc/pg_misc.sql.go +++ b/riverdriver/riverdatabasesql/internal/dbsqlc/pg_misc.sql.go @@ -21,6 +21,26 @@ func (q *Queries) PGAdvisoryXactLock(ctx context.Context, db DBTX, key int64) er return err } +const pGGetProductAndVersion = `-- name: PGGetProductAndVersion :one +SELECT + version()::text AS product, + current_setting('server_version_num')::int AS version_num, + coalesce(current_setting('yb_enable_listen_notify', true), 'off')::boolean AS yb_listen_notify_enabled +` + +type PGGetProductAndVersionRow struct { + Product string + VersionNum int32 + YbListenNotifyEnabled bool +} + +func (q *Queries) PGGetProductAndVersion(ctx context.Context, db DBTX) (*PGGetProductAndVersionRow, error) { + row := db.QueryRowContext(ctx, pGGetProductAndVersion) + var i PGGetProductAndVersionRow + err := row.Scan(&i.Product, &i.VersionNum, &i.YbListenNotifyEnabled) + return &i, err +} + const pGNotifyMany = `-- name: PGNotifyMany :exec WITH topic_to_notify AS ( SELECT diff --git a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go index 024a45a7d..e43848ee9 100644 --- a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go @@ -24,10 +24,10 @@ WITH locked_job AS ( notification AS ( SELECT id, - pg_notify( - concat(coalesce($2::text, current_schema()), '.', $3::text), + CASE WHEN $2::boolean THEN pg_notify( + concat(coalesce($3::text, current_schema()), '.', $4::text), json_build_object('action', 'cancel', 'job_id', id, 'queue', queue)::text - ) + ) END FROM locked_job WHERE @@ -40,10 +40,10 @@ updated_job AS ( -- If the job is actively running, we want to let its current client and -- producer handle the cancellation. Otherwise, immediately cancel it. state = CASE WHEN state = 'running' THEN state ELSE 'cancelled' END, - finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE coalesce($4::timestamptz, now()) END, + finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE coalesce($5::timestamptz, now()) END, -- Mark the job as cancelled by query so that the rescuer knows not to -- rescue it, even if it gets stuck in the running state: - metadata = jsonb_set(metadata, '{cancel_attempted_at}'::text[], $5::jsonb, true) + metadata = jsonb_set(metadata, '{cancel_attempted_at}'::text[], $6::jsonb, true) FROM notification WHERE river_job.id = notification.id RETURNING river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states @@ -59,6 +59,7 @@ FROM updated_job type JobCancelParams struct { ID int64 + Notify bool Schema sql.NullString ControlTopic string Now *time.Time @@ -68,6 +69,7 @@ type JobCancelParams struct { func (q *Queries) JobCancel(ctx context.Context, db DBTX, arg *JobCancelParams) (*RiverJob, error) { row := db.QueryRowContext(ctx, jobCancel, arg.ID, + arg.Notify, arg.Schema, arg.ControlTopic, arg.Now, @@ -605,6 +607,38 @@ func (q *Queries) JobGetByKindMany(ctx context.Context, db DBTX, kind []string) return items, nil } +const jobGetCancelRequested = `-- name: JobGetCancelRequested :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE id = any($1::bigint[]) + AND metadata ? 'cancel_attempted_at' + AND state = 'running' +ORDER BY id +` + +func (q *Queries) JobGetCancelRequested(ctx context.Context, db DBTX, id []int64) ([]int64, error) { + rows, err := db.QueryContext(ctx, jobGetCancelRequested, pq.Array(id)) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const jobGetStuck = `-- name: JobGetStuck :many SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states FROM /* TEMPLATE: schema */river_job @@ -718,7 +752,9 @@ ON CONFLICT (unique_key) AND /* TEMPLATE: schema */river_job_state_in_bitmask(unique_states, state) -- Something needs to be updated for a row to be returned on a conflict. DO UPDATE SET kind = EXCLUDED.kind -RETURNING river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states, (xmax != 0) AS unique_skipped_as_duplicate +RETURNING + river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states, + /* TEMPLATE_BEGIN: unique_skipped_as_duplicate */ (xmax != 0) /* TEMPLATE_END */ AS unique_skipped_as_duplicate ` type JobInsertFastManyParams struct { diff --git a/riverdriver/riverdatabasesql/internal/dbsqlc/river_leader.sql.go b/riverdriver/riverdatabasesql/internal/dbsqlc/river_leader.sql.go index a54d2bd7b..8ec6a55f7 100644 --- a/riverdriver/riverdatabasesql/internal/dbsqlc/river_leader.sql.go +++ b/riverdriver/riverdatabasesql/internal/dbsqlc/river_leader.sql.go @@ -157,10 +157,10 @@ WITH currently_held_leaders AS ( FOR UPDATE ), notified_resignations AS ( - SELECT pg_notify( - concat(coalesce($3::text, current_schema()), '.', $4::text), + SELECT CASE WHEN $3::boolean THEN pg_notify( + concat(coalesce($4::text, current_schema()), '.', $5::text), json_build_object('leader_id', leader_id, 'action', 'resigned')::text - ) + ) END FROM currently_held_leaders ) DELETE FROM /* TEMPLATE: schema */river_leader USING notified_resignations @@ -169,6 +169,7 @@ DELETE FROM /* TEMPLATE: schema */river_leader USING notified_resignations type LeaderResignParams struct { ElectedAt time.Time LeaderID string + Notify bool Schema sql.NullString LeadershipTopic string } @@ -177,6 +178,7 @@ func (q *Queries) LeaderResign(ctx context.Context, db DBTX, arg *LeaderResignPa result, err := db.ExecContext(ctx, leaderResign, arg.ElectedAt, arg.LeaderID, + arg.Notify, arg.Schema, arg.LeadershipTopic, ) diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver.go b/riverdriver/riverdatabasesql/river_database_sql_driver.go index 07024be99..fa4e510ae 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver.go @@ -16,6 +16,7 @@ import ( "io/fs" "math" "strings" + "sync/atomic" "time" "github.com/jackc/pgx/v5/pgxpool" @@ -28,6 +29,7 @@ import ( "github.com/riverqueue/river/rivershared/uniquestates" "github.com/riverqueue/river/rivershared/util/dbutil" "github.com/riverqueue/river/rivershared/util/ptrutil" + "github.com/riverqueue/river/rivershared/util/randutil" "github.com/riverqueue/river/rivershared/util/savepointutil" "github.com/riverqueue/river/rivershared/util/sliceutil" "github.com/riverqueue/river/rivertype" @@ -38,9 +40,10 @@ var migrationFS embed.FS // Driver is an implementation of riverdriver.Driver for database/sql. type Driver struct { - dbPool *sql.DB - listenerDriver *riverpgxv5.Driver - replacer sqlctemplate.Replacer + dbPool *sql.DB + listenerDriver *riverpgxv5.Driver + postgresCapabilities atomic.Pointer[riverdriver.PostgresCapabilities] + replacer sqlctemplate.Replacer } // New returns a new database/sql River driver for use with River. @@ -135,8 +138,12 @@ func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, er return fmt.Sprintf("%s = any(@%s)", column, column), pq.Array(values), nil } -func (d *Driver) SupportsListener() bool { return d.listenerDriver != nil } -func (d *Driver) SupportsListenNotify() bool { return true } +func (d *Driver) SupportsListener() bool { return d.listenerDriver != nil && d.SupportsListenNotify() } + +func (d *Driver) SupportsListenNotify() bool { + capabilities := d.postgresCapabilities.Load() + return capabilities == nil || capabilities.SupportsListenNotify +} func (d *Driver) TimePrecision() time.Duration { return time.Microsecond } func (d *Driver) UnwrapExecutor(tx *sql.Tx) riverdriver.ExecutorTx { @@ -247,7 +254,17 @@ func (e *Executor) IndexesExist(ctx context.Context, params *riverdriver.Indexes return exists, nil } +func (e *Executor) InitDriver(ctx context.Context) error { + _, err := e.getPostgresCapabilities(ctx) + return err +} + func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelParams) (*rivertype.JobRow, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return nil, err + } + cancelledAt, err := params.CancelAttemptedAt.MarshalJSON() if err != nil { return nil, err @@ -257,6 +274,7 @@ func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelP ID: params.ID, CancelAttemptedAt: string(cancelledAt), ControlTopic: params.ControlTopic, + Notify: capabilities.SupportsListenNotify, Now: params.Now, Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, }) @@ -389,6 +407,11 @@ func (e *Executor) JobGetByKindMany(ctx context.Context, params *riverdriver.Job return sliceutil.MapError(jobs, jobRowFromInternal) } +func (e *Executor) JobGetCancelRequested(ctx context.Context, params *riverdriver.JobGetCancelRequestedParams) ([]int64, error) { + ids, err := dbsqlc.New().JobGetCancelRequested(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.ID) + return ids, interpretError(err) +} + func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetStuckParams) ([]*rivertype.JobRow, error) { jobs, err := dbsqlc.New().JobGetStuck(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobGetStuckParams{ AfterID: params.AfterID, @@ -402,6 +425,17 @@ func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetSt } func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.JobInsertFastManyParams) ([]*riverdriver.JobInsertFastResult, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return nil, err + } + uniqueInsertMode := capabilities.UniqueInsertMode + + var uniqueNonce string + if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce { + uniqueNonce = randutil.Hex(8) + } + insertJobsParams := &dbsqlc.JobInsertFastManyParams{ ID: make([]int64, len(params.Jobs)), Args: make([]string, len(params.Jobs)), @@ -442,7 +476,16 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo insertJobsParams.CreatedAt[i] = createdAt insertJobsParams.Kind[i] = params.Kind insertJobsParams.MaxAttempts[i] = int16(min(params.MaxAttempts, math.MaxInt16)) //nolint:gosec - insertJobsParams.Metadata[i] = cmp.Or(string(params.Metadata), "{}") + metadata := []byte(cmp.Or(string(params.Metadata), "{}")) + if uniqueNonce != "" { + var err error + metadata, err = riverdriver.UniqueInsertMetadataWithNonce(metadata, uniqueNonce) + if err != nil { + return nil, err + } + } + + insertJobsParams.Metadata[i] = string(metadata) insertJobsParams.Priority[i] = int16(min(params.Priority, math.MaxInt16)) //nolint:gosec insertJobsParams.Queue[i] = params.Queue insertJobsParams.ScheduledAt[i] = scheduledAt @@ -452,6 +495,9 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo insertJobsParams.UniqueStates[i] = int32(params.UniqueStates) } + ctx = sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "unique_skipped_as_duplicate": {Value: uniqueInsertMode.SQL(), Stable: true}, + }, nil) items, err := dbsqlc.New().JobInsertFastMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, insertJobsParams) if err != nil { return nil, interpretError(err) @@ -462,7 +508,13 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo if err != nil { return nil, err } - return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: row.UniqueSkippedAsDuplicate}, nil + + uniqueSkippedAsDuplicate := row.UniqueSkippedAsDuplicate + if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce { + uniqueSkippedAsDuplicate = riverdriver.UniqueInsertMetadataIsDuplicate(job.Metadata, uniqueNonce) + } + + return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: uniqueSkippedAsDuplicate}, nil }) } @@ -833,10 +885,16 @@ func (e *Executor) LeaderInsert(ctx context.Context, params *riverdriver.LeaderI } func (e *Executor) LeaderResign(ctx context.Context, params *riverdriver.LeaderResignParams) (bool, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return false, err + } + numResigned, err := dbsqlc.New().LeaderResign(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.LeaderResignParams{ ElectedAt: params.ElectedAt, LeaderID: params.LeaderID, LeadershipTopic: params.LeadershipTopic, + Notify: capabilities.SupportsListenNotify, Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, }) if err != nil { @@ -929,6 +987,14 @@ func (e *Executor) NotificationDeleteBefore(ctx context.Context, params *riverdr } func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyManyParams) error { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return err + } + if !capabilities.SupportsListenNotify { + return nil + } + return dbsqlc.New().PGNotifyMany(ctx, e.dbtx, &dbsqlc.PGNotifyManyParams{ Payload: params.Payload, Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, @@ -936,6 +1002,10 @@ func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyMan }) } +func (e *Executor) Ping(ctx context.Context) error { + return e.Exec(ctx, "SELECT 1") +} + func (e *Executor) PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) { err := dbsqlc.New().PGAdvisoryXactLock(ctx, e.dbtx, key) return &struct{}{}, interpretError(err) @@ -1094,6 +1164,29 @@ func (e *Executor) TableTruncate(ctx context.Context, params *riverdriver.TableT return interpretError(err) } +func (e *Executor) getPostgresCapabilities(ctx context.Context) (*riverdriver.PostgresCapabilities, error) { + if e.driver != nil { + if capabilities := e.driver.postgresCapabilities.Load(); capabilities != nil { + return capabilities, nil + } + } + + productAndVersion, err := dbsqlc.New().PGGetProductAndVersion(ctx, e.dbtx) + if err != nil { + return nil, interpretError(err) + } + + capabilities := riverdriver.NewPostgresCapabilities(productAndVersion.Product, productAndVersion.VersionNum, productAndVersion.YbListenNotifyEnabled) + if e.driver != nil { + // Concurrent callers may both detect, but the first successful result + // becomes the driver's cached capabilities. Don't hold a lock while + // querying: a caller may already hold the pool's only connection. + e.driver.postgresCapabilities.CompareAndSwap(nil, capabilities) + capabilities = e.driver.postgresCapabilities.Load() + } + return capabilities, nil +} + type ExecutorTx struct { Executor diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver_test.go b/riverdriver/riverdatabasesql/river_database_sql_driver_test.go index a452dd017..a81d9fe50 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver_test.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver_test.go @@ -5,18 +5,34 @@ import ( "database/sql" "errors" "testing" + "time" "github.com/jackc/pgx/v5/pgxpool" + _ "github.com/jackc/pgx/v5/stdlib" "github.com/stretchr/testify/require" "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/riversharedtest" "github.com/riverqueue/river/rivershared/sqlctemplate" + "github.com/riverqueue/river/rivershared/testsignal" + "github.com/riverqueue/river/rivershared/util/urlutil" "github.com/riverqueue/river/rivertype" ) // Verify interface compliance. var _ riverdriver.Driver[*sql.Tx] = New(nil) +type executorInitDriverTestDBTX struct { + *sql.DB + + QueryRowStarted testsignal.TestSignal[struct{}] +} + +func (d *executorInitDriverTestDBTX) QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row { + d.QueryRowStarted.Signal(struct{}{}) + return d.DB.QueryRowContext(ctx, query, args...) +} + func TestNew(t *testing.T) { t.Parallel() @@ -48,6 +64,51 @@ func TestNew(t *testing.T) { }) } +func TestExecutor_InitDriverDoesNotBlockTransaction(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dbPool, err := sql.Open("pgx", urlutil.DatabaseSQLCompatibleURL(riversharedtest.TestDatabaseURL())) + require.NoError(t, err) + dbPool.SetMaxOpenConns(1) + t.Cleanup(func() { require.NoError(t, dbPool.Close()) }) + + driver := New(dbPool) + tx, err := dbPool.BeginTx(ctx, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback() }) + + poolDBTX := &executorInitDriverTestDBTX{DB: dbPool} + poolDBTX.QueryRowStarted.Init(t) + poolExecutor := &Executor{ + dbPool: dbPool, + dbtx: templateReplaceWrapper{dbtx: poolDBTX, replacer: &driver.replacer}, + driver: driver, + } + + initCtx, initCancel := context.WithTimeout(ctx, 10*time.Second) + t.Cleanup(initCancel) + + var poolInitFinished testsignal.TestSignal[error] + poolInitFinished.Init(t) + go func() { poolInitFinished.Signal(poolExecutor.InitDriver(initCtx)) }() + poolDBTX.QueryRowStarted.WaitOrTimeout() + + var txInitFinished testsignal.TestSignal[error] + txInitFinished.Init(t) + go func() { txInitFinished.Signal(driver.UnwrapExecutor(tx).InitDriver(ctx)) }() + + select { + case err := <-txInitFinished.WaitC(): + require.NoError(t, err) + case <-time.After(2 * time.Second): + require.FailNow(t, "transactional driver initialization blocked behind pool initialization") + } + + require.NoError(t, tx.Rollback()) + require.NoError(t, poolInitFinished.WaitOrTimeout()) +} + func TestNewWithPgxListener(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverdatabasesql/yugabyte_compatibility_test.go b/riverdriver/riverdatabasesql/yugabyte_compatibility_test.go new file mode 100644 index 000000000..a4d9fcf51 --- /dev/null +++ b/riverdriver/riverdatabasesql/yugabyte_compatibility_test.go @@ -0,0 +1,57 @@ +package riverdatabasesql + +import ( + "fmt" + "io/fs" + "os" + "regexp" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestYugabyteCompatibility(t *testing.T) { + t.Parallel() + + // YugabyteDB doesn't expose PostgreSQL's transaction-related system columns: + // https://docs.yugabyte.com/stable/yugabyte-voyager/known-issues/postgresql/#system-columns-is-not-yet-supported + // The unique insert query may use xmax only because the entire expression is + // replaced when the driver detects YugabyteDB. + var ( + uniqueInsertModeTemplateRE = regexp.MustCompile(`(?s)/\*\s*TEMPLATE_BEGIN: unique_skipped_as_duplicate\s*\*/.*?/\*\s*TEMPLATE_END\s*\*/`) + unsupportedSystemColumnRE = regexp.MustCompile(`(?i)\b(?:cmax|cmin|ctid|xmax|xmin)\b`) + ) + + sourceRoot, err := os.OpenRoot(".") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, sourceRoot.Close()) }) + + var violations []string + err = fs.WalkDir(sourceRoot.FS(), ".", func(path string, entry fs.DirEntry, err error) error { + if err != nil { + return err + } + if entry.IsDir() || (!strings.HasSuffix(path, ".go") && !strings.HasSuffix(path, ".sql")) || strings.HasSuffix(path, "_test.go") { + return nil + } + + contents, err := sourceRoot.ReadFile(path) + if err != nil { + return err + } + contents = uniqueInsertModeTemplateRE.ReplaceAll(contents, nil) + + for lineNum, line := range strings.Split(string(contents), "\n") { + for _, column := range unsupportedSystemColumnRE.FindAllString(line, -1) { + violations = append(violations, fmt.Sprintf("%s:%d: %s", path, lineNum+1, column)) + } + } + return nil + }) + require.NoError(t, err) + + require.Empty(t, violations, + "YugabyteDB-incompatible PostgreSQL system columns must only appear inside SQL templates that replace them for YugabyteDB", + ) +} diff --git a/riverdriver/riverdrivertest/client_cancel_test.go b/riverdriver/riverdrivertest/client_cancel_test.go new file mode 100644 index 000000000..601c5c882 --- /dev/null +++ b/riverdriver/riverdrivertest/client_cancel_test.go @@ -0,0 +1,102 @@ +package riverdrivertest + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/riverqueue/river" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/testsignal" + "github.com/riverqueue/river/rivershared/util/testutil" + "github.com/riverqueue/river/rivertype" +) + +type cancelRunningJobArgs struct { + testutil.JobArgsReflectKind[cancelRunningJobArgs] +} + +func exerciseClientCancelRunningJob[TTx any](ctx context.Context, t *testing.T, driver riverdriver.Driver[TTx], schema string, pollOnly, transactional bool) { + t.Helper() + + config := newTestConfig(t, schema) + config.FetchPollInterval = time.Minute + config.PollOnly = pollOnly + config.Queues = map[string]river.QueueConfig{river.QueueDefault: {MaxWorkers: 1}} + + var jobStarted, jobContextCancelled testsignal.TestSignal[int64] + jobStarted.Init(t) + jobContextCancelled.Init(t) + + river.AddWorker(config.Workers, river.WorkFunc(func(ctx context.Context, job *river.Job[cancelRunningJobArgs]) error { + jobStarted.Signal(job.ID) + <-ctx.Done() + jobContextCancelled.Signal(job.ID) + return ctx.Err() + })) + + client, err := river.NewClient(driver, config) + require.NoError(t, err) + + // An independent insert-only client cannot use the worker client's local + // cancellation shortcut, exercising the same path as another process. + controller, err := river.NewClient(driver, &river.Config{Schema: schema}) + require.NoError(t, err) + + // Insert before starting so initial fetching finds the job even with a long + // fetch interval. Once running, the sole worker slot is occupied. + insertRes, err := controller.Insert(ctx, &cancelRunningJobArgs{}, nil) + require.NoError(t, err) + + events := subscribe(t, client) + startClient(ctx, t, client) + t.Cleanup(func() { require.NoError(t, client.StopAndCancel(ctx)) }) + require.Equal(t, insertRes.Job.ID, jobStarted.WaitOrTimeout()) + + if transactional { + // Rolling back must leave both the durable cancellation marker and the + // worker untouched. + execTx, err := driver.GetExecutor().Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = execTx.Rollback(ctx) }) + _, err = controller.JobCancelTx(ctx, driver.UnwrapTx(execTx), insertRes.Job.ID) + require.NoError(t, err) + require.NoError(t, execTx.Rollback(ctx)) + + jobAfterRollback, err := controller.JobGet(ctx, insertRes.Job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateRunning, jobAfterRollback.State) + require.NotContains(t, string(jobAfterRollback.Metadata), `"cancel_attempted_at"`) + jobContextCancelled.RequireEmpty() + + execTx, err = driver.GetExecutor().Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = execTx.Rollback(ctx) }) + updatedJob, err := controller.JobCancelTx(ctx, driver.UnwrapTx(execTx), insertRes.Job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateRunning, updatedJob.State) + + // Keep the transaction open across a cancellation poll. Neither polling + // nor notifications may cancel the worker before commit. + select { + case <-jobContextCancelled.WaitC(): + t.Fatal("worker cancelled before transaction committed") + case <-time.After(2200 * time.Millisecond): + } + + require.NoError(t, execTx.Commit(ctx)) + } else { + updatedJob, err := controller.JobCancel(ctx, insertRes.Job.ID) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateRunning, updatedJob.State) + } + + require.Equal(t, insertRes.Job.ID, jobContextCancelled.WaitOrTimeout()) + event := riversharedtest.WaitOrTimeout(t, events) + require.Equal(t, river.EventKindJobCancelled, event.Kind) + require.Equal(t, insertRes.Job.ID, event.Job.ID) + require.Equal(t, rivertype.JobStateCancelled, event.Job.State) +} diff --git a/riverdriver/riverdrivertest/client_maintenance_test.go b/riverdriver/riverdrivertest/client_maintenance_test.go new file mode 100644 index 000000000..77ef4e00a --- /dev/null +++ b/riverdriver/riverdrivertest/client_maintenance_test.go @@ -0,0 +1,78 @@ +package riverdrivertest + +import ( + "context" + "errors" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/riverqueue/river" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/testsignal" + "github.com/riverqueue/river/rivertype" +) + +func exerciseClientMaintenanceStartRecovery[TTx any](ctx context.Context, t *testing.T, driver riverdriver.Driver[TTx], schema string, pollOnly bool) { + t.Helper() + + var startAttempts atomic.Int32 + var attemptLeaders testsignal.TestSignal[*riverdriver.Leader] + attemptLeaders.Init(t) + + config := newTestConfig(t, schema) + config.PollOnly = pollOnly + config.Hooks = []rivertype.Hook{ + river.HookPeriodicJobsStartFunc(func(ctx context.Context, _ *rivertype.HookPeriodicJobsStartParams) error { + leader, err := driver.GetExecutor().LeaderGetElectedLeader(ctx, &riverdriver.LeaderGetElectedLeaderParams{Schema: schema}) + if err != nil { + return err + } + + attempt := startAttempts.Add(1) + attemptLeaders.Signal(leader) + if attempt <= 3 { + return errors.New("maintenance start error") + } + return nil + }), + } + config.PeriodicJobs = []*river.PeriodicJob{ + river.NewPeriodicJob(river.PeriodicInterval(time.Hour), func() (river.JobArgs, *river.InsertOpts) { + return noOpArgs{}, nil + }, &river.PeriodicJobOpts{RunOnStart: true}), + } + + client, err := river.NewClient(driver, config) + require.NoError(t, err) + + events := subscribe(t, client) + startClient(ctx, t, client) + + first := attemptLeaders.WaitOrTimeout() + require.NotNil(t, first) + require.Equal(t, client.ID(), first.LeaderID) + for range 2 { + attemptLeader := attemptLeaders.WaitOrTimeout() + require.NotNil(t, attemptLeader) + require.Equal(t, first.LeaderID, attemptLeader.LeaderID) + require.Equal(t, first.ElectedAt, attemptLeader.ElectedAt) + } + + // Exhausting startup retries must end the term. With only one client, it + // wins again and retries maintenance in a fresh term without a notification. + recovered := attemptLeaders.WaitOrTimeout() + require.NotNil(t, recovered) + require.Equal(t, client.ID(), recovered.LeaderID) + require.True(t, recovered.ElectedAt.After(first.ElectedAt)) + + // Prove maintenance actually recovered, rather than just receiving a + // resignation request: its run-on-start periodic job must be worked. + event := riversharedtest.WaitOrTimeout(t, events) + require.Equal(t, river.EventKindJobCompleted, event.Kind) + require.Equal(t, (noOpArgs{}).Kind(), event.Job.Kind) + require.Equal(t, int32(4), startAttempts.Load()) +} diff --git a/riverdriver/riverdrivertest/driver_client_test.go b/riverdriver/riverdrivertest/driver_client_test.go index f678935a8..43c0a7b63 100644 --- a/riverdriver/riverdrivertest/driver_client_test.go +++ b/riverdriver/riverdrivertest/driver_client_test.go @@ -389,112 +389,29 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, require.Equal(t, insertRes.Job.Kind, event.Job.Kind) }) - t.Run("CancelRunningJobWithListener", func(t *testing.T) { - t.Parallel() - - config, bundle := setupConfig(t) - if !bundle.driver.SupportsListener() { - t.Skip("requires a listener") - } - config.FetchPollInterval = time.Minute - - client, err := river.NewClient(bundle.driver, config) - require.NoError(t, err) - - var jobStarted, jobContextCancelled testsignal.TestSignal[int64] - jobStarted.Init(t) - jobContextCancelled.Init(t) - - type JobArgs struct { - testutil.JobArgsReflectKind[JobArgs] - } - - river.AddWorker(bundle.config.Workers, river.WorkFunc(func(ctx context.Context, job *river.Job[JobArgs]) error { - jobStarted.Signal(job.ID) - <-ctx.Done() - jobContextCancelled.Signal(job.ID) - return ctx.Err() - })) - - subscribeChan := subscribe(t, client) - startClient(ctx, t, client) - - insertRes, err := client.Insert(ctx, &JobArgs{}, nil) - require.NoError(t, err) - - // Cancel only after work starts so the listener must reach an active - // worker, rather than just changing a queued job's state. - require.Equal(t, insertRes.Job.ID, jobStarted.WaitOrTimeout()) - - updatedJob, err := client.JobCancel(ctx, insertRes.Job.ID) - require.NoError(t, err) - require.Equal(t, rivertype.JobStateRunning, updatedJob.State) - require.Equal(t, insertRes.Job.ID, jobContextCancelled.WaitOrTimeout()) - - event := riversharedtest.WaitOrTimeout(t, subscribeChan) - require.Equal(t, river.EventKindJobCancelled, event.Kind) - require.Equal(t, rivertype.JobStateCancelled, event.Job.State) - }) - - t.Run("CancelRunningJobWithListenerTx", func(t *testing.T) { - t.Parallel() - - config, bundle := setupConfig(t) - if !bundle.driver.SupportsListener() { - t.Skip("requires a listener") - } - config.FetchPollInterval = time.Minute - - client, err := river.NewClient(bundle.driver, config) - require.NoError(t, err) - - var jobStarted, jobContextCancelled testsignal.TestSignal[int64] - jobStarted.Init(t) - jobContextCancelled.Init(t) - - type JobArgs struct { - testutil.JobArgsReflectKind[JobArgs] + for _, transactional := range []bool{false, true} { + name := "CancelRunningJob" + if transactional { + name += "Tx" } + t.Run(name, func(t *testing.T) { + t.Parallel() - river.AddWorker(bundle.config.Workers, river.WorkFunc(func(ctx context.Context, job *river.Job[JobArgs]) error { - jobStarted.Signal(job.ID) - <-ctx.Done() - jobContextCancelled.Signal(job.ID) - return ctx.Err() - })) - - events := subscribe(t, client) - startClient(ctx, t, client) - - insertRes, err := client.Insert(ctx, &JobArgs{}, nil) - require.NoError(t, err) - - // Wait until the worker is running so a state update alone cannot pass. - require.Equal(t, insertRes.Job.ID, jobStarted.WaitOrTimeout()) - - // A rolled-back cancellation must leave both the job and worker active. - tx, execTx := beginTx(ctx, t, bundle) - _, err = client.JobCancelTx(ctx, tx, insertRes.Job.ID) - require.NoError(t, err) - require.NoError(t, execTx.Rollback(ctx)) - - jobAfterRollback, err := client.JobGet(ctx, insertRes.Job.ID) - require.NoError(t, err) - require.Equal(t, rivertype.JobStateRunning, jobAfterRollback.State) - require.NotContains(t, string(jobAfterRollback.Metadata), `"cancel_attempted_at"`) - jobContextCancelled.RequireEmpty() + for _, pollOnly := range []bool{false, true} { + name := "Default" + if pollOnly { + name = "PollOnly" + } + t.Run(name, func(t *testing.T) { + t.Parallel() - // Committing the same request should deliver the control notification. - tx, execTx = beginTx(ctx, t, bundle) - updatedJob, err := client.JobCancelTx(ctx, tx, insertRes.Job.ID) - require.NoError(t, err) - require.Equal(t, rivertype.JobStateRunning, updatedJob.State) - require.NoError(t, execTx.Commit(ctx)) - require.Equal(t, insertRes.Job.ID, jobContextCancelled.WaitOrTimeout()) + _, bundle := setupConfig(t) - event := riversharedtest.WaitOrTimeout(t, events) - require.Equal(t, river.EventKindJobCancelled, event.Kind) - }) + exerciseClientCancelRunningJob(ctx, t, bundle.driver, bundle.schema, pollOnly, transactional) + }) + } + }) + } // Keys containing gjson/sjson path syntax (and the empty key) are distinct // keys when unique by all args, so args differing in their values aren't @@ -1498,6 +1415,26 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, } }) + t.Run("MaintenanceStartRecovery", func(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + name string + pollOnly bool + }{ + {name: "Default"}, + {name: "PollOnly", pollOnly: true}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + _, bundle := setupConfig(t) + + exerciseClientMaintenanceStartRecovery(ctx, t, bundle.driver, bundle.schema, testCase.pollOnly) + }) + } + }) + t.Run("QueueGet", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverdrivertest/executor_tx.go b/riverdriver/riverdrivertest/executor_tx.go index e43139320..df952daed 100644 --- a/riverdriver/riverdrivertest/executor_tx.go +++ b/riverdriver/riverdrivertest/executor_tx.go @@ -2,7 +2,9 @@ package riverdrivertest import ( "context" + "sync" "testing" + "time" "github.com/stretchr/testify/require" @@ -54,6 +56,56 @@ func exerciseExecutorTx[TTx any](ctx context.Context, t *testing.T, } }) + t.Run("CancelledBeginLeavesPoolUsable", func(t *testing.T) { + t.Parallel() + + driver, _ := driverWithSchema(ctx, t, nil) + exec := driver.GetExecutor() + + // Race cancellation against BEGIN. A driver may start the + // transaction but return an error if cancellation arrives just + // afterwards. Subsequent transactions must still work. + for range 100 { + beginCtx, cancel := context.WithCancel(ctx) + var cancelGroup sync.WaitGroup + cancelGroup.Go(cancel) + tx, err := exec.Begin(beginCtx) + cancelGroup.Wait() + if err == nil { + _ = tx.Rollback(ctx) + } + + tx, err = exec.Begin(ctx) + require.NoError(t, err) + require.NoError(t, tx.Commit(ctx)) + } + }) + + t.Run("CancelledSQLiteTransactionReleasesConnection", func(t *testing.T) { + t.Parallel() + + driver, _ := driverWithSchema(ctx, t, nil) + if driver.DatabaseName() != riverdriver.DatabaseNameSQLite { + t.Skip("SQLite pools use one connection and database/sql rolls back automatically on cancellation") + } + exec := driver.GetExecutor() + + beginCtx, cancel := context.WithCancel(ctx) + defer cancel() + tx, err := exec.Begin(beginCtx) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback(ctx) }) + cancel() + + // No explicit rollback: cancellation alone must return the only + // connection to the pool after database/sql rolls back. + nextCtx, nextCancel := context.WithTimeout(ctx, 5*time.Second) + defer nextCancel() + nextTx, err := exec.Begin(nextCtx) + require.NoError(t, err) + require.NoError(t, nextTx.Rollback(ctx)) + }) + t.Run("NestedTransactions", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverdrivertest/job_insert.go b/riverdriver/riverdrivertest/job_insert.go index 6b625a0e5..9bc8cf64a 100644 --- a/riverdriver/riverdrivertest/job_insert.go +++ b/riverdriver/riverdrivertest/job_insert.go @@ -88,7 +88,7 @@ func exerciseJobInsert[TTx any](ctx context.Context, t *testing.T, // SQLite needs to set a special metadata key to be able to // check for duplicates. Remove this for purposes of comparing // inserted metadata. - job.Metadata, err = sjson.DeleteBytes(job.Metadata, rivercommon.MetadataKeyUniqueNonce) + job.Metadata, err = sjson.DeleteBytes(job.Metadata, riverdriver.UniqueInsertMetadataKey) require.NoError(t, err) require.Equal(t, idStart+int64(i), job.ID) diff --git a/riverdriver/riverdrivertest/job_read.go b/riverdriver/riverdrivertest/job_read.go index f8fe44182..4934b94da 100644 --- a/riverdriver/riverdrivertest/job_read.go +++ b/riverdriver/riverdrivertest/job_read.go @@ -55,7 +55,7 @@ func exerciseJobRead[TTx any](ctx context.Context, t *testing.T, executorWithTx for _, state := range rivertype.JobStates() { require.Contains(t, countsByState, state) - switch state { //nolint:exhaustive + switch state { case rivertype.JobStateAvailable: require.Equal(t, 2, countsByState[state]) case rivertype.JobStateCancelled: @@ -64,8 +64,10 @@ func exerciseJobRead[TTx any](ctx context.Context, t *testing.T, executorWithTx require.Equal(t, 1, countsByState[state]) case rivertype.JobStateDiscarded: require.Equal(t, 1, countsByState[state]) - default: + case rivertype.JobStatePending, rivertype.JobStateRetryable, rivertype.JobStateRunning, rivertype.JobStateScheduled: require.Equal(t, 0, countsByState[state]) + default: + require.FailNow(t, "unknown job state", state) } } }) @@ -535,6 +537,58 @@ func exerciseJobRead[TTx any](ctx context.Context, t *testing.T, executorWithTx sliceutil.Map(jobs, func(j *rivertype.JobRow) int64 { return j.ID })) }) + t.Run("JobGetCancelRequested", func(t *testing.T) { + t.Parallel() + + t.Run("AlternateSchema", func(t *testing.T) { + t.Parallel() + + exec, _ := setup(ctx, t) + + _, err := exec.JobGetCancelRequested(ctx, &riverdriver.JobGetCancelRequestedParams{ + ID: []int64{1}, + Schema: "custom_schema", + }) + requireMissingRelation(t, err, "custom_schema", "river_job") + }) + + t.Run("FiltersRunningJobsAndRequestedIDs", func(t *testing.T) { + t.Parallel() + + exec, _ := setup(ctx, t) + + jobIDs := make([]int64, 0, len(rivertype.JobStates())+2) + var expectedIDs []int64 + for _, state := range rivertype.JobStates() { + job := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{ + Metadata: []byte(`{"cancel_attempted_at":"2026-09-28T00:00:00Z"}`), + State: new(state), + }) + jobIDs = append(jobIDs, job.ID) + if state == rivertype.JobStateRunning { + expectedIDs = append(expectedIDs, job.ID) + } + } + + // An active job without a cancellation and a cancelled active job + // outside the requested IDs must both be excluded. + job := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{State: new(rivertype.JobStateRunning)}) + jobIDs = append(jobIDs, job.ID, 0) + _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{ + Metadata: []byte(`{"cancel_attempted_at":"2026-09-28T00:00:00Z"}`), + State: new(rivertype.JobStateRunning), + }) + + ids, err := exec.JobGetCancelRequested(ctx, &riverdriver.JobGetCancelRequestedParams{ID: jobIDs}) + require.NoError(t, err) + require.Equal(t, expectedIDs, ids) + + ids, err = exec.JobGetCancelRequested(ctx, &riverdriver.JobGetCancelRequestedParams{}) + require.NoError(t, err) + require.Empty(t, ids) + }) + }) + t.Run("JobGetStuck", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverdrivertest/riverdrivertest.go b/riverdriver/riverdrivertest/riverdrivertest.go index aa66762c7..41c50b1a7 100644 --- a/riverdriver/riverdrivertest/riverdrivertest.go +++ b/riverdriver/riverdrivertest/riverdrivertest.go @@ -54,6 +54,26 @@ func exerciseDriverPool[TTx any](ctx context.Context, t *testing.T, ) { t.Helper() + t.Run("InitDriver", func(t *testing.T) { + t.Parallel() + + exec, _ := executorWithTx(ctx, t) + require.NoError(t, exec.InitDriver(ctx)) + require.NoError(t, exec.InitDriver(ctx)) + }) + + t.Run("Ping", func(t *testing.T) { + t.Parallel() + + exec, _ := executorWithTx(ctx, t) + require.NoError(t, exec.InitDriver(ctx)) + require.NoError(t, exec.Ping(ctx)) + + cancelledCtx, cancel := context.WithCancel(ctx) + cancel() + require.ErrorIs(t, exec.Ping(cancelledCtx), context.Canceled) + }) + t.Run("PoolIsSet", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverdrivertest/yugabyte_test.go b/riverdriver/riverdrivertest/yugabyte_test.go new file mode 100644 index 000000000..f70dab5c2 --- /dev/null +++ b/riverdriver/riverdrivertest/yugabyte_test.go @@ -0,0 +1,132 @@ +package riverdrivertest + +import ( + "context" + "database/sql" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/stdlib" + "github.com/stretchr/testify/require" + + "github.com/riverqueue/river/riverdbtest" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/riverdriver/riverdatabasesql" + "github.com/riverqueue/river/riverdriver/riverpgxv5" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/testfactory" + "github.com/riverqueue/river/rivertype" +) + +func TestDriverYugabyteNotifications(t *testing.T) { + t.Parallel() + + ctx := context.Background() + for _, testCase := range []struct { + enabled *bool + name string + }{ + {enabled: new(false), name: "Disabled"}, + {enabled: new(true), name: "Enabled"}, + {name: "Unavailable"}, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + for _, driverName := range []string{"DatabaseSQL", "DatabaseSQLWithListener", "Pgx"} { + t.Run(driverName, func(t *testing.T) { + t.Parallel() + + basePool := riversharedtest.DBPool(ctx, t) + schema := riverdbtest.TestSchema(ctx, t, riverpgxv5.New(basePool), nil) + pool := riversharedtest.DBPoolWithYugabyteVersion(ctx, t, schema, testCase.enabled) + enabled := testCase.enabled != nil && *testCase.enabled + if driverName == "Pgx" { + exerciseYugabyteNotifications(ctx, t, schema, true, enabled, func() riverdriver.Driver[pgx.Tx] { + return riverpgxv5.New(pool) + }) + } else { + sqlPool := stdlib.OpenDBFromPool(pool) + t.Cleanup(func() { require.NoError(t, sqlPool.Close()) }) + withListener := driverName == "DatabaseSQLWithListener" + exerciseYugabyteNotifications(ctx, t, schema, withListener, enabled, func() riverdriver.Driver[*sql.Tx] { + if withListener { + return riverdatabasesql.NewWithPgxListener(sqlPool, pool) + } + return riverdatabasesql.New(sqlPool) + }) + } + }) + } + }) + } +} + +func exerciseYugabyteNotifications[TTx any](ctx context.Context, t *testing.T, schema string, withListener, enabled bool, newDriverFunc func() riverdriver.Driver[TTx]) { + t.Helper() + + driver := newDriverFunc() + require.Equal(t, withListener, driver.SupportsListener()) + require.True(t, driver.SupportsListenNotify()) + + cancelledCtx, cancel := context.WithCancel(ctx) + cancel() + require.ErrorIs(t, driver.GetExecutor().InitDriver(cancelledCtx), context.Canceled) + require.NoError(t, driver.GetExecutor().InitDriver(ctx)) + // Successful detection is cached, even when the next context is cancelled. + require.NoError(t, driver.GetExecutor().InitDriver(cancelledCtx)) + require.Equal(t, withListener && enabled, driver.SupportsListener()) + require.Equal(t, enabled, driver.SupportsListenNotify()) + + // These executors haven't been initialized explicitly, as with insert-only + // clients. Each must detect support before attempting a notification. + notifyDriver := newDriverFunc() + require.NoError(t, notifyDriver.GetExecutor().NotifyMany(ctx, &riverdriver.NotifyManyParams{ + Payload: []string{`{"action":"pause","queue":"default"}`}, + Schema: schema, + Topic: "river_control", + })) + require.Equal(t, enabled, notifyDriver.SupportsListenNotify()) + + cancelExec := newDriverFunc().GetExecutor() + job := testfactory.Job(ctx, t, cancelExec, &testfactory.JobOpts{Schema: schema}) + cancelledJob, err := cancelExec.JobCancel(ctx, &riverdriver.JobCancelParams{ + ID: job.ID, + CancelAttemptedAt: time.Now(), + ControlTopic: "river_control", + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCancelled, cancelledJob.State) + + leaderExec := newDriverFunc().GetExecutor() + leader, err := leaderExec.LeaderInsert(ctx, &riverdriver.LeaderInsertParams{ + LeaderID: "yugabyte-test", + Schema: schema, + TTL: time.Minute, + }) + require.NoError(t, err) + resigned, err := leaderExec.LeaderResign(ctx, &riverdriver.LeaderResignParams{ + ElectedAt: leader.ElectedAt, + LeaderID: leader.LeaderID, + LeadershipTopic: "river_leadership", + Schema: schema, + }) + require.NoError(t, err) + require.True(t, resigned) + _, err = leaderExec.LeaderGetElectedLeader(ctx, &riverdriver.LeaderGetElectedLeaderParams{Schema: schema}) + require.ErrorIs(t, err, rivertype.ErrNotFound) + + // Default PollOnly=false must still cancel a running job from another + // client when startup detection disables LISTEN/NOTIFY. + t.Run("CancelRunningJob", func(t *testing.T) { + // Sequential: these clients share a schema and must stop before the + // maintenance recovery test can elect its own leader. + exerciseClientCancelRunningJob(ctx, t, newDriverFunc(), schema, false, true) + }) + t.Run("MaintenanceStartRecovery", func(t *testing.T) { + // Sequential because the cancellation client must release leadership. + exerciseClientMaintenanceStartRecovery(ctx, t, newDriverFunc(), schema, false) + }) +} diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql b/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql index 19a7b99f6..2f6445b4d 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql +++ b/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql @@ -1,6 +1,12 @@ -- name: PGAdvisoryXactLock :exec SELECT pg_advisory_xact_lock(@key); +-- name: PGGetProductAndVersion :one +SELECT + version()::text AS product, + current_setting('server_version_num')::int AS version_num, + coalesce(current_setting('yb_enable_listen_notify', true), 'off')::boolean AS yb_listen_notify_enabled; + -- name: PGNotifyMany :exec WITH topic_to_notify AS ( SELECT diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql.go b/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql.go index 9215c089b..fc299af20 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql.go +++ b/riverdriver/riverpgxv5/internal/dbsqlc/pg_misc.sql.go @@ -20,6 +20,26 @@ func (q *Queries) PGAdvisoryXactLock(ctx context.Context, db DBTX, key int64) er return err } +const pGGetProductAndVersion = `-- name: PGGetProductAndVersion :one +SELECT + version()::text AS product, + current_setting('server_version_num')::int AS version_num, + coalesce(current_setting('yb_enable_listen_notify', true), 'off')::boolean AS yb_listen_notify_enabled +` + +type PGGetProductAndVersionRow struct { + Product string + VersionNum int32 + YbListenNotifyEnabled bool +} + +func (q *Queries) PGGetProductAndVersion(ctx context.Context, db DBTX) (*PGGetProductAndVersionRow, error) { + row := db.QueryRow(ctx, pGGetProductAndVersion) + var i PGGetProductAndVersionRow + err := row.Scan(&i.Product, &i.VersionNum, &i.YbListenNotifyEnabled) + return &i, err +} + const pGNotifyMany = `-- name: PGNotifyMany :exec WITH topic_to_notify AS ( SELECT diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql index 26e023fa2..a1d4b5592 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql @@ -48,10 +48,10 @@ WITH locked_job AS ( notification AS ( SELECT id, - pg_notify( + CASE WHEN @notify::boolean THEN pg_notify( concat(coalesce(sqlc.narg('schema')::text, current_schema()), '.', @control_topic::text), json_build_object('action', 'cancel', 'job_id', id, 'queue', queue)::text - ) + ) END FROM locked_job WHERE @@ -254,6 +254,14 @@ FROM /* TEMPLATE: schema */river_job WHERE kind = any(@kind::text[]) ORDER BY id; +-- name: JobGetCancelRequested :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE id = any(@id::bigint[]) + AND metadata ? 'cancel_attempted_at' + AND state = 'running' +ORDER BY id; + -- name: JobGetStuck :many SELECT * FROM /* TEMPLATE: schema */river_job @@ -318,7 +326,9 @@ ON CONFLICT (unique_key) AND /* TEMPLATE: schema */river_job_state_in_bitmask(unique_states, state) -- Something needs to be updated for a row to be returned on a conflict. DO UPDATE SET kind = EXCLUDED.kind -RETURNING sqlc.embed(river_job), (xmax != 0) AS unique_skipped_as_duplicate; +RETURNING + sqlc.embed(river_job), + /* TEMPLATE_BEGIN: unique_skipped_as_duplicate */ (xmax != 0) /* TEMPLATE_END */ AS unique_skipped_as_duplicate; -- name: JobInsertFastManyNoReturning :execrows INSERT INTO /* TEMPLATE: schema */river_job( diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go index 62b07c2ba..af7a92a5e 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go @@ -24,10 +24,10 @@ WITH locked_job AS ( notification AS ( SELECT id, - pg_notify( - concat(coalesce($2::text, current_schema()), '.', $3::text), + CASE WHEN $2::boolean THEN pg_notify( + concat(coalesce($3::text, current_schema()), '.', $4::text), json_build_object('action', 'cancel', 'job_id', id, 'queue', queue)::text - ) + ) END FROM locked_job WHERE @@ -40,10 +40,10 @@ updated_job AS ( -- If the job is actively running, we want to let its current client and -- producer handle the cancellation. Otherwise, immediately cancel it. state = CASE WHEN state = 'running' THEN state ELSE 'cancelled' END, - finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE coalesce($4::timestamptz, now()) END, + finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE coalesce($5::timestamptz, now()) END, -- Mark the job as cancelled by query so that the rescuer knows not to -- rescue it, even if it gets stuck in the running state: - metadata = jsonb_set(metadata, '{cancel_attempted_at}'::text[], $5::jsonb, true) + metadata = jsonb_set(metadata, '{cancel_attempted_at}'::text[], $6::jsonb, true) FROM notification WHERE river_job.id = notification.id RETURNING river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states @@ -59,6 +59,7 @@ FROM updated_job type JobCancelParams struct { ID int64 + Notify bool Schema pgtype.Text ControlTopic string Now *time.Time @@ -68,6 +69,7 @@ type JobCancelParams struct { func (q *Queries) JobCancel(ctx context.Context, db DBTX, arg *JobCancelParams) (*RiverJob, error) { row := db.QueryRow(ctx, jobCancel, arg.ID, + arg.Notify, arg.Schema, arg.ControlTopic, arg.Now, @@ -587,6 +589,35 @@ func (q *Queries) JobGetByKindMany(ctx context.Context, db DBTX, kind []string) return items, nil } +const jobGetCancelRequested = `-- name: JobGetCancelRequested :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE id = any($1::bigint[]) + AND metadata ? 'cancel_attempted_at' + AND state = 'running' +ORDER BY id +` + +func (q *Queries) JobGetCancelRequested(ctx context.Context, db DBTX, id []int64) ([]int64, error) { + rows, err := db.Query(ctx, jobGetCancelRequested, id) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const jobGetStuck = `-- name: JobGetStuck :many SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states FROM /* TEMPLATE: schema */river_job @@ -697,7 +728,9 @@ ON CONFLICT (unique_key) AND /* TEMPLATE: schema */river_job_state_in_bitmask(unique_states, state) -- Something needs to be updated for a row to be returned on a conflict. DO UPDATE SET kind = EXCLUDED.kind -RETURNING river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states, (xmax != 0) AS unique_skipped_as_duplicate +RETURNING + river_job.id, river_job.args, river_job.attempt, river_job.attempted_at, river_job.attempted_by, river_job.created_at, river_job.errors, river_job.finalized_at, river_job.kind, river_job.max_attempts, river_job.metadata, river_job.priority, river_job.queue, river_job.state, river_job.scheduled_at, river_job.tags, river_job.unique_key, river_job.unique_states, + /* TEMPLATE_BEGIN: unique_skipped_as_duplicate */ (xmax != 0) /* TEMPLATE_END */ AS unique_skipped_as_duplicate ` type JobInsertFastManyParams struct { diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql b/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql index cea2195f1..b883042da 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql @@ -60,10 +60,10 @@ WITH currently_held_leaders AS ( FOR UPDATE ), notified_resignations AS ( - SELECT pg_notify( + SELECT CASE WHEN @notify::boolean THEN pg_notify( concat(coalesce(sqlc.narg('schema')::text, current_schema()), '.', @leadership_topic::text), json_build_object('leader_id', leader_id, 'action', 'resigned')::text - ) + ) END FROM currently_held_leaders ) DELETE FROM /* TEMPLATE: schema */river_leader USING notified_resignations; diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql.go b/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql.go index 1976cf4d7..9abe69cdd 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql.go +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_leader.sql.go @@ -158,10 +158,10 @@ WITH currently_held_leaders AS ( FOR UPDATE ), notified_resignations AS ( - SELECT pg_notify( - concat(coalesce($3::text, current_schema()), '.', $4::text), + SELECT CASE WHEN $3::boolean THEN pg_notify( + concat(coalesce($4::text, current_schema()), '.', $5::text), json_build_object('leader_id', leader_id, 'action', 'resigned')::text - ) + ) END FROM currently_held_leaders ) DELETE FROM /* TEMPLATE: schema */river_leader USING notified_resignations @@ -170,6 +170,7 @@ DELETE FROM /* TEMPLATE: schema */river_leader USING notified_resignations type LeaderResignParams struct { ElectedAt time.Time LeaderID string + Notify bool Schema pgtype.Text LeadershipTopic string } @@ -178,6 +179,7 @@ func (q *Queries) LeaderResign(ctx context.Context, db DBTX, arg *LeaderResignPa result, err := db.Exec(ctx, leaderResign, arg.ElectedAt, arg.LeaderID, + arg.Notify, arg.Schema, arg.LeadershipTopic, ) diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver.go b/riverdriver/riverpgxv5/river_pgx_v5_driver.go index 5167ff6c6..62a29a493 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver.go @@ -16,6 +16,7 @@ import ( "math" "strings" "sync" + "sync/atomic" "time" "github.com/jackc/pgx/v5" @@ -30,6 +31,7 @@ import ( "github.com/riverqueue/river/rivershared/uniquestates" "github.com/riverqueue/river/rivershared/util/dbutil" "github.com/riverqueue/river/rivershared/util/ptrutil" + "github.com/riverqueue/river/rivershared/util/randutil" "github.com/riverqueue/river/rivershared/util/sliceutil" "github.com/riverqueue/river/rivertype" ) @@ -39,8 +41,9 @@ var migrationFS embed.FS // Driver is an implementation of riverdriver.Driver for Pgx v5. type Driver struct { - dbPool *pgxpool.Pool - replacer sqlctemplate.Replacer + dbPool *pgxpool.Pool + postgresCapabilities atomic.Pointer[riverdriver.PostgresCapabilities] + replacer sqlctemplate.Replacer } // New returns a new Pgx v5 River driver for use with River. @@ -105,8 +108,11 @@ func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, er return fmt.Sprintf("%s = any(@%s)", column, column), values, nil } -func (d *Driver) SupportsListener() bool { return true } -func (d *Driver) SupportsListenNotify() bool { return true } +func (d *Driver) SupportsListener() bool { return d.SupportsListenNotify() } +func (d *Driver) SupportsListenNotify() bool { + capabilities := d.postgresCapabilities.Load() + return capabilities == nil || capabilities.SupportsListenNotify +} func (d *Driver) TimePrecision() time.Duration { return time.Microsecond } func (d *Driver) UnwrapExecutor(tx pgx.Tx) riverdriver.ExecutorTx { @@ -216,7 +222,17 @@ func (e *Executor) IndexesExist(ctx context.Context, params *riverdriver.Indexes return exists, nil } +func (e *Executor) InitDriver(ctx context.Context) error { + _, err := e.getPostgresCapabilities(ctx) + return err +} + func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelParams) (*rivertype.JobRow, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return nil, err + } + cancelledAt, err := params.CancelAttemptedAt.MarshalJSON() if err != nil { return nil, err @@ -226,6 +242,7 @@ func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelP ID: params.ID, CancelAttemptedAt: cancelledAt, ControlTopic: params.ControlTopic, + Notify: capabilities.SupportsListenNotify, Now: params.Now, Schema: pgtype.Text{String: params.Schema, Valid: params.Schema != ""}, }) @@ -354,6 +371,11 @@ func (e *Executor) JobGetByKindMany(ctx context.Context, params *riverdriver.Job return sliceutil.MapError(jobs, jobRowFromInternal) } +func (e *Executor) JobGetCancelRequested(ctx context.Context, params *riverdriver.JobGetCancelRequestedParams) ([]int64, error) { + ids, err := dbsqlc.New().JobGetCancelRequested(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.ID) + return ids, interpretError(err) +} + func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetStuckParams) ([]*rivertype.JobRow, error) { jobs, err := dbsqlc.New().JobGetStuck(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobGetStuckParams{ AfterID: params.AfterID, @@ -367,6 +389,17 @@ func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetSt } func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.JobInsertFastManyParams) ([]*riverdriver.JobInsertFastResult, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return nil, err + } + uniqueInsertMode := capabilities.UniqueInsertMode + + var uniqueNonce string + if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce { + uniqueNonce = randutil.Hex(8) + } + insertJobsParams := &dbsqlc.JobInsertFastManyParams{ ID: make([]int64, len(params.Jobs)), Args: make([][]byte, len(params.Jobs)), @@ -408,7 +441,16 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo insertJobsParams.CreatedAt[i] = createdAt insertJobsParams.Kind[i] = params.Kind insertJobsParams.MaxAttempts[i] = int16(min(params.MaxAttempts, math.MaxInt16)) //nolint:gosec - insertJobsParams.Metadata[i] = sliceutil.FirstNonEmpty(params.Metadata, defaultObject) + metadata := sliceutil.FirstNonEmpty(params.Metadata, defaultObject) + if uniqueNonce != "" { + var err error + metadata, err = riverdriver.UniqueInsertMetadataWithNonce(metadata, uniqueNonce) + if err != nil { + return nil, err + } + } + + insertJobsParams.Metadata[i] = metadata insertJobsParams.Priority[i] = int16(min(params.Priority, math.MaxInt16)) //nolint:gosec insertJobsParams.Queue[i] = params.Queue insertJobsParams.ScheduledAt[i] = scheduledAt @@ -418,6 +460,9 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo insertJobsParams.UniqueStates[i] = int32(params.UniqueStates) } + ctx = sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "unique_skipped_as_duplicate": {Value: uniqueInsertMode.SQL(), Stable: true}, + }, nil) items, err := dbsqlc.New().JobInsertFastMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, insertJobsParams) if err != nil { return nil, interpretError(err) @@ -428,7 +473,13 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo if err != nil { return nil, err } - return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: row.UniqueSkippedAsDuplicate}, nil + + uniqueSkippedAsDuplicate := row.UniqueSkippedAsDuplicate + if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce { + uniqueSkippedAsDuplicate = riverdriver.UniqueInsertMetadataIsDuplicate(job.Metadata, uniqueNonce) + } + + return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: uniqueSkippedAsDuplicate}, nil }) } @@ -779,10 +830,16 @@ func (e *Executor) LeaderInsert(ctx context.Context, params *riverdriver.LeaderI } func (e *Executor) LeaderResign(ctx context.Context, params *riverdriver.LeaderResignParams) (bool, error) { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return false, err + } + numResigned, err := dbsqlc.New().LeaderResign(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.LeaderResignParams{ ElectedAt: params.ElectedAt, LeaderID: params.LeaderID, LeadershipTopic: params.LeadershipTopic, + Notify: capabilities.SupportsListenNotify, Schema: pgtype.Text{String: params.Schema, Valid: params.Schema != ""}, }) if err != nil { @@ -875,6 +932,14 @@ func (e *Executor) NotificationDeleteBefore(ctx context.Context, params *riverdr } func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyManyParams) error { + capabilities, err := e.getPostgresCapabilities(ctx) + if err != nil { + return err + } + if !capabilities.SupportsListenNotify { + return nil + } + return dbsqlc.New().PGNotifyMany(ctx, e.dbtx, &dbsqlc.PGNotifyManyParams{ Payload: params.Payload, Schema: pgtype.Text{String: params.Schema, Valid: params.Schema != ""}, @@ -882,6 +947,10 @@ func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyMan }) } +func (e *Executor) Ping(ctx context.Context) error { + return e.Exec(ctx, "SELECT 1") +} + func (e *Executor) PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) { err := dbsqlc.New().PGAdvisoryXactLock(ctx, e.dbtx, key) return &struct{}{}, interpretError(err) @@ -1040,6 +1109,29 @@ func (e *Executor) TableTruncate(ctx context.Context, params *riverdriver.TableT return interpretError(err) } +func (e *Executor) getPostgresCapabilities(ctx context.Context) (*riverdriver.PostgresCapabilities, error) { + if e.driver != nil { + if capabilities := e.driver.postgresCapabilities.Load(); capabilities != nil { + return capabilities, nil + } + } + + productAndVersion, err := dbsqlc.New().PGGetProductAndVersion(ctx, e.dbtx) + if err != nil { + return nil, interpretError(err) + } + + capabilities := riverdriver.NewPostgresCapabilities(productAndVersion.Product, productAndVersion.VersionNum, productAndVersion.YbListenNotifyEnabled) + if e.driver != nil { + // Concurrent callers may both detect, but the first successful result + // becomes the driver's cached capabilities. Don't hold a lock while + // querying: a caller may already hold the pool's only connection. + e.driver.postgresCapabilities.CompareAndSwap(nil, capabilities) + capabilities = e.driver.postgresCapabilities.Load() + } + return capabilities, nil +} + type ExecutorTx struct { Executor diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver_test.go b/riverdriver/riverpgxv5/river_pgx_v5_driver_test.go index ae366d5a7..34d20da5a 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver_test.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver_test.go @@ -18,12 +18,24 @@ import ( "github.com/riverqueue/river/riverdriver" "github.com/riverqueue/river/rivershared/sqlctemplate" + "github.com/riverqueue/river/rivershared/testsignal" "github.com/riverqueue/river/rivertype" ) // Verify interface compliance. var _ riverdriver.Driver[pgx.Tx] = New(nil) +type executorInitDriverTestDBTX struct { + *pgxpool.Pool + + QueryRowStarted testsignal.TestSignal[struct{}] +} + +func (d *executorInitDriverTestDBTX) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row { + d.QueryRowStarted.Signal(struct{}{}) + return d.Pool.QueryRow(ctx, sql, args...) +} + func TestNew(t *testing.T) { t.Parallel() @@ -43,6 +55,49 @@ func TestNew(t *testing.T) { }) } +func TestExecutor_InitDriverDoesNotBlockTransaction(t *testing.T) { + t.Parallel() + + ctx := context.Background() + config := testPoolConfig() + config.MaxConns = 1 + dbPool := testPool(ctx, t, config) + driver := New(dbPool) + + tx, err := dbPool.Begin(ctx) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback(ctx) }) + + poolDBTX := &executorInitDriverTestDBTX{Pool: dbPool} + poolDBTX.QueryRowStarted.Init(t) + poolExecutor := &Executor{ + dbtx: templateReplaceWrapper{dbtx: poolDBTX, replacer: &driver.replacer}, + driver: driver, + } + + initCtx, initCancel := context.WithTimeout(ctx, 10*time.Second) + t.Cleanup(initCancel) + + var poolInitFinished testsignal.TestSignal[error] + poolInitFinished.Init(t) + go func() { poolInitFinished.Signal(poolExecutor.InitDriver(initCtx)) }() + poolDBTX.QueryRowStarted.WaitOrTimeout() + + var txInitFinished testsignal.TestSignal[error] + txInitFinished.Init(t) + go func() { txInitFinished.Signal(driver.UnwrapExecutor(tx).InitDriver(ctx)) }() + + select { + case err := <-txInitFinished.WaitC(): + require.NoError(t, err) + case <-time.After(2 * time.Second): + require.FailNow(t, "transactional driver initialization blocked behind pool initialization") + } + + require.NoError(t, tx.Rollback(ctx)) + require.NoError(t, poolInitFinished.WaitOrTimeout()) +} + func TestListener_Close(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverpgxv5/yugabyte_compatibility_test.go b/riverdriver/riverpgxv5/yugabyte_compatibility_test.go new file mode 100644 index 000000000..57aafa691 --- /dev/null +++ b/riverdriver/riverpgxv5/yugabyte_compatibility_test.go @@ -0,0 +1,57 @@ +package riverpgxv5 + +import ( + "fmt" + "io/fs" + "os" + "regexp" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestYugabyteCompatibility(t *testing.T) { + t.Parallel() + + // YugabyteDB doesn't expose PostgreSQL's transaction-related system columns: + // https://docs.yugabyte.com/stable/yugabyte-voyager/known-issues/postgresql/#system-columns-is-not-yet-supported + // The unique insert query may use xmax only because the entire expression is + // replaced when the driver detects YugabyteDB. + var ( + uniqueInsertModeTemplateRE = regexp.MustCompile(`(?s)/\*\s*TEMPLATE_BEGIN: unique_skipped_as_duplicate\s*\*/.*?/\*\s*TEMPLATE_END\s*\*/`) + unsupportedSystemColumnRE = regexp.MustCompile(`(?i)\b(?:cmax|cmin|ctid|xmax|xmin)\b`) + ) + + sourceRoot, err := os.OpenRoot(".") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, sourceRoot.Close()) }) + + var violations []string + err = fs.WalkDir(sourceRoot.FS(), ".", func(path string, entry fs.DirEntry, err error) error { + if err != nil { + return err + } + if entry.IsDir() || (!strings.HasSuffix(path, ".go") && !strings.HasSuffix(path, ".sql")) || strings.HasSuffix(path, "_test.go") { + return nil + } + + contents, err := sourceRoot.ReadFile(path) + if err != nil { + return err + } + contents = uniqueInsertModeTemplateRE.ReplaceAll(contents, nil) + + for lineNum, line := range strings.Split(string(contents), "\n") { + for _, column := range unsupportedSystemColumnRE.FindAllString(line, -1) { + violations = append(violations, fmt.Sprintf("%s:%d: %s", path, lineNum+1, column)) + } + } + return nil + }) + require.NoError(t, err) + + require.Empty(t, violations, + "YugabyteDB-incompatible PostgreSQL system columns must only appear inside SQL templates that replace them for YugabyteDB", + ) +} diff --git a/riverdriver/riversqlite/go.mod b/riverdriver/riversqlite/go.mod index fd1a95df8..2d71bdaa0 100644 --- a/riverdriver/riversqlite/go.mod +++ b/riverdriver/riversqlite/go.mod @@ -5,19 +5,14 @@ go 1.26.0 toolchain go1.26.6 require ( - github.com/riverqueue/river v0.47.0 github.com/riverqueue/river/riverdriver v0.47.0 github.com/riverqueue/river/rivershared v0.47.0 github.com/riverqueue/river/rivertype v0.47.0 github.com/stretchr/testify v1.12.1 - github.com/tidwall/gjson v1.19.0 - github.com/tidwall/sjson v1.2.5 ) require ( github.com/jackc/pgx/v5 v5.11.0 // indirect - github.com/tidwall/match v1.2.0 // indirect - github.com/tidwall/pretty v1.2.1 // indirect go.yaml.in/yaml/v3 v3.0.5 // indirect golang.org/x/sync v0.23.0 // indirect golang.org/x/text v0.42.0 // indirect diff --git a/riverdriver/riversqlite/go.sum b/riverdriver/riversqlite/go.sum index e8af2804f..c47995de6 100644 --- a/riverdriver/riversqlite/go.sum +++ b/riverdriver/riversqlite/go.sum @@ -18,17 +18,6 @@ github.com/riverqueue/river/rivertype v0.47.0 h1:SzNavtLGR4nMT1QkrEYQ7n96OMatYsn github.com/riverqueue/river/rivertype v0.47.0/go.mod h1:XKkcRQR6zm8RR/JQa1Q2ywpj8uXQu21quPa4Lpw1Xhw= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= -github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU= -github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc= -github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= -github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM= -github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= -github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= -github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= -github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= -github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= -github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql index 1cf46567b..00e2a7c8d 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql @@ -182,6 +182,14 @@ FROM /* TEMPLATE: schema */river_job WHERE kind IN (sqlc.slice('kind')) ORDER BY id; +-- name: JobGetCancelRequested :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE id IN (sqlc.slice('id')) + AND (metadata -> 'cancel_attempted_at') IS NOT NULL + AND state = 'running' +ORDER BY id; + -- name: JobGetStuck :many SELECT * FROM /* TEMPLATE: schema */river_job diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go index 4848272ac..9f255817c 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go @@ -563,6 +563,48 @@ func (q *Queries) JobGetByKindMany(ctx context.Context, db DBTX, kind []string) return items, nil } +const jobGetCancelRequested = `-- name: JobGetCancelRequested :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE id IN (/*SLICE:id*/?) + AND (metadata -> 'cancel_attempted_at') IS NOT NULL + AND state = 'running' +ORDER BY id +` + +func (q *Queries) JobGetCancelRequested(ctx context.Context, db DBTX, id []int64) ([]int64, error) { + query := jobGetCancelRequested + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const jobGetStuck = `-- name: JobGetStuck :many SELECT id, json(args), attempt, attempted_at, json(attempted_by), created_at, json(errors), finalized_at, kind, max_attempts, json(metadata), priority, queue, state, scheduled_at, json(tags), unique_key, unique_states FROM /* TEMPLATE: schema */river_job diff --git a/riverdriver/riversqlite/river_sqlite_driver.go b/riverdriver/riversqlite/river_sqlite_driver.go index b9bbb00f0..112f686a0 100644 --- a/riverdriver/riversqlite/river_sqlite_driver.go +++ b/riverdriver/riversqlite/river_sqlite_driver.go @@ -20,6 +20,7 @@ package riversqlite import ( "context" "database/sql" + "database/sql/driver" "embed" "encoding/hex" "encoding/json" @@ -38,10 +39,6 @@ import ( "sync" "time" - "github.com/tidwall/gjson" - "github.com/tidwall/sjson" - - "github.com/riverqueue/river/internal/rivercommon" "github.com/riverqueue/river/riverdriver" "github.com/riverqueue/river/riverdriver/riversqlite/internal/dbsqlc" "github.com/riverqueue/river/rivershared/sqlctemplate" @@ -199,12 +196,32 @@ func (e *Executor) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { return e.execTx.Begin(ctx) } - tx, err := e.dbPool.BeginTx(ctx, nil) + conn, err := e.dbPool.Conn(ctx) if err != nil { return nil, err } - executorTx := &ExecutorTx{tx: tx} + tx, err := conn.BeginTx(ctx, nil) + if err != nil { + // Turso can execute BEGIN, then report context cancellation without + // rolling back. Discard the connection while we still own it so an + // open transaction can't be returned to the pool. + _ = conn.Raw(func(any) error { return driver.ErrBadConn }) + _ = conn.Close() + return nil, err + } + + // database/sql rolls back on cancellation, but a transaction begun on a + // dedicated connection doesn't return that connection to the pool. Close + // it on cancellation too; Close waits for the rollback to finish. + stopCloseFunc := context.AfterFunc(ctx, func() { _ = conn.Close() }) + executorTx := &ExecutorTx{ + releaseConnFunc: func() { + stopCloseFunc() + _ = conn.Close() + }, + tx: tx, + } executorTx.Executor = Executor{nil, templateReplaceWrapper{tx, &e.driver.replacer}, e.driver, executorTx} return executorTx, nil @@ -294,6 +311,8 @@ func (e *Executor) IndexesExist(ctx context.Context, params *riverdriver.Indexes return exists, nil } +func (e *Executor) InitDriver(context.Context) error { return nil } + func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelParams) (*rivertype.JobRow, error) { // Unlike Postgres, this must be carried out in two operations because // SQLite doesn't support CTEs containing `UPDATE`. As long as the job @@ -573,6 +592,11 @@ func (e *Executor) JobGetByKindMany(ctx context.Context, params *riverdriver.Job return sliceutil.MapError(jobs, jobRowFromInternal) } +func (e *Executor) JobGetCancelRequested(ctx context.Context, params *riverdriver.JobGetCancelRequestedParams) ([]int64, error) { + ids, err := dbsqlc.New().JobGetCancelRequested(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.ID) + return ids, interpretError(err) +} + func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetStuckParams) ([]*rivertype.JobRow, error) { jobs, err := dbsqlc.New().JobGetStuck(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobGetStuckParams{ AfterID: params.AfterID, @@ -611,7 +635,7 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo return &riverdriver.JobInsertFastResult{ Job: job, - UniqueSkippedAsDuplicate: gjson.GetBytes(job.Metadata, rivercommon.MetadataKeyUniqueNonce).Str != uniqueNonce, + UniqueSkippedAsDuplicate: riverdriver.UniqueInsertMetadataIsDuplicate(job.Metadata, uniqueNonce), }, nil }) } @@ -1216,6 +1240,10 @@ func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyMan return dbsqlc.New().NotificationInsertMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, notifications) } +func (e *Executor) Ping(ctx context.Context) error { + return e.Exec(ctx, "SELECT 1") +} + func (e *Executor) PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) { return nil, riverdriver.ErrNotImplemented } @@ -1420,7 +1448,8 @@ func (e *Executor) TableTruncate(ctx context.Context, params *riverdriver.TableT type ExecutorTx struct { Executor - tx *sql.Tx + releaseConnFunc func() + tx *sql.Tx } func (t *ExecutorTx) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { @@ -1434,11 +1463,17 @@ func (t *ExecutorTx) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) } func (t *ExecutorTx) Commit(ctx context.Context) error { + if t.releaseConnFunc != nil { + defer t.releaseConnFunc() + } // unfortunately, `database/sql` does not take a context ... return t.tx.Commit() } func (t *ExecutorTx) Rollback(ctx context.Context) error { + if t.releaseConnFunc != nil { + defer t.releaseConnFunc() + } // unfortunately, `database/sql` does not take a context ... return t.tx.Rollback() } @@ -1546,7 +1581,7 @@ func sqliteJobInsertFastManyJobsParam(jobs []*riverdriver.JobInsertFastParams, u metadata := sliceutil.FirstNonEmpty(job.Metadata, []byte("{}")) if uniqueNonce != "" { var err error - metadata, err = sjson.SetBytes(metadata, rivercommon.MetadataKeyUniqueNonce, uniqueNonce) + metadata, err = riverdriver.UniqueInsertMetadataWithNonce(metadata, uniqueNonce) if err != nil { return nil, err } diff --git a/riverdriver/riversqlite/river_sqlite_driver_test.go b/riverdriver/riversqlite/river_sqlite_driver_test.go index b2e86b3ca..9873917fc 100644 --- a/riverdriver/riversqlite/river_sqlite_driver_test.go +++ b/riverdriver/riversqlite/river_sqlite_driver_test.go @@ -3,6 +3,7 @@ package riversqlite import ( "context" "database/sql" + "database/sql/driver" "errors" "testing" "time" @@ -24,6 +25,75 @@ func TestDurationAsString(t *testing.T) { require.Equal(t, "3.255 seconds", durationAsString(3*time.Second+255*time.Millisecond)) } +func TestExecutorBegin(t *testing.T) { + t.Parallel() + + t.Run("FailedBeginDiscardsConnection", func(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var opened, closed int + dbPool := sql.OpenDB(&beginTestConnector{ + connectFunc: func() (driver.Conn, error) { + opened++ + return &beginTestConn{ + beginFunc: func() (driver.Tx, error) { + // Simulate a driver that starts a transaction, but + // returns a cancellation error without rolling back. + cancel() + return nil, context.Canceled + }, + closeFunc: func() { closed++ }, + }, nil + }, + }) + t.Cleanup(func() { require.NoError(t, dbPool.Close()) }) + dbPool.SetMaxOpenConns(1) + exec := New(dbPool).GetExecutor() + + tx, err := exec.Begin(ctx) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, tx) + require.Equal(t, 1, opened) + require.Equal(t, 1, closed) + + // The next checkout must open a new connection instead of reusing + // the one whose transaction state is unknown. + conn, err := dbPool.Conn(context.Background()) + require.NoError(t, err) + require.NoError(t, conn.Close()) + require.Equal(t, 2, opened) + }) +} + +type beginTestConn struct { + driver.Conn + + beginFunc func() (driver.Tx, error) + closeFunc func() +} + +func (c *beginTestConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + return c.beginFunc() +} + +func (c *beginTestConn) Close() error { + c.closeFunc() + return nil +} + +type beginTestConnector struct { + driver.Connector + + connectFunc func() (driver.Conn, error) +} + +func (c *beginTestConnector) Connect(context.Context) (driver.Conn, error) { + return c.connectFunc() +} + func TestInterpretError(t *testing.T) { t.Parallel() diff --git a/riverdriver/unique_insert.go b/riverdriver/unique_insert.go new file mode 100644 index 000000000..bdeaa07d6 --- /dev/null +++ b/riverdriver/unique_insert.go @@ -0,0 +1,110 @@ +package riverdriver + +import ( + "encoding/json" + "fmt" +) + +// UniqueInsertMetadataKey is a reserved job metadata key used to detect unique +// insert conflicts on databases that don't expose PostgreSQL system columns. +const UniqueInsertMetadataKey = "river:unique_nonce" + +// UniqueInsertMode is a database-specific strategy for detecting whether a +// unique insert returned a newly inserted job or an existing one. +type UniqueInsertMode uint32 + +const ( + // UniqueInsertModeUnknown indicates that a database's mode hasn't been + // detected yet. + UniqueInsertModeUnknown UniqueInsertMode = iota + + // UniqueInsertModeMetadataNonce detects conflicts by putting a nonce in the + // metadata of the proposed job and checking whether the returned job + // contains it. + UniqueInsertModeMetadataNonce + + // UniqueInsertModeReturningOld uses PostgreSQL 18's OLD row support in + // RETURNING. + UniqueInsertModeReturningOld + + // UniqueInsertModeXmax uses PostgreSQL's xmax system column. + UniqueInsertModeXmax +) + +// SQL returns the SQL expression for the mode. UniqueInsertModeMetadataNonce +// always returns false because duplicate detection is performed in Go instead. +func (m UniqueInsertMode) SQL() string { + switch m { + case UniqueInsertModeMetadataNonce: + return "false" + + case UniqueInsertModeReturningOld: + return "(OLD.id IS NOT NULL)" + + case UniqueInsertModeXmax: + return "(xmax != 0)" + + case UniqueInsertModeUnknown: + panic("unique insert mode has not been detected") + + default: + panic(fmt.Sprintf("invalid unique insert mode: %d", m)) + } +} + +// UniqueInsertMetadataIsDuplicate returns whether metadata lacks the nonce +// from a proposed insert, indicating that an existing row was returned +// instead. +func UniqueInsertMetadataIsDuplicate(metadata []byte, nonce string) bool { + var metadataMap map[string]json.RawMessage + if err := json.Unmarshal(metadata, &metadataMap); err != nil { + return true + } + + var metadataNonce string + if err := json.Unmarshal(metadataMap[UniqueInsertMetadataKey], &metadataNonce); err != nil { + return true + } + return metadataNonce != nonce +} + +// UniqueInsertMetadataWithNonce returns metadata with nonce set under +// UniqueInsertMetadataKey. +func UniqueInsertMetadataWithNonce(metadata []byte, nonce string) ([]byte, error) { + if len(metadata) == 0 { + metadata = []byte("{}") + } + + var metadataMap map[string]json.RawMessage + if err := json.Unmarshal(metadata, &metadataMap); err != nil { + return nil, fmt.Errorf("error unmarshaling job metadata: %w", err) + } + if metadataMap == nil { + metadataMap = make(map[string]json.RawMessage) + } + + nonceJSON, err := json.Marshal(nonce) + if err != nil { + return nil, fmt.Errorf("error marshaling unique insert nonce: %w", err) + } + metadataMap[UniqueInsertMetadataKey] = nonceJSON + + metadata, err = json.Marshal(metadataMap) + if err != nil { + return nil, fmt.Errorf("error marshaling job metadata: %w", err) + } + return metadata, nil +} + +// UniqueInsertModeFromProductAndVersion returns the unique insert mode +// appropriate for a database product and its PostgreSQL-compatible server +// version number. +func UniqueInsertModeFromProductAndVersion(product string, version int32) UniqueInsertMode { + if postgresProductIsYugabyte(product) { + return UniqueInsertModeMetadataNonce + } + if version >= 180_000 { + return UniqueInsertModeReturningOld + } + return UniqueInsertModeXmax +} diff --git a/riverdriver/unique_insert_test.go b/riverdriver/unique_insert_test.go new file mode 100644 index 000000000..3c78de985 --- /dev/null +++ b/riverdriver/unique_insert_test.go @@ -0,0 +1,126 @@ +package riverdriver + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestUniqueInsertMetadataIsDuplicate(t *testing.T) { + t.Parallel() + + t.Run("DifferentNonce", func(t *testing.T) { + t.Parallel() + + require.True(t, UniqueInsertMetadataIsDuplicate([]byte(`{"river:unique_nonce":"old"}`), "new")) + }) + + t.Run("InvalidMetadata", func(t *testing.T) { + t.Parallel() + + require.True(t, UniqueInsertMetadataIsDuplicate([]byte(`{`), "nonce")) + }) + + t.Run("MatchingNonce", func(t *testing.T) { + t.Parallel() + + require.False(t, UniqueInsertMetadataIsDuplicate([]byte(`{"river:unique_nonce":"nonce"}`), "nonce")) + }) + + t.Run("MissingNonce", func(t *testing.T) { + t.Parallel() + + require.True(t, UniqueInsertMetadataIsDuplicate([]byte(`{"existing":123}`), "nonce")) + }) +} + +func TestUniqueInsertMetadataWithNonce(t *testing.T) { + t.Parallel() + + t.Run("EmptyMetadata", func(t *testing.T) { + t.Parallel() + + metadata, err := UniqueInsertMetadataWithNonce(nil, "nonce") + require.NoError(t, err) + require.JSONEq(t, `{"river:unique_nonce":"nonce"}`, string(metadata)) + }) + + t.Run("ExistingMetadata", func(t *testing.T) { + t.Parallel() + + metadata, err := UniqueInsertMetadataWithNonce([]byte(`{"existing":123}`), "nonce") + require.NoError(t, err) + require.JSONEq(t, `{"existing":123,"river:unique_nonce":"nonce"}`, string(metadata)) + }) + + t.Run("ExistingNonce", func(t *testing.T) { + t.Parallel() + + metadata, err := UniqueInsertMetadataWithNonce([]byte(`{"river:unique_nonce":"old"}`), "new") + require.NoError(t, err) + require.JSONEq(t, `{"river:unique_nonce":"new"}`, string(metadata)) + }) + + t.Run("InvalidMetadata", func(t *testing.T) { + t.Parallel() + + _, err := UniqueInsertMetadataWithNonce([]byte(`{`), "nonce") + require.ErrorContains(t, err, "error unmarshaling job metadata") + }) +} + +func TestUniqueInsertModeFromProductAndVersion(t *testing.T) { + t.Parallel() + + t.Run("PostgreSQL17", func(t *testing.T) { + t.Parallel() + + require.Equal(t, UniqueInsertModeXmax, UniqueInsertModeFromProductAndVersion("PostgreSQL 17.5", 170_005)) + }) + + t.Run("PostgreSQL18", func(t *testing.T) { + t.Parallel() + + require.Equal(t, UniqueInsertModeReturningOld, UniqueInsertModeFromProductAndVersion("PostgreSQL 18.0", 180_000)) + }) + + t.Run("YugabyteByName", func(t *testing.T) { + t.Parallel() + + require.Equal(t, UniqueInsertModeMetadataNonce, UniqueInsertModeFromProductAndVersion("YugabyteDB", 180_000)) + }) + + t.Run("YugabytePostgreSQLVersion", func(t *testing.T) { + t.Parallel() + + require.Equal(t, UniqueInsertModeMetadataNonce, UniqueInsertModeFromProductAndVersion("PostgreSQL 15.2-YB-2.25.1.0-b0", 150_002)) + }) +} + +func TestUniqueInsertModeSQL(t *testing.T) { + t.Parallel() + + t.Run("MetadataNonce", func(t *testing.T) { + t.Parallel() + + require.Equal(t, "false", UniqueInsertModeMetadataNonce.SQL()) + }) + + t.Run("ReturningOld", func(t *testing.T) { + t.Parallel() + + require.Equal(t, "(OLD.id IS NOT NULL)", UniqueInsertModeReturningOld.SQL()) + }) + + t.Run("Unknown", func(t *testing.T) { + t.Parallel() + + require.PanicsWithValue(t, "unique insert mode has not been detected", func() { UniqueInsertModeUnknown.SQL() }) + }) + + t.Run("Xmax", func(t *testing.T) { + t.Parallel() + + require.Equal(t, "(xmax != 0)", UniqueInsertModeXmax.SQL()) + }) +} diff --git a/rivershared/riversharedtest/yugabyte.go b/rivershared/riversharedtest/yugabyte.go new file mode 100644 index 000000000..abb5c870e --- /dev/null +++ b/rivershared/riversharedtest/yugabyte.go @@ -0,0 +1,67 @@ +package riversharedtest + +import ( + "context" + "fmt" + "testing" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +// DBPoolWithYugabyteVersion returns a PostgreSQL pool that reports a Yugabyte +// version and LISTEN/NOTIFY setting. A nil setting simulates versions where it +// doesn't exist. This exercises detection on ordinary PostgreSQL; it does not +// emulate Yugabyte's storage or transaction semantics. +// +// The schema must be isolated to this test. When notifications are disabled, +// pg_notify raises an exception to catch accidental attempts to broadcast. +func DBPoolWithYugabyteVersion(ctx context.Context, t *testing.T, schema string, listenNotifyEnabled *bool) *pgxpool.Pool { + t.Helper() + + pool := DBPool(ctx, t) + setting := "NULL::text" + version := "2025.2.1.0" + if listenNotifyEnabled != nil { + version = "2025.2.3.0" + setting = "'off'::text" + if *listenNotifyEnabled { + setting = "'on'::text" + } + } + safeSchema := pgx.Identifier{schema}.Sanitize() + _, err := pool.Exec(ctx, fmt.Sprintf(` +CREATE FUNCTION %s.version() RETURNS text LANGUAGE sql AS $$ + SELECT 'PostgreSQL 15.12-YB-%s-b1'::text +$$; +CREATE FUNCTION %s.current_setting(setting_name text, missing_ok boolean) RETURNS text LANGUAGE sql AS $$ + SELECT CASE WHEN setting_name = 'yb_enable_listen_notify' THEN %s + ELSE pg_catalog.current_setting(setting_name, missing_ok) END +$$;`, safeSchema, version, safeSchema, setting)) + require.NoError(t, err) + t.Cleanup(func() { + _, err := pool.Exec(ctx, "DROP FUNCTION "+safeSchema+".version(), "+safeSchema+".current_setting(text, boolean)") + require.NoError(t, err) + }) + + if listenNotifyEnabled == nil || !*listenNotifyEnabled { + _, err := pool.Exec(ctx, "CREATE FUNCTION "+safeSchema+`.pg_notify(text, text) RETURNS void LANGUAGE plpgsql AS $$ +BEGIN RAISE EXCEPTION 'LISTEN/NOTIFY is unavailable'; END +$$;`) + require.NoError(t, err) + t.Cleanup(func() { + _, err := pool.Exec(ctx, "DROP FUNCTION "+safeSchema+".pg_notify(text, text)") + require.NoError(t, err) + }) + } + + config := pool.Config().Copy() + config.AfterConnect = nil // DBPool normally clears the search path. + config.MaxConns = 2 + config.ConnConfig.RuntimeParams["search_path"] = safeSchema + ", pg_catalog" + yugabytePool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + t.Cleanup(yugabytePool.Close) + return yugabytePool +}