From 29d09b376acafae3b9185d7e309403cf2b0a7f8c Mon Sep 17 00:00:00 2001 From: comicrime Date: Sun, 30 Aug 2026 15:36:42 +0300 Subject: [PATCH 1/4] tmp --- postgres/go.mod | 20 ++++++ postgres/go.sum | 39 +++++++++++ postgres/interfaces.go | 20 ++++++ postgres/postgres.go | 154 +++++++++++++++++++++++++++++++++++++++++ 4 files changed, 233 insertions(+) create mode 100644 postgres/go.mod create mode 100644 postgres/go.sum create mode 100644 postgres/interfaces.go create mode 100644 postgres/postgres.go diff --git a/postgres/go.mod b/postgres/go.mod new file mode 100644 index 0000000..88219e7 --- /dev/null +++ b/postgres/go.mod @@ -0,0 +1,20 @@ +module github.com/draincloud/callpack/postgres + +go 1.26.3 + +require ( + github.com/draincloud/logger v0.0.4 + github.com/jackc/pgx/v5 v5.10.0 +) + +require ( + github.com/fatih/color v1.18.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + github.com/mattn/go-colorable v0.1.13 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + golang.org/x/sync v0.17.0 // indirect + golang.org/x/sys v0.25.0 // indirect + golang.org/x/text v0.29.0 // indirect +) diff --git a/postgres/go.sum b/postgres/go.sum new file mode 100644 index 0000000..aef17eb --- /dev/null +++ b/postgres/go.sum @@ -0,0 +1,39 @@ +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/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= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA= +github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg= +github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug= +golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.25.0 h1:r+8e+loiHxRqhXVl6ML1nO3l1+oFoWbnlu2Ehimmi34= +golang.org/x/sys v0.25.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk= +golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/postgres/interfaces.go b/postgres/interfaces.go new file mode 100644 index 0000000..b6fac42 --- /dev/null +++ b/postgres/interfaces.go @@ -0,0 +1,20 @@ +package postgres + +import ( + "context" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" +) + +type DBTX interface { + Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) + QueryRow(ctx context.Context, sql string, args ...any) pgx.Row + Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) +} + +type TXer interface { + BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) + Commit(ctx context.Context) error + Rollback(ctx context.Context) error +} diff --git a/postgres/postgres.go b/postgres/postgres.go new file mode 100644 index 0000000..3d0d36e --- /dev/null +++ b/postgres/postgres.go @@ -0,0 +1,154 @@ +package postgres + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/draincloud/logger" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +type Config struct { + Host string + Port string + Username string + Password string + Database string + AllowSSL bool + + MaxConns int + MaxConnIdleTime time.Duration + MaxConnLifetime time.Duration +} + +type DB struct { + db *pgxpool.Pool +} + +type ConnectOpt func(c *Config) + +func ConnectDSN(ctx context.Context, dsn string, opts ...ConnectOpt) (*DB, func() error, error) { + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + return nil, func() error { return nil }, fmt.Errorf("failed to parse config: %w", err) + } + + cfg := Config{} + for _, o := range opts { + o(&cfg) + } + + return connect(ctx, config, cfg) +} + +func Connect(ctx context.Context, cfg Config, opts ...ConnectOpt) (*DB, func() error, error) { + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + return nil, func() error { return nil }, fmt.Errorf("failed to parse config: %w", err) + } + + config.MaxConnIdleTime = time.Minute * 5 + config.MaxConnLifetime = time.Second * 30 + config.MaxConns = 20 + + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + logger.FatalKV(ctx, "failed to connect to postgres: %s", err.Error()) + } + + if err := pool.Ping(ctx); err != nil { + logger.FatalKV(ctx, "failed to ping postgres: %s", err.Error()) + } + + d := &DB{db: pool} + go d.asyncPing(ctx) + + return d +} + +func connect(ctx context.Context, pgconfig *pgxpool.Config, config Config) (*Database, func() error, error) { + +} + +func (d *Database) Ping(ctx context.Context) error { + return d.db.Ping(ctx) +} + +func (d *Database) asyncPing(ctx context.Context) { + dur := time.Second + t := time.NewTicker(dur) + defer t.Stop() + + for { + <-t.C + func() { + defer t.Reset(dur) + if err := d.Ping(ctx); err != nil { + logger.Error(ctx, "Database.asyncPing error", logger.Err(err)) + } + }() + } +} + +func (d *Database) RunWith(ctx context.Context) ports.DBTX { + if tx := txFromContext(ctx); tx != nil { + return tx + } + + return d.db +} + +type txKey struct{} + +var ctxKey txKey = txKey{} + +var _ ports.DBTX = (*pgx.Conn)(nil) +var _ ports.DBTX = func() pgx.Tx { return nil }() + +func (d *Database) WithTransaction(ctx context.Context, fn func(context.Context) error, opts pgx.TxOptions) (err error) { + tx := txFromContext(ctx) + if tx == nil { + tx, err = d.db.BeginTx(ctx, opts) + if err != nil { + return fmt.Errorf("failed to begin tx: %w", err) + } + + defer func() { + if err == nil { + err = tx.Commit(ctx) + } + if err != nil { + if rbErr := tx.Rollback(ctx); rbErr != nil { + err = errors.Join(err, rbErr) + } + } + }() + + ctx = txContext(ctx, tx) + } + + return fn(ctx) +} + +func Conn(ctx context.Context, db ports.DBTX) ports.DBTX { + if tx := txFromContext(ctx); tx != nil { + return tx + } + + return db +} + +func txFromContext(ctx context.Context) pgx.Tx { + if tx, ok := ctx.Value(ctxKey).(pgx.Tx); ok { + return tx + } + + return nil +} + +func txContext(parent context.Context, tx pgx.Tx) context.Context { + return context.WithValue(parent, ctxKey, tx) +} From 5632c3ba9783328538fbd81330c9bd016ccc06aa Mon Sep 17 00:00:00 2001 From: comicrime Date: Mon, 31 Aug 2026 14:49:44 +0300 Subject: [PATCH 2/4] added postgres wrapper --- README.md | 8 +- go.work | 1 + postgres/README.md | 51 ++++++++++++ postgres/interfaces.go | 6 -- postgres/postgres.go | 162 +++++++++++++++++++++++++++----------- postgres/postgres_test.go | 93 ++++++++++++++++++++++ postgres/txutils.go | 12 +++ 7 files changed, 273 insertions(+), 60 deletions(-) create mode 100644 postgres/README.md create mode 100644 postgres/postgres_test.go create mode 100644 postgres/txutils.go diff --git a/README.md b/README.md index 75d2da5..143b2c7 100644 --- a/README.md +++ b/README.md @@ -6,12 +6,8 @@ Modules that I constantly implement in my projects: and which I decided to colle | --- | --- | --- | | `caller` | `github.com/draincloud/callpack/caller` | `http.Client` fancy wrapper with a round-tripper middleware chain: headers, logging, etc | | `registry` | `github.com/draincloud/callpack/registry` | Registers a service instance in consul and heartbeats its TTL check. | +| `postgres` | `github.com/draincloud/callpack/postgres` | Easy to use postgres wrapper | -```sh -go get github.com/draincloud/callpack/caller -go get github.com/draincloud/callpack/registry -``` - -Each is tagged with its own prefix — `caller/v0.1.0`, `registry/v0.1.0` +Each is tagged with its own prefix — `caller/v0.1.0`, `registry/v0.1.0`, etc... `integration/` holds the tests that exercise the two against each other. \ No newline at end of file diff --git a/go.work b/go.work index d189fae..32bf2bd 100644 --- a/go.work +++ b/go.work @@ -5,5 +5,6 @@ use ( ./caller ./closer ./integration + ./postgres ./registry ) diff --git a/postgres/README.md b/postgres/README.md new file mode 100644 index 0000000..265c0fb --- /dev/null +++ b/postgres/README.md @@ -0,0 +1,51 @@ +# Postgres wrapper +Usage example: +main.go: +```go +func main() { + config := readConfig(ctx) // read from app config. + + db, closeDB, err := postgres.Connect(ctx, config) + if err != nil { + panic(err) + } + defer closeDB(ctx) + + data := getDataToFill(ctx) // get some data + + storage := storage.New(db) + + if err = storage.Backfill(ctx, data); err != nil { + panic(err) + } +} +``` +storage.go: +```go +type Storage struct { + db *postgres.DB +} + +func New(db *postgres.DB) *Storage { + return &Storage{db: db} +} + +func (s *Storage) Backfill(ctx context.Context, data map[int64]string) error { + return s.db.WithTransaction(ctx, func(ctx context.Context) error { + query := `update table set key = $1 where id = $2;` + for id, key := range data { + if _, err := s.db.RunWith(ctx).Exec(ctx, query, key, id); err != nil { + return err + } + } + return nil + }, pgx.TxOptions{}) +} +``` + +`closeDB` shuts the pool down and stops the background health check; it matches +`closer.CloseFunc`, so it can be handed to `closer.Add` directly. + +If WithTransaction called inside WithTransaction callback, it will reuse top-level transaction. +Passing non-zero `pgx.TxOptions` to a nested call returns `ErrNestedTxOptions`, since the +options cannot be applied to the transaction that is already running. diff --git a/postgres/interfaces.go b/postgres/interfaces.go index b6fac42..d351764 100644 --- a/postgres/interfaces.go +++ b/postgres/interfaces.go @@ -12,9 +12,3 @@ type DBTX interface { QueryRow(ctx context.Context, sql string, args ...any) pgx.Row Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) } - -type TXer interface { - BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) - Commit(ctx context.Context) error - Rollback(ctx context.Context) error -} diff --git a/postgres/postgres.go b/postgres/postgres.go index 3d0d36e..d0c8c10 100644 --- a/postgres/postgres.go +++ b/postgres/postgres.go @@ -4,6 +4,8 @@ import ( "context" "errors" "fmt" + "net" + "net/url" "time" "github.com/draincloud/logger" @@ -11,29 +13,54 @@ import ( "github.com/jackc/pgx/v5/pgxpool" ) +var ErrNestedTxOptions = errors.New("tx options cannot be applied to an already running transaction") + type Config struct { Host string Port string Username string Password string Database string - AllowSSL bool + SSLMode string MaxConns int MaxConnIdleTime time.Duration MaxConnLifetime time.Duration } +func (c Config) dsn() string { + host := c.Host + if c.Port != "" { + host = net.JoinHostPort(c.Host, c.Port) + } + + u := url.URL{ + Scheme: "postgres", + Host: host, + Path: "/" + c.Database, + } + + if c.Username != "" { + u.User = url.UserPassword(c.Username, c.Password) + } + + if c.SSLMode != "" { + u.RawQuery = url.Values{"sslmode": {c.SSLMode}}.Encode() + } + + return u.String() +} + type DB struct { db *pgxpool.Pool } type ConnectOpt func(c *Config) -func ConnectDSN(ctx context.Context, dsn string, opts ...ConnectOpt) (*DB, func() error, error) { - config, err := pgxpool.ParseConfig(dsn) +func ConnectDSN(ctx context.Context, dsn string, opts ...ConnectOpt) (*DB, func(context.Context) error, error) { + pgconfig, err := pgxpool.ParseConfig(dsn) if err != nil { - return nil, func() error { return nil }, fmt.Errorf("failed to parse config: %w", err) + return nil, noopCloser, fmt.Errorf("failed to parse config: %w", err) } cfg := Config{} @@ -41,59 +68,91 @@ func ConnectDSN(ctx context.Context, dsn string, opts ...ConnectOpt) (*DB, func( o(&cfg) } - return connect(ctx, config, cfg) + return connect(ctx, pgconfig, cfg) } -func Connect(ctx context.Context, cfg Config, opts ...ConnectOpt) (*DB, func() error, error) { - config, err := pgxpool.ParseConfig(dsn) +func Connect(ctx context.Context, cfg Config, opts ...ConnectOpt) (*DB, func(context.Context) error, error) { + for _, o := range opts { + o(&cfg) + } + + pgconfig, err := pgxpool.ParseConfig(cfg.dsn()) if err != nil { - return nil, func() error { return nil }, fmt.Errorf("failed to parse config: %w", err) + return nil, noopCloser, fmt.Errorf("failed to parse config: %w", err) + } + + return connect(ctx, pgconfig, cfg) +} + +func noopCloser(context.Context) error { return nil } + +func connect(ctx context.Context, pgconfig *pgxpool.Config, cfg Config) (*DB, func(context.Context) error, error) { + if cfg.MaxConns > 0 { + pgconfig.MaxConns = int32(cfg.MaxConns) + } + + if cfg.MaxConnIdleTime > 0 { + pgconfig.MaxConnIdleTime = cfg.MaxConnIdleTime } - config.MaxConnIdleTime = time.Minute * 5 - config.MaxConnLifetime = time.Second * 30 - config.MaxConns = 20 + if cfg.MaxConnLifetime > 0 { + pgconfig.MaxConnLifetime = cfg.MaxConnLifetime + } - pool, err := pgxpool.NewWithConfig(ctx, config) + pool, err := pgxpool.NewWithConfig(ctx, pgconfig) if err != nil { - logger.FatalKV(ctx, "failed to connect to postgres: %s", err.Error()) + return nil, noopCloser, fmt.Errorf("failed to connect to postgres: %w", err) } if err := pool.Ping(ctx); err != nil { - logger.FatalKV(ctx, "failed to ping postgres: %s", err.Error()) + pool.Close() + + return nil, noopCloser, fmt.Errorf("failed to ping postgres: %w", err) } - d := &DB{db: pool} - go d.asyncPing(ctx) + pingCtx, cancelPing := context.WithCancel(context.WithoutCancel(ctx)) - return d -} + d := &DB{db: pool} + go d.asyncPing(pingCtx) -func connect(ctx context.Context, pgconfig *pgxpool.Config, config Config) (*Database, func() error, error) { + return d, func(context.Context) error { + cancelPing() + pool.Close() + return nil + }, nil } -func (d *Database) Ping(ctx context.Context) error { - return d.db.Ping(ctx) +func (d *DB) Ping(ctx context.Context) error { + if err := d.db.Ping(ctx); err != nil { + return fmt.Errorf("failed to ping postgres: %w", err) + } + + return nil } -func (d *Database) asyncPing(ctx context.Context) { +func (d *DB) asyncPing(ctx context.Context) { dur := time.Second t := time.NewTicker(dur) defer t.Stop() for { - <-t.C + select { + case <-ctx.Done(): + return + case <-t.C: + } + func() { defer t.Reset(dur) if err := d.Ping(ctx); err != nil { - logger.Error(ctx, "Database.asyncPing error", logger.Err(err)) + logger.Error(ctx, "DB.asyncPing error", logger.Err(err)) } }() } } -func (d *Database) RunWith(ctx context.Context) ports.DBTX { +func (d *DB) RunWith(ctx context.Context) DBTX { if tx := txFromContext(ctx); tx != nil { return tx } @@ -101,39 +160,46 @@ func (d *Database) RunWith(ctx context.Context) ports.DBTX { return d.db } -type txKey struct{} +func (d *DB) WithTransaction(ctx context.Context, fn func(context.Context) error, opts pgx.TxOptions) (err error) { + if tx := txFromContext(ctx); tx != nil { + if opts != (pgx.TxOptions{}) { + return ErrNestedTxOptions + } -var ctxKey txKey = txKey{} + return fn(ctx) + } -var _ ports.DBTX = (*pgx.Conn)(nil) -var _ ports.DBTX = func() pgx.Tx { return nil }() + tx, err := d.db.BeginTx(ctx, opts) + if err != nil { + return fmt.Errorf("failed to begin tx: %w", err) + } -func (d *Database) WithTransaction(ctx context.Context, fn func(context.Context) error, opts pgx.TxOptions) (err error) { - tx := txFromContext(ctx) - if tx == nil { - tx, err = d.db.BeginTx(ctx, opts) - if err != nil { - return fmt.Errorf("failed to begin tx: %w", err) + defer func() { + closeCtx := context.WithoutCancel(ctx) + + if p := recover(); p != nil { + _ = tx.Rollback(closeCtx) + + panic(p) } - defer func() { - if err == nil { - err = tx.Commit(ctx) - } - if err != nil { - if rbErr := tx.Rollback(ctx); rbErr != nil { - err = errors.Join(err, rbErr) - } + if err != nil { + if rbErr := tx.Rollback(closeCtx); rbErr != nil { + err = errors.Join(err, rbErr) } - }() - ctx = txContext(ctx, tx) - } + return + } + + if cErr := tx.Commit(closeCtx); cErr != nil { + err = fmt.Errorf("failed to commit tx: %w", cErr) + } + }() - return fn(ctx) + return fn(txContext(ctx, tx)) } -func Conn(ctx context.Context, db ports.DBTX) ports.DBTX { +func Conn(ctx context.Context, db DBTX) DBTX { if tx := txFromContext(ctx); tx != nil { return tx } diff --git a/postgres/postgres_test.go b/postgres/postgres_test.go new file mode 100644 index 0000000..47a82bb --- /dev/null +++ b/postgres/postgres_test.go @@ -0,0 +1,93 @@ +package postgres + +import ( + "context" + "errors" + "testing" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestConfigDSN(t *testing.T) { + tests := []struct { + name string + cfg Config + want string + }{ + { + name: "full", + cfg: Config{Host: "db.local", Port: "5432", Username: "us er", Password: "p@ss:/word", Database: "app", SSLMode: "verify-full"}, + want: "postgres://us%20er:p%40ss%3A%2Fword@db.local:5432/app?sslmode=verify-full", + }, + { + name: "no ssl mode leaves the pgx default", + cfg: Config{Host: "localhost", Database: "app"}, + want: "postgres://localhost/app", + }, + { + name: "no port", + cfg: Config{Host: "localhost", Username: "u", Password: "p", Database: "app"}, + want: "postgres://u:p@localhost/app", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := tt.cfg.dsn(); got != tt.want { + t.Fatalf("dsn() = %q, want %q", got, tt.want) + } + + if _, err := pgxpool.ParseConfig(tt.cfg.dsn()); err != nil { + t.Fatalf("ParseConfig(%q) = %v", tt.cfg.dsn(), err) + } + }) + } +} + +func TestConfigDSNDefaultsToPrefer(t *testing.T) { + pgconfig, err := pgxpool.ParseConfig(Config{Host: "localhost", Database: "app"}.dsn()) + if err != nil { + t.Fatal(err) + } + + if pgconfig.ConnConfig.TLSConfig == nil { + t.Fatal("expected TLS to remain available when SSLMode is unset") + } +} + +type stubTx struct{ pgx.Tx } + +func TestWithTransactionNestedRejectsOptions(t *testing.T) { + ctx := txContext(context.Background(), stubTx{}) + + err := (&DB{}).WithTransaction(ctx, func(context.Context) error { + t.Fatal("callback must not run when options conflict") + + return nil + }, pgx.TxOptions{IsoLevel: pgx.Serializable}) + + if !errors.Is(err, ErrNestedTxOptions) { + t.Fatalf("err = %v, want ErrNestedTxOptions", err) + } +} + +func TestWithTransactionNestedReusesOuterTx(t *testing.T) { + tx := stubTx{} + ctx := txContext(context.Background(), tx) + + var got pgx.Tx + + err := (&DB{}).WithTransaction(ctx, func(ctx context.Context) error { + got = txFromContext(ctx) + + return nil + }, pgx.TxOptions{}) + if err != nil { + t.Fatalf("err = %v, want nil", err) + } + + if got != pgx.Tx(tx) { + t.Fatal("nested call did not reuse the outer transaction") + } +} diff --git a/postgres/txutils.go b/postgres/txutils.go new file mode 100644 index 0000000..8da55c3 --- /dev/null +++ b/postgres/txutils.go @@ -0,0 +1,12 @@ +package postgres + +import ( + "github.com/jackc/pgx/v5" +) + +type txKey struct{} + +var ctxKey txKey = txKey{} + +var _ DBTX = (*pgx.Conn)(nil) +var _ DBTX = func() pgx.Tx { return nil }() From 0a41fd5b6619ec2dab5ed4c75ebdc4a6cc42ce7a Mon Sep 17 00:00:00 2001 From: comicrime Date: Mon, 31 Aug 2026 14:57:44 +0300 Subject: [PATCH 3/4] added tests --- postgres/postgres_test.go | 201 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 201 insertions(+) diff --git a/postgres/postgres_test.go b/postgres/postgres_test.go index 47a82bb..e5d4ba9 100644 --- a/postgres/postgres_test.go +++ b/postgres/postgres_test.go @@ -3,7 +3,9 @@ package postgres import ( "context" "errors" + "strings" "testing" + "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" @@ -58,6 +60,8 @@ func TestConfigDSNDefaultsToPrefer(t *testing.T) { type stubTx struct{ pgx.Tx } +type stubDBTX struct{ DBTX } + func TestWithTransactionNestedRejectsOptions(t *testing.T) { ctx := txContext(context.Background(), stubTx{}) @@ -91,3 +95,200 @@ func TestWithTransactionNestedReusesOuterTx(t *testing.T) { t.Fatal("nested call did not reuse the outer transaction") } } + +func TestConnectAppliesOptsBeforeBuildingDSN(t *testing.T) { + cfg := Config{Host: "127.0.0.1", Port: "1", Database: "app"} + + db, closeDB, err := Connect(t.Context(), cfg, func(c *Config) { c.SSLMode = "bogus" }) + if err == nil { + closeDB(t.Context()) + t.Fatal("expected the option to reach the DSN and fail to parse") + } + + if db != nil { + t.Fatalf("db = %v, want nil", db) + } + + if err := closeDB(t.Context()); err != nil { + t.Fatalf("closer returned on failure: %v", err) + } +} + +func TestConnectUnreachableHost(t *testing.T) { + db, closeDB, err := Connect(t.Context(), Config{Host: "127.0.0.1", Port: "1", Database: "app", SSLMode: "disable"}) + if err == nil { + closeDB(t.Context()) + t.Fatal("expected the ping to fail") + } + + if db != nil { + t.Fatalf("db = %v, want nil", db) + } + + if err := closeDB(t.Context()); err != nil { + t.Fatalf("closer returned on failure: %v", err) + } +} + +func TestConnectDSNRejectsUnparseableDSN(t *testing.T) { + db, closeDB, err := ConnectDSN(t.Context(), "postgres://host:port/app") + if err == nil { + closeDB(t.Context()) + t.Fatal("expected an unparseable DSN to fail") + } + + if db != nil { + t.Fatalf("db = %v, want nil", db) + } + + if err := closeDB(t.Context()); err != nil { + t.Fatalf("closer returned on failure: %v", err) + } +} + +func TestConnectAppliesPoolTuning(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{MaxConns: 7, MaxConnIdleTime: time.Minute, MaxConnLifetime: time.Hour} + if _, closeDB, err := connect(t.Context(), pgconfig, cfg); err == nil { + closeDB(t.Context()) + t.Fatal("expected the ping to fail") + } + + if pgconfig.MaxConns != int32(cfg.MaxConns) { + t.Fatalf("MaxConns = %d, want %d", pgconfig.MaxConns, cfg.MaxConns) + } + + if pgconfig.MaxConnIdleTime != cfg.MaxConnIdleTime { + t.Fatalf("MaxConnIdleTime = %s, want %s", pgconfig.MaxConnIdleTime, cfg.MaxConnIdleTime) + } + + if pgconfig.MaxConnLifetime != cfg.MaxConnLifetime { + t.Fatalf("MaxConnLifetime = %s, want %s", pgconfig.MaxConnLifetime, cfg.MaxConnLifetime) + } +} + +func TestConnectLeavesPgxDefaultsWhenUntuned(t *testing.T) { + pgconfig, err := pgxpool.ParseConfig(Config{Host: "127.0.0.1", Port: "1", Database: "app"}.dsn()) + if err != nil { + t.Fatal(err) + } + + want := *pgconfig + + if _, closeDB, err := connect(t.Context(), pgconfig, Config{}); err == nil { + closeDB(t.Context()) + t.Fatal("expected the ping to fail") + } + + if pgconfig.MaxConns != want.MaxConns || + pgconfig.MaxConnIdleTime != want.MaxConnIdleTime || + pgconfig.MaxConnLifetime != want.MaxConnLifetime { + t.Fatalf("zero-value Config overwrote the pgx defaults: %+v", pgconfig) + } +} + +func TestRunWith(t *testing.T) { + db := &DB{} + + got := db.RunWith(t.Context()) + if pool, ok := got.(*pgxpool.Pool); !ok || pool != db.db { + t.Fatalf("RunWith without a transaction = %#v, want the pool", got) + } + + tx := stubTx{} + if got := db.RunWith(txContext(t.Context(), tx)); got != DBTX(tx) { + t.Fatalf("RunWith inside a transaction = %#v, want the transaction", got) + } +} + +func TestConn(t *testing.T) { + db := stubDBTX{} + + if got := Conn(t.Context(), db); got != DBTX(db) { + t.Fatalf("Conn without a transaction = %#v, want the given DBTX", got) + } + + tx := stubTx{} + if got := Conn(txContext(t.Context(), tx), db); got != DBTX(tx) { + t.Fatalf("Conn inside a transaction = %#v, want the transaction", got) + } +} + +func TestWithTransactionNestedPropagatesError(t *testing.T) { + errCallback := errors.New("callback failed") + ctx := txContext(t.Context(), stubTx{}) + + err := (&DB{}).WithTransaction(ctx, func(context.Context) error { + return errCallback + }, pgx.TxOptions{}) + + if !errors.Is(err, errCallback) { + t.Fatalf("err = %v, want %v", err, errCallback) + } +} + +func TestConnectDSNAppliesOpts(t *testing.T) { + applied := false + + db, closeDB, err := ConnectDSN(t.Context(), "postgres://127.0.0.1:1/app?sslmode=disable", func(c *Config) { + applied = true + c.MaxConns = 3 + }) + if err == nil { + closeDB(t.Context()) + t.Fatal("expected the ping to fail") + } + + if db != nil { + t.Fatalf("db = %v, want nil", db) + } + + if !applied { + t.Fatal("options are not applied on the DSN path") + } +} + +func newLazyPool(t *testing.T) *pgxpool.Pool { + t.Helper() + + pool, err := pgxpool.New(t.Context(), "postgres://127.0.0.1:1/app?sslmode=disable") + if err != nil { + t.Fatal(err) + } + + t.Cleanup(pool.Close) + + return pool +} + +func TestPingWrapsError(t *testing.T) { + err := (&DB{db: newLazyPool(t)}).Ping(t.Context()) + if err == nil { + t.Fatal("expected the ping to fail") + } + + if !strings.Contains(err.Error(), "failed to ping postgres") { + t.Fatalf("err = %v, want it named by this package", err) + } +} + +func TestAsyncPingStopsWhenContextIsCancelled(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + done := make(chan struct{}) + go func() { + defer close(done) + (&DB{db: newLazyPool(t)}).asyncPing(ctx) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("asyncPing outlived its context") + } +} From 2aa6ecc3700fe693ca001d292aa8fa3b72b145d6 Mon Sep 17 00:00:00 2001 From: comicrime Date: Sun, 6 Sep 2026 00:30:47 +0300 Subject: [PATCH 4/4] +1 --- postgres/README.md | 13 +++ postgres/go.mod | 2 +- postgres/go.sum | 4 +- postgres/postgres.go | 70 ++++++++++++++- postgres/postgres_test.go | 185 +++++++++++++++++++++++++++++++++++++- 5 files changed, 267 insertions(+), 7 deletions(-) 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") + } +}