diff --git a/.gitignore b/.gitignore index b68996b..662310d 100644 --- a/.gitignore +++ b/.gitignore @@ -21,4 +21,6 @@ go.work # Mock gen files -mock*.go \ No newline at end of file +mock*.go + +.claude/settings.local.json diff --git a/outbox.go b/outbox.go index 97b2ace..a83b145 100644 --- a/outbox.go +++ b/outbox.go @@ -10,6 +10,7 @@ import ( "log/slog" "os" "reflect" + "sync" "time" "github.com/rs/xid" @@ -147,11 +148,20 @@ func (o outbox[T]) SendTx(ctx context.Context, tx *sql.Tx, msg T) error { func (o outbox[T]) dispatch() { tokens := make(chan struct{}, o.numRoutines) + var inFlight sync.Map + for id := range o.store.Listen() { + if _, loaded := inFlight.LoadOrStore(id, struct{}{}); loaded { + continue + } + tokens <- struct{}{} go func(id xid.ID) { + defer func() { + inFlight.Delete(id) + <-tokens + }() o.process(id) - <-tokens }(id) } } @@ -249,7 +259,7 @@ func (o outbox[T]) processMessageTx(ctx context.Context, id xid.ID) func(s Store } if errors.Is(err, errSkippingRecord) { - logger.InfoContext( + logger.DebugContext( ctx, "Skipping message", slog.String("reason", err.Error()), diff --git a/store/pg/pg.go b/store/pg/pg.go index 088d825..9686ddc 100644 --- a/store/pg/pg.go +++ b/store/pg/pg.go @@ -22,11 +22,15 @@ type execQuerier interface { } type Store struct { - db execQuerier - tableName string - connStr string - chanName string - logger *slog.Logger + db execQuerier + tableName string + connStr string + chanName string + logger *slog.Logger + pageSize int + chanBufferSize int + maxConsumerConns int + done chan struct{} } type option func(s *Store) @@ -47,15 +51,43 @@ func WithLogger(logger *slog.Logger) option { } } +func WithPageSize(n int) option { + return func(s *Store) { + if n > 0 { + s.pageSize = n + } + } +} + +func WithChannelBufferSize(n int) option { + return func(s *Store) { + if n > 0 { + s.chanBufferSize = n + } + } +} + +func WithMaxConsumerConns(n int) option { + return func(s *Store) { + if n > 0 { + s.maxConsumerConns = n + } + } +} + var _ outbox.Store = &Store{} func NewStore(db execQuerier, connStr string, opts ...option) (*Store, error) { s := &Store{ - db, - "outbox", - connStr, - "", - slog.Default(), + db: db, + tableName: "outbox", + connStr: connStr, + chanName: "", + logger: slog.Default(), + pageSize: 6, + chanBufferSize: 5, + maxConsumerConns: 8, + done: make(chan struct{}), } for _, o := range opts { @@ -73,13 +105,47 @@ func NewStore(db execQuerier, connStr string, opts ...option) (*Store, error) { "_", ) + // dedicated pool for consumer operations (queryPage, ProcessTx) so that Listen can keep polling + // for new work even if the caller is doing long-running work or has a slow connection + consumerDB, err := sql.Open("postgres", connStr) + if err != nil { + return nil, fmt.Errorf("open consumer db pool: %w", err) + } + consumerDB.SetMaxOpenConns(s.maxConsumerConns) + consumerDB.SetMaxIdleConns(s.maxConsumerConns) + if err := consumerDB.Ping(); err != nil { + consumerDB.Close() + return nil, fmt.Errorf("ping consumer db pool: %w", err) + } + s.db = consumerDB + if err := s.init(); err != nil { + consumerDB.Close() return nil, err } return s, nil } +// Close signals the Listen goroutine to stop and shuts down the dedicated +// consumer connection pool. It does not close the caller-owned db passed to NewStore. +func (s *Store) Close() error { + select { + case <-s.done: + // already closed + default: + close(s.done) + } + if s.db != nil { + if db, ok := s.db.(interface { + Close() error + }); ok { + db.Close() + } + } + return nil +} + func (s Store) CreateRecordTx(ctx context.Context, tx *sql.Tx, r outbox.Record) (*outbox.Record, error) { query := fmt.Sprintf(` INSERT INTO %s VALUES ($1, $2); @@ -121,20 +187,18 @@ func (s Store) Listen() <-chan xid.ID { s.chanName, ) - idChan := make(chan xid.ID, 1) + idChan := make(chan xid.ID, s.chanBufferSize) go func(l *pq.Listener) { + defer l.Close() for { - ids, err := s.getRecordIDs() + err := s.getRecordIDs(idChan) if err != nil { s.logger.Error("unable to get record ids", "error", err) - continue - } - - for _, i := range ids { - idChan <- i } select { + case <-s.done: + return case <-l.Notify: // New record(s) available to process case <-time.After(90 * time.Second): @@ -148,21 +212,65 @@ func (s Store) Listen() <-chan xid.ID { return idChan } -func (s Store) getRecordIDs() ([]xid.ID, error) { - var res []xid.ID +func (s Store) getRecordIDs(idChan chan xid.ID) error { + var lastID string + for { + n, err := s.fetchPage(idChan, &lastID) + if err != nil { + return err + } + if n == 0 { + return nil + } + } +} + +func (s Store) fetchPage(idChan chan xid.ID, lastID *string) (int, error) { + ids, err := s.queryPage(*lastID) + if err != nil { + return 0, err + } + + // Send to the channel outside of the DB context so that blocking on a + // full channel does not hold open rows or trigger a context timeout. + for _, id := range ids { + select { + case idChan <- id: + case <-s.done: + return 0, nil + } + } + + if len(ids) > 0 { + *lastID = ids[len(ids)-1].String() + } + + return len(ids), nil +} + +// queryPage executes a single keyset-paginated query and returns up to +// pageSize IDs. The context timeout only covers the DB round-trip; channel +// backpressure cannot cause it to expire. +func (s Store) queryPage(afterID string) ([]xid.ID, error) { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - query := fmt.Sprintf(` - SELECT id FROM %s; - `, s.tableName) + var query string + var args []any + if afterID == "" { + query = fmt.Sprintf(`SELECT id FROM %s ORDER BY id LIMIT %d;`, s.tableName, s.pageSize) + } else { + query = fmt.Sprintf(`SELECT id FROM %s WHERE id > $1 ORDER BY id LIMIT %d;`, s.tableName, s.pageSize) + args = []any{afterID} + } - rows, err := s.db.QueryContext(ctx, query) + rows, err := s.db.QueryContext(ctx, query, args...) if err != nil { return nil, err } defer rows.Close() + var ids []xid.ID for rows.Next() { var rawID string if err := rows.Scan(&rawID); err != nil { @@ -174,10 +282,14 @@ func (s Store) getRecordIDs() ([]xid.ID, error) { return nil, err } - res = append(res, id) + ids = append(ids, id) + } + + if err := rows.Err(); err != nil { + return nil, err } - return res, nil + return ids, nil } func (s Store) GetWithLock(ctx context.Context, id xid.ID) (*outbox.Record, error) { @@ -234,14 +346,18 @@ func (s Store) ProcessTx(ctx context.Context, fn func(outbox.Store) bool) error if err != nil { return fmt.Errorf("unable to create transaction: %v", err) } + defer tx.Rollback() // no-op after Commit; silently handles context cancellation store := Store{ - db: tx, - tableName: s.tableName, + db: tx, + tableName: s.tableName, + logger: s.logger, + pageSize: s.pageSize, + chanBufferSize: s.chanBufferSize, } if success := fn(store); !success { - return tx.Rollback() + return nil // rollback handled by defer; real error already logged in callback } return tx.Commit() diff --git a/store/pg/pg_suite_test.go b/store/pg/pg_suite_test.go index 2583521..5e974a7 100644 --- a/store/pg/pg_suite_test.go +++ b/store/pg/pg_suite_test.go @@ -13,7 +13,7 @@ import ( var ( db *sql.DB - connStr = getConfig().GenerateAddress() + connStr string ) func TestIntegration(t *testing.T) { @@ -28,6 +28,7 @@ var _ = BeforeEach(func() { var err error cfg := getConfig() cfg.Database = "outbox_test" + connStr = cfg.GenerateAddress() db, err = pghelpers.ConnectPostgres(*cfg) Expect(err).To(Succeed()) }) diff --git a/store/pg/pg_test.go b/store/pg/pg_test.go index 9d7a58f..aedb943 100644 --- a/store/pg/pg_test.go +++ b/store/pg/pg_test.go @@ -21,6 +21,38 @@ var _ = Describe("pgStore", func() { subject, err := NewStore(db, connStr) Expect(err).To(Succeed()) Expect(subject).ToNot(BeNil()) + defer subject.Close() + }) + + It("should create a separate consumer pool", func() { + subject, err := NewStore(db, connStr) + Expect(err).To(Succeed()) + defer subject.Close() + + Expect(subject.db).ToNot(BeNil()) + Expect(subject.db.(interface{ Ping() error }).Ping()).To(Succeed()) + }) + + It("should respect WithMaxConsumerConns option", func() { + subject, err := NewStore(db, connStr, WithMaxConsumerConns(3)) + Expect(err).To(Succeed()) + defer subject.Close() + + stats := subject.db.(interface{ Stats() sql.DBStats }).Stats() + Expect(stats.MaxOpenConnections).To(Equal(3)) + }) + }) + + Describe("#Close", func() { + It("should close the consumer pool without error", func() { + subject, err := NewStore(db, connStr) + Expect(err).To(Succeed()) + + err = subject.Close() + Expect(err).To(Succeed()) + + err = subject.db.(interface{ Ping() error }).Ping() + Expect(err).To(HaveOccurred()) }) }) @@ -39,6 +71,10 @@ var _ = Describe("pgStore", func() { tx, _ = db.BeginTx(ctx, nil) }) + AfterEach(func() { + subject.Close() + }) + It("should save the provided record on a successfull transaction", func() { res, err := subject.CreateRecordTx(ctx, tx, record) Expect(err).To(Succeed()) @@ -74,6 +110,10 @@ var _ = Describe("pgStore", func() { Expect(err).To(Succeed()) }) + AfterEach(func() { + subject.Close() + }) + It("should return a record on a valid id", func() { subject.ProcessTx(ctx, func(s outbox.Store) bool { res, err := s.GetWithLock(ctx, id) @@ -102,6 +142,7 @@ var _ = Describe("pgStore", func() { subject = createStore() ids = []xid.ID{xid.New(), xid.New(), xid.New()} ) + defer subject.Close() for _, id := range ids { _ = insertRecord(db, outbox.Record{ID: id, Message: []byte("data")}) @@ -132,6 +173,7 @@ var _ = Describe("pgStore", func() { ctx = context.Background() tx, _ = db.BeginTx(ctx, nil) ) + defer subject.Close() res, err := subject.CreateRecordTx(ctx, tx, outbox.Record{}) tx.Commit() @@ -147,6 +189,7 @@ var _ = Describe("pgStore", func() { subject = createStore() id = xid.New() ) + defer subject.Close() _ = insertRecord(db, outbox.Record{ID: id, Message: []byte("data")}) @@ -164,6 +207,7 @@ var _ = Describe("pgStore", func() { subject = createStore() id = xid.New() ) + defer subject.Close() _ = insertRecord(db, outbox.Record{ID: id, Message: []byte("data")}) @@ -220,6 +264,7 @@ var _ = Describe("pgStore", func() { ) _ = insertRecord(db, record) + defer subject.Close() err := subject.Update(context.Background(), &record) Expect(err).To(Succeed())