Skip to content
Merged
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
13 changes: 13 additions & 0 deletions postgres/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,19 @@ func (s *Storage) Backfill(ctx context.Context, data map[int64]string) error {
}
```

`Connect` and `ConnectDSN` take options:
```go
db, closeDB, err := postgres.Connect(ctx, config,
postgres.WithLogger(slog.Default().ErrorContext),
postgres.WithTracer(&tracelog.TraceLog{Logger: tracelog.LoggerFunc(logQuery)}),
)
```

`WithLogger` redirects the background health check, which otherwise reports failures to
the global logger. `WithTracer` takes a `pgx.QueryTracer`; pgx also uses it for the batch,
copy, prepare, connect, acquire and release traces it implements, so `otelpgx` and
`pgx/v5/tracelog` both go through this one option.

`closeDB` shuts the pool down and stops the background health check; it matches
`closer.CloseFunc`, so it can be handed to `closer.Add` directly.

Expand Down
2 changes: 1 addition & 1 deletion postgres/go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ module github.com/draincloud/callpack/postgres
go 1.26.3

require (
github.com/draincloud/logger v0.0.4
github.com/draincloud/logger v0.0.5
github.com/jackc/pgx/v5 v5.10.0
)

Expand Down
4 changes: 2 additions & 2 deletions postgres/go.sum
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/draincloud/logger v0.0.4 h1:BEbNyy8AkuqQy60uST7O9IgW8ofl1zI6iN4tygFTwAg=
github.com/draincloud/logger v0.0.4/go.mod h1:t/vRd80p4+3cQg7HlmRCm5/GSTHfTWDC0WWANdtKYMw=
github.com/draincloud/logger v0.0.5 h1:4saeda/sm5E6S0lnfwCIdgmqZYiIGuTIUN0HpIlm3js=
github.com/draincloud/logger v0.0.5/go.mod h1:Z/GP5qHAC+MrhTMs5FckdQahiP2MjugaeQxjaMMaqWY=
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=
github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU=
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
Expand Down
70 changes: 67 additions & 3 deletions postgres/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"log/slog"
"net"
"net/url"
"time"
Expand All @@ -26,6 +27,10 @@ type Config struct {
MaxConns int
MaxConnIdleTime time.Duration
MaxConnLifetime time.Duration

logger *slog.Logger
tracer pgx.QueryTracer
afterConnect func(context.Context, *pgx.Conn) error
}

func (c Config) dsn() string {
Expand All @@ -52,11 +57,48 @@ func (c Config) dsn() string {
}

type DB struct {
db *pgxpool.Pool
db *pgxpool.Pool
log *slog.Logger
}

type ConnectOpt func(c *Config)

func WithLogger(logger *slog.Logger) ConnectOpt {
return func(c *Config) {
c.logger = logger
}
}

func WithTracer(t pgx.QueryTracer) ConnectOpt {
return func(c *Config) {
c.tracer = t
}
}

func WithAfterConnect(fn func(ctx context.Context, conn *pgx.Conn) error) ConnectOpt {
return func(c *Config) {
c.afterConnect = fn
}
}

func WithMaxConns(n int) ConnectOpt {
return func(c *Config) {
c.MaxConns = n
}
}

func WithMaxConnIdleTime(d time.Duration) ConnectOpt {
return func(c *Config) {
c.MaxConnIdleTime = d
}
}

func WithMaxConnLifetime(d time.Duration) ConnectOpt {
return func(c *Config) {
c.MaxConnLifetime = d
}
}

func ConnectDSN(ctx context.Context, dsn string, opts ...ConnectOpt) (*DB, func(context.Context) error, error) {
pgconfig, err := pgxpool.ParseConfig(dsn)
if err != nil {
Expand Down Expand Up @@ -99,6 +141,14 @@ func connect(ctx context.Context, pgconfig *pgxpool.Config, cfg Config) (*DB, fu
pgconfig.MaxConnLifetime = cfg.MaxConnLifetime
}

if cfg.tracer != nil {
pgconfig.ConnConfig.Tracer = cfg.tracer
}

if cfg.afterConnect != nil {
pgconfig.AfterConnect = cfg.afterConnect
}

pool, err := pgxpool.NewWithConfig(ctx, pgconfig)
if err != nil {
return nil, noopCloser, fmt.Errorf("failed to connect to postgres: %w", err)
Expand All @@ -112,7 +162,13 @@ func connect(ctx context.Context, pgconfig *pgxpool.Config, cfg Config) (*DB, fu

pingCtx, cancelPing := context.WithCancel(context.WithoutCancel(ctx))

d := &DB{db: pool}
log := cfg.logger
if log == nil {
log = logger.FromContext(ctx)
}
setupLogger(log)

d := &DB{db: pool, log: log}
go d.asyncPing(pingCtx)

return d, func(context.Context) error {
Expand Down Expand Up @@ -146,7 +202,7 @@ func (d *DB) asyncPing(ctx context.Context) {
func() {
defer t.Reset(dur)
if err := d.Ping(ctx); err != nil {
logger.Error(ctx, "DB.asyncPing error", logger.Err(err))
d.log.Error("DB.asyncPing error", logger.Err(err))
}
}()
}
Expand Down Expand Up @@ -199,6 +255,10 @@ func (d *DB) WithTransaction(ctx context.Context, fn func(context.Context) error
return fn(txContext(ctx, tx))
}

func InTransaction(ctx context.Context) bool {
return txFromContext(ctx) != nil
}

func Conn(ctx context.Context, db DBTX) DBTX {
if tx := txFromContext(ctx); tx != nil {
return tx
Expand All @@ -218,3 +278,7 @@ func txFromContext(ctx context.Context) pgx.Tx {
func txContext(parent context.Context, tx pgx.Tx) context.Context {
return context.WithValue(parent, ctxKey, tx)
}

func setupLogger(log *slog.Logger) {
log.With(slog.String("system", "postgres.database"))
}
185 changes: 184 additions & 1 deletion postgres/postgres_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package postgres
import (
"context"
"errors"
"log/slog"
"strings"
"testing"
"time"
Expand Down Expand Up @@ -283,7 +284,7 @@ func TestAsyncPingStopsWhenContextIsCancelled(t *testing.T) {
done := make(chan struct{})
go func() {
defer close(done)
(&DB{db: newLazyPool(t)}).asyncPing(ctx)
(&DB{db: newLazyPool(t), log: slog.New(slog.DiscardHandler)}).asyncPing(ctx)
}()

select {
Expand All @@ -292,3 +293,185 @@ func TestAsyncPingStopsWhenContextIsCancelled(t *testing.T) {
t.Fatal("asyncPing outlived its context")
}
}

type stubTracer struct{}

func (stubTracer) TraceQueryStart(ctx context.Context, _ *pgx.Conn, _ pgx.TraceQueryStartData) context.Context {
return ctx
}

func (stubTracer) TraceQueryEnd(context.Context, *pgx.Conn, pgx.TraceQueryEndData) {}

func TestWithTracerReachesTheConnConfig(t *testing.T) {
pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn())
if err != nil {
t.Fatal(err)
}

tracer := stubTracer{}

cfg := Config{}
WithTracer(tracer)(&cfg)

if _, closeDB, err := connect(t.Context(), pgconfig, cfg); err == nil {
closeDB(t.Context())
t.Fatal("expected the ping to fail")
}

if pgconfig.ConnConfig.Tracer != pgx.QueryTracer(tracer) {
t.Fatalf("Tracer = %#v, want the tracer the option was given", pgconfig.ConnConfig.Tracer)
}
}

func TestWithTracerLeavesTheConnConfigAloneWhenUnset(t *testing.T) {
pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn())
if err != nil {
t.Fatal(err)
}

if _, closeDB, err := connect(t.Context(), pgconfig, Config{}); err == nil {
closeDB(t.Context())
t.Fatal("expected the ping to fail")
}

if pgconfig.ConnConfig.Tracer != nil {
t.Fatalf("Tracer = %#v, want nil", pgconfig.ConnConfig.Tracer)
}
}

// capturingHandler forwards each record's message to a channel, so a test can
// assert on what the configured logger was asked to write.
type capturingHandler struct {
msgs chan<- string
}

func (capturingHandler) Enabled(context.Context, slog.Level) bool { return true }

func (h capturingHandler) Handle(_ context.Context, r slog.Record) error {
select {
case h.msgs <- r.Message:
default:
}

return nil
}

func (h capturingHandler) WithAttrs([]slog.Attr) slog.Handler { return h }

func (h capturingHandler) WithGroup(string) slog.Handler { return h }

func TestWithLoggerReceivesHealthCheckFailures(t *testing.T) {
logged := make(chan string, 1)

cfg := Config{}
WithLogger(slog.New(capturingHandler{msgs: logged}))(&cfg)

ctx, cancel := context.WithCancel(t.Context())
defer cancel()

db := &DB{db: newLazyPool(t), log: cfg.logger}

done := make(chan struct{})
go func() {
defer close(done)
db.asyncPing(ctx)
}()

select {
case msg := <-logged:
if !strings.Contains(msg, "asyncPing") {
t.Fatalf("logged %q, want the health check failure", msg)
}
case <-time.After(5 * time.Second):
t.Fatal("the configured logger never saw the failed health check")
}

cancel()
<-done
}

func TestWithAfterConnectReachesThePoolConfig(t *testing.T) {
pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn())
if err != nil {
t.Fatal(err)
}

sentinel := errors.New("sentinel")

cfg := Config{}
WithAfterConnect(func(context.Context, *pgx.Conn) error { return sentinel })(&cfg)

if _, closeDB, err := connect(t.Context(), pgconfig, cfg); err == nil {
closeDB(t.Context())
t.Fatal("expected the ping to fail")
}

if pgconfig.AfterConnect == nil {
t.Fatal("AfterConnect = nil, want the hook the option was given")
}

// Compared by what it does: func values are not comparable.
if err := pgconfig.AfterConnect(t.Context(), nil); !errors.Is(err, sentinel) {
t.Fatalf("AfterConnect returned %v, want the hook the option was given", err)
}
}

func TestWithAfterConnectLeavesThePoolConfigAloneWhenUnset(t *testing.T) {
pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn())
if err != nil {
t.Fatal(err)
}

if _, closeDB, err := connect(t.Context(), pgconfig, Config{}); err == nil {
closeDB(t.Context())
t.Fatal("expected the ping to fail")
}

if pgconfig.AfterConnect != nil {
t.Fatal("AfterConnect is set, want nil")
}
}

func TestPoolLimitOptionsReachThePoolConfig(t *testing.T) {
pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn())
if err != nil {
t.Fatal(err)
}

cfg := Config{}
for _, o := range []ConnectOpt{
WithMaxConns(20),
WithMaxConnIdleTime(5 * time.Minute),
WithMaxConnLifetime(30 * time.Second),
} {
o(&cfg)
}

if _, closeDB, err := connect(t.Context(), pgconfig, cfg); err == nil {
closeDB(t.Context())
t.Fatal("expected the ping to fail")
}

if pgconfig.MaxConns != 20 {
t.Fatalf("MaxConns = %d, want 20", pgconfig.MaxConns)
}

if pgconfig.MaxConnIdleTime != 5*time.Minute {
t.Fatalf("MaxConnIdleTime = %v, want 5m", pgconfig.MaxConnIdleTime)
}

if pgconfig.MaxConnLifetime != 30*time.Second {
t.Fatalf("MaxConnLifetime = %v, want 30s", pgconfig.MaxConnLifetime)
}
}

func TestInTransactionTracksTheContext(t *testing.T) {
if InTransaction(t.Context()) {
t.Fatal("InTransaction = true on a bare context")
}

ctx := txContext(t.Context(), nil)
if InTransaction(ctx) {
t.Fatal("InTransaction = true for a nil transaction")
}
}
Loading