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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 47 additions & 11 deletions driver/natsqueue/worker_nats_impl.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,9 +31,11 @@ type natsWorker struct {
sub natsWorkerSubscription
sem chan struct{}

running sync.WaitGroup
delayed sync.WaitGroup
observer queue.Observer
running sync.WaitGroup
delayed sync.WaitGroup
delayedMu sync.Mutex
delayedJobs map[*time.Timer]natsMessage
observer queue.Observer
}

type natsWorkerSubscription interface {
Expand Down Expand Up @@ -84,6 +86,7 @@ func newNATSWorkerWithConfig(cfg natsWorkerConfig) *natsWorker {
defaultQueue: cfg.DefaultQueue,
workers: cfg.Workers,
handlers: make(map[string]queue.Handler),
delayedJobs: make(map[*time.Timer]natsMessage),
observer: cfg.Observer,
}
}
Expand Down Expand Up @@ -177,7 +180,7 @@ func (w *natsWorker) Shutdown(ctx context.Context) error {
stopErr = sub.Drain()
}
w.running.Wait()
w.delayed.Wait()
w.flushDelayedJobs()
if conn != nil {
if drainErr := conn.Drain(); stopErr == nil {
stopErr = drainErr
Expand Down Expand Up @@ -216,13 +219,7 @@ func (w *natsWorker) processMessage(message *nats.Msg) {
if incoming.AvailableAtMS > 0 {
remaining := time.Until(time.UnixMilli(incoming.AvailableAtMS))
if remaining > 0 {
w.delayed.Add(1)
time.AfterFunc(remaining, func() {
defer w.delayed.Done()
if err := w.republish(incoming); err != nil {
w.observeRepublishFailure(context.Background(), incoming, err)
}
})
w.scheduleDelayedJob(incoming, remaining)
return
}
}
Expand Down Expand Up @@ -268,6 +265,45 @@ func (w *natsWorker) processMessage(message *nats.Msg) {
}
}

// scheduleDelayedJob keeps Core NATS delayed work recoverable because the broker cannot retain a claimed callback until its due time.
func (w *natsWorker) scheduleDelayedJob(message natsMessage, remaining time.Duration) {
w.delayedMu.Lock()
w.delayed.Add(1)
var timer *time.Timer
timer = time.AfterFunc(remaining, func() {
w.delayedMu.Lock()
delete(w.delayedJobs, timer)
w.delayedMu.Unlock()
defer w.delayed.Done()
if err := w.republish(message); err != nil {
w.observeRepublishFailure(context.Background(), message, err)
}
})
w.delayedJobs[timer] = message
w.delayedMu.Unlock()
}

// flushDelayedJobs returns timer-owned work to Core NATS before the consumer connection closes so maintenance does not wait for long delays.
func (w *natsWorker) flushDelayedJobs() {
w.delayedMu.Lock()
messages := make([]natsMessage, 0, len(w.delayedJobs))
for timer, message := range w.delayedJobs {
if !timer.Stop() {
continue
}
delete(w.delayedJobs, timer)
messages = append(messages, message)
w.delayed.Done()
}
w.delayedMu.Unlock()
for _, message := range messages {
if err := w.republish(message); err != nil {
w.observeRepublishFailure(context.Background(), message, err)
}
}
w.delayed.Wait()
}

// connectNATSWorker creates the Core NATS subscription owned by one worker lifecycle.
func connectNATSWorker(url, subject string, callback nats.MsgHandler) (natsConnection, natsWorkerSubscription, error) {
nc, err := nats.Connect(url)
Expand Down
10 changes: 6 additions & 4 deletions driver/natsqueue/worker_nats_impl_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -473,8 +473,8 @@ func TestNATSWorkerShutdownWaitsForInFlightRepublish(t *testing.T) {
}
}

// TestNATSWorkerShutdownTracksDelayedRepublish verifies timer-backed accepted work finishes before connection drain.
func TestNATSWorkerShutdownTracksDelayedRepublish(t *testing.T) {
// TestNATSWorkerShutdownFlushesDelayedRepublish verifies maintenance can return long-delay work to Core NATS without waiting for its due time.
func TestNATSWorkerShutdownFlushesDelayedRepublish(t *testing.T) {
w := newNATSWorker("nats://example:4222")
connection, subscription := newNATSWorkerLifecycleStubs()
w.conn = connection
Expand All @@ -483,13 +483,15 @@ func TestNATSWorkerShutdownTracksDelayedRepublish(t *testing.T) {
payload, err := json.Marshal(natsMessage{
Type: "job:delayed-shutdown",
Queue: "default",
AvailableAtMS: time.Now().Add(25 * time.Millisecond).UnixMilli(),
AvailableAtMS: time.Now().Add(time.Hour).UnixMilli(),
})
if err != nil {
t.Fatalf("marshal: %v", err)
}
w.processMessage(&nats.Msg{Data: payload})
if err := w.Shutdown(context.Background()); err != nil {
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel()
if err := w.Shutdown(ctx); err != nil {
t.Fatalf("shutdown: %v", err)
}
select {
Expand Down
105 changes: 105 additions & 0 deletions driver/sqlitequeue/worker_lifecycle_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,105 @@
package sqlitequeue

import (
"context"
"path/filepath"
"sync/atomic"
"testing"
"time"

"github.com/goforj/queue"
)

// TestSQLiteWorkerLifecycleLeavesPausedJobsPending verifies maintenance-style pauses happen before durable claims.
func TestSQLiteWorkerLifecycleLeavesPausedJobsPending(t *testing.T) {
dsn := filepath.Join(t.TempDir(), "queue.db")
q, err := New(dsn, queue.WithWorkers(1))
if err != nil {
t.Fatalf("new SQLite queue: %v", err)
}
t.Cleanup(func() { _ = q.Shutdown(context.Background()) })

handled := make(chan struct{}, 1)
var calls atomic.Int64
q.Register("reports:build", func(context.Context, queue.Message) error {
calls.Add(1)
handled <- struct{}{}
return nil
})
if err := q.PauseWorkers(context.Background()); err != nil {
t.Fatalf("pause workers before startup: %v", err)
}
if _, err := q.Dispatch(queue.NewJob("reports:build").OnQueue("default")); err != nil {
t.Fatalf("dispatch while paused: %v", err)
}
select {
case <-handled:
t.Fatal("paused SQLite worker claimed the job")
case <-time.After(150 * time.Millisecond):
}
stats, err := q.Stats(context.Background())
if err != nil {
t.Fatalf("read paused queue stats: %v", err)
}
counters := stats.ByQueue["default"]
if counters.Pending != 1 || counters.Active != 0 || counters.Failed != 0 || counters.Archived != 0 {
t.Fatalf("paused queue counters = %+v, want one untouched pending job", counters)
}

if err := q.ResumeWorkers(context.Background()); err != nil {
t.Fatalf("resume workers: %v", err)
}
select {
case <-handled:
case <-time.After(2 * time.Second):
t.Fatal("resumed SQLite worker did not execute the pending job")
}
counters = waitForSQLiteQueueSettlement(t, q)
if calls.Load() != 1 || counters.Failed != 0 || counters.Archived != 0 {
t.Fatalf("resumed queue counters = %+v, want one successful execution", counters)
}

if err := q.PauseWorkers(context.Background()); err != nil {
t.Fatalf("pause live workers: %v", err)
}
if _, err := q.Dispatch(queue.NewJob("reports:build").OnQueue("default")); err != nil {
t.Fatalf("dispatch during live pause: %v", err)
}
select {
case <-handled:
t.Fatal("live-paused SQLite worker claimed the job")
case <-time.After(150 * time.Millisecond):
}
if err := q.ResumeWorkers(context.Background()); err != nil {
t.Fatalf("resume live workers: %v", err)
}
select {
case <-handled:
case <-time.After(2 * time.Second):
t.Fatal("live-resumed SQLite worker did not execute the pending job")
}
counters = waitForSQLiteQueueSettlement(t, q)
if calls.Load() != 2 || counters.Failed != 0 || counters.Archived != 0 {
t.Fatalf("live-resumed queue counters = %+v, want two successful executions", counters)
}
}

// waitForSQLiteQueueSettlement waits for the durable success update that follows handler return.
func waitForSQLiteQueueSettlement(t *testing.T, q *queue.Queue) queue.QueueCounters {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for {
stats, err := q.Stats(context.Background())
if err != nil {
t.Fatalf("read resumed queue stats: %v", err)
}
counters := stats.ByQueue["default"]
if counters.Pending == 0 && counters.Active == 0 {
return counters
}
if time.Now().After(deadline) {
t.Fatalf("queue did not settle after handler completion; counters = %+v", counters)
}
time.Sleep(10 * time.Millisecond)
}
}
54 changes: 41 additions & 13 deletions driver/sqlqueuecore/queue_database_impl.go
Original file line number Diff line number Diff line change
Expand Up @@ -109,12 +109,14 @@ type databaseQueue struct {
cfg localDatabaseConfig
db *sql.DB

ownsDB bool
ownsDB bool
prepareSchemaOnUse bool

mu sync.RWMutex
handlers map[string]queue.Handler

startMu sync.Mutex
schemaMu sync.Mutex
shutdownOnce sync.Once
shutdownDone chan struct{}
closeOnce sync.Once
Expand All @@ -124,6 +126,7 @@ type databaseQueue struct {
shutdownCh chan struct{}

started atomic.Bool
schemaReady atomic.Bool
shuttingDown atomic.Bool
uniqueClaims atomic.Uint64
continuation *busruntime.ContinuationScope
Expand Down Expand Up @@ -204,13 +207,14 @@ func New(cfg queue.DatabaseConfig) (*databaseQueue, error) {
}

d := &databaseQueue{
cfg: local,
db: cfg.DB,
handlers: make(map[string]queue.Handler),
shutdownCh: make(chan struct{}),
ownsDB: ownsDB,
continuation: busruntime.NewContinuationScope(),
observer: cfg.Observer,
cfg: local,
db: cfg.DB,
handlers: make(map[string]queue.Handler),
shutdownCh: make(chan struct{}),
ownsDB: ownsDB,
prepareSchemaOnUse: true,
continuation: busruntime.NewContinuationScope(),
observer: cfg.Observer,
}
if cfg.DriverName == "sqlite" {
d.db.SetMaxOpenConns(1)
Expand Down Expand Up @@ -267,11 +271,7 @@ func (d *databaseQueue) StartWorkers(ctx context.Context) error {
if d.started.Load() {
return nil
}
if d.cfg.AutoMigrate && !d.cfg.DisableAutoMigrate {
if err := d.ensureSchema(ctx); err != nil {
return err
}
} else if err := d.requireManagedQueueSchema(ctx); err != nil {
if err := d.prepareSchema(ctx); err != nil {
return err
}
for i := 0; i < d.cfg.Workers; i++ {
Expand All @@ -282,6 +282,29 @@ func (d *databaseQueue) StartWorkers(ctx context.Context) error {
return nil
}

// prepareSchema keeps producer-only and worker database handles on the same lazy, retryable schema boundary.
func (d *databaseQueue) prepareSchema(ctx context.Context) error {
if d.schemaReady.Load() {
return nil
}
d.schemaMu.Lock()
defer d.schemaMu.Unlock()
if d.schemaReady.Load() {
return nil
}
var err error
if d.cfg.AutoMigrate && !d.cfg.DisableAutoMigrate {
err = d.ensureSchema(ctx)
} else {
err = d.requireManagedQueueSchema(ctx)
}
if err != nil {
return err
}
d.schemaReady.Store(true)
return nil
}

// Shutdown drains workers before closing only database handles opened by this queue.
func (d *databaseQueue) Shutdown(ctx context.Context) error {
if err := d.DrainWorkers(ctx); err != nil {
Expand Down Expand Up @@ -374,6 +397,11 @@ func (d *databaseQueue) Dispatch(ctx context.Context, job queue.Job) error {
return err
}
}
if d.prepareSchemaOnUse {
if err := d.prepareSchema(ctx); err != nil {
return err
}
}
parsed := queuecore.DriverOptions(job)
payloadBytes := job.PayloadBytes()
if payloadBytes == nil {
Expand Down
12 changes: 9 additions & 3 deletions driver/sqlqueuecore/wrapper.go
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ type ModuleConfig struct {
// for a specific SQL driver name.
func NewQueue(driverName string, cfg ModuleConfig, opts ...queue.Option) (*queue.Queue, error) {
observer := driverbridge.NewObserverSink(cfg.Observer)
backend, err := New(queue.DatabaseConfig{
databaseConfig := queue.DatabaseConfig{
DB: cfg.DB,
DriverName: driverName,
DSN: cfg.DSN,
Expand All @@ -34,7 +34,8 @@ func NewQueue(driverName string, cfg ModuleConfig, opts ...queue.Option) (*queue
ProcessingLeaseNoTimeout: cfg.ProcessingLeaseNoTimeout,
Observer: observer,
Logger: cfg.Logger,
})
}
backend, err := New(databaseConfig)
if err != nil {
return nil, err
}
Expand All @@ -43,5 +44,10 @@ func NewQueue(driverName string, cfg ModuleConfig, opts ...queue.Option) (*queue
DefaultQueue: cfg.DefaultQueue,
Logger: cfg.Logger,
}
return driverbridge.NewQueueFromDriver(rootCfg, observer, backend, nil, opts...)
workerFactory := func(workers int) (any, error) {
workerConfig := databaseConfig
workerConfig.Workers = workers
return New(workerConfig)
}
return driverbridge.NewQueueFromDriver(rootCfg, observer, backend, workerFactory, opts...)
}
2 changes: 1 addition & 1 deletion driver_runtime.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ func newQueueFromDriver(cfg Config, observer Observer, backend driverQueueBacken

var q queueBackend
var runtime runtimeQueueBackend
if native, ok := backend.(driverRuntimeQueueBackend); ok {
if native, ok := backend.(driverRuntimeQueueBackend); ok && workerFactory == nil {
runtime = driverRuntimeQueueBackendAdapter{native}
q = runtime
} else {
Expand Down
8 changes: 8 additions & 0 deletions fake_queue.go
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,14 @@ func (f *FakeQueue) Register(string, Handler) {}
// _ = err
func (f *FakeQueue) StartWorkers(context.Context) error { return nil }

// PauseWorkers stops fake worker intake.
// @group Testing
func (f *FakeQueue) PauseWorkers(context.Context) error { return nil }

// ResumeWorkers restarts fake worker intake.
// @group Testing
func (f *FakeQueue) ResumeWorkers(context.Context) error { return nil }

// Workers preserves fluent lifecycle compatibility without creating workers.
// @group Testing
//
Expand Down
Loading
Loading