diff --git a/postgres/README.md b/postgres/README.md index 265c0fb..1a10a39 100644 --- a/postgres/README.md +++ b/postgres/README.md @@ -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. diff --git a/postgres/go.mod b/postgres/go.mod index 88219e7..e0059f0 100644 --- a/postgres/go.mod +++ b/postgres/go.mod @@ -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 ) diff --git a/postgres/go.sum b/postgres/go.sum index aef17eb..e40cbc3 100644 --- a/postgres/go.sum +++ b/postgres/go.sum @@ -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= diff --git a/postgres/postgres.go b/postgres/postgres.go index d0c8c10..5e163c0 100644 --- a/postgres/postgres.go +++ b/postgres/postgres.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "log/slog" "net" "net/url" "time" @@ -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 { @@ -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 { @@ -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) @@ -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 { @@ -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)) } }() } @@ -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 @@ -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")) +} diff --git a/postgres/postgres_test.go b/postgres/postgres_test.go index e5d4ba9..2a45a7f 100644 --- a/postgres/postgres_test.go +++ b/postgres/postgres_test.go @@ -3,6 +3,7 @@ package postgres import ( "context" "errors" + "log/slog" "strings" "testing" "time" @@ -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 { @@ -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") + } +}