From 4852ca9ce6e76d4168bf5a20d97206caad50a972 Mon Sep 17 00:00:00 2001 From: zhuyh1606-oss Date: Fri, 24 Jul 2026 16:01:37 +0800 Subject: [PATCH] Fix deadlock in pgxpool under concurrent context cancellation Fixes concurrent context cancellation + semaphore contention causing pool to livelock. Two-part fix: 1. puddle/pool.go acquire(): check ctx.Done() after semaphore acquire to prevent goroutine leak when context is canceled between TryAcquire success and resource creation. 2. pgxpool/pool.go Acquire(): check ctx.Err() between retry iterations so a cancelled context terminates the retry loop immediately instead of cycling through all maxConns attempts. Includes regression test TestPoolAcquireStressCancelContention that exercises high-concurrency acquire with rapid context cancellation. --- go.mod | 22 + go.sum | 39 + pgxpool/batch_results.go | 52 + pgxpool/bench_test.go | 81 + pgxpool/common_test.go | 203 +++ pgxpool/conn.go | 133 ++ pgxpool/conn_test.go | 96 ++ pgxpool/doc.go | 27 + pgxpool/helper_test.go | 39 + pgxpool/pool.go | 852 +++++++++++ pgxpool/pool_test.go | 1445 ++++++++++++++++++ pgxpool/rows.go | 116 ++ pgxpool/stat.go | 91 ++ pgxpool/tracer.go | 33 + pgxpool/tracer_test.go | 130 ++ pgxpool/tx.go | 83 ++ pgxpool/tx_test.go | 96 ++ puddle/.devcontainer/devcontainer.json | 22 + puddle/CHANGELOG.md | 79 + puddle/LICENSE | 22 + puddle/README.md | 96 ++ puddle/context.go | 24 + puddle/doc.go | 11 + puddle/export_test.go | 9 + puddle/go.mod | 14 + puddle/go.sum | 19 + puddle/internal/genstack/gen_stack.go | 85 ++ puddle/internal/genstack/gen_stack_test.go | 90 ++ puddle/internal/genstack/stack.go | 39 + puddle/nanotime.go | 16 + puddle/pool.go | 728 +++++++++ puddle/pool_test.go | 1576 ++++++++++++++++++++ puddle/resource_list.go | 28 + puddle/resource_list_test.go | 62 + 34 files changed, 6458 insertions(+) create mode 100644 go.mod create mode 100644 go.sum create mode 100644 pgxpool/batch_results.go create mode 100644 pgxpool/bench_test.go create mode 100644 pgxpool/common_test.go create mode 100644 pgxpool/conn.go create mode 100644 pgxpool/conn_test.go create mode 100644 pgxpool/doc.go create mode 100644 pgxpool/helper_test.go create mode 100644 pgxpool/pool.go create mode 100644 pgxpool/pool_test.go create mode 100644 pgxpool/rows.go create mode 100644 pgxpool/stat.go create mode 100644 pgxpool/tracer.go create mode 100644 pgxpool/tracer_test.go create mode 100644 pgxpool/tx.go create mode 100644 pgxpool/tx_test.go create mode 100644 puddle/.devcontainer/devcontainer.json create mode 100644 puddle/CHANGELOG.md create mode 100644 puddle/LICENSE create mode 100644 puddle/README.md create mode 100644 puddle/context.go create mode 100644 puddle/doc.go create mode 100644 puddle/export_test.go create mode 100644 puddle/go.mod create mode 100644 puddle/go.sum create mode 100644 puddle/internal/genstack/gen_stack.go create mode 100644 puddle/internal/genstack/gen_stack_test.go create mode 100644 puddle/internal/genstack/stack.go create mode 100644 puddle/nanotime.go create mode 100644 puddle/pool.go create mode 100644 puddle/pool_test.go create mode 100644 puddle/resource_list.go create mode 100644 puddle/resource_list_test.go diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..c1d2a02 --- /dev/null +++ b/go.mod @@ -0,0 +1,22 @@ +module github.com/jackc/pgx/v5 + +go 1.25.0 + +require ( + github.com/jackc/pgpassfile v1.0.0 + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 + github.com/jackc/puddle/v2 v2.2.2 + github.com/stretchr/testify v1.11.1 + golang.org/x/sync v0.17.0 + golang.org/x/text v0.29.0 +) + +replace github.com/jackc/puddle/v2 => ./puddle + +require ( + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/kr/pretty v0.3.0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..1d17aab --- /dev/null +++ b/go.sum @@ -0,0 +1,39 @@ +github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= +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/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/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/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= +github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= +github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk= +github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= +github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +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/rogpeppe/go-internal v1.6.1 h1:/FiVV8dS/e+YqF2JvO3yXRFbBLTIuSDkuC7aBOAvL+k= +github.com/rogpeppe/go-internal v1.6.1/go.mod h1:xXDCJY+GAPziupqXw64V24skbSoqbTEfhy4qGm1nDQc= +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/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/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI= +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/pgxpool/batch_results.go b/pgxpool/batch_results.go new file mode 100644 index 0000000..5d5c681 --- /dev/null +++ b/pgxpool/batch_results.go @@ -0,0 +1,52 @@ +package pgxpool + +import ( + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" +) + +type errBatchResults struct { + err error +} + +func (br errBatchResults) Exec() (pgconn.CommandTag, error) { + return pgconn.CommandTag{}, br.err +} + +func (br errBatchResults) Query() (pgx.Rows, error) { + return errRows{err: br.err}, br.err +} + +func (br errBatchResults) QueryRow() pgx.Row { + return errRow{err: br.err} +} + +func (br errBatchResults) Close() error { + return br.err +} + +type poolBatchResults struct { + br pgx.BatchResults + c *Conn +} + +func (br *poolBatchResults) Exec() (pgconn.CommandTag, error) { + return br.br.Exec() +} + +func (br *poolBatchResults) Query() (pgx.Rows, error) { + return br.br.Query() +} + +func (br *poolBatchResults) QueryRow() pgx.Row { + return br.br.QueryRow() +} + +func (br *poolBatchResults) Close() error { + err := br.br.Close() + if br.c != nil { + br.c.Release() + br.c = nil + } + return err +} diff --git a/pgxpool/bench_test.go b/pgxpool/bench_test.go new file mode 100644 index 0000000..2748ddb --- /dev/null +++ b/pgxpool/bench_test.go @@ -0,0 +1,81 @@ +package pgxpool_test + +import ( + "context" + "os" + "testing" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +func BenchmarkAcquireAndRelease(b *testing.B) { + pool, err := pgxpool.New(context.Background(), os.Getenv("PGX_TEST_DATABASE")) + require.NoError(b, err) + defer pool.Close() + + for b.Loop() { + c, err := pool.Acquire(context.Background()) + if err != nil { + b.Fatal(err) + } + c.Release() + } +} + +func BenchmarkMinimalPreparedSelectBaseline(b *testing.B) { + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(b, err) + + config.AfterConnect = func(ctx context.Context, c *pgx.Conn) error { + _, err := c.Prepare(ctx, "ps1", "select $1::int8") + return err + } + + db, err := pgxpool.NewWithConfig(context.Background(), config) + require.NoError(b, err) + + conn, err := db.Acquire(context.Background()) + require.NoError(b, err) + defer conn.Release() + + var n int64 + + for i := 0; b.Loop(); i++ { + err = conn.QueryRow(context.Background(), "ps1", i).Scan(&n) + if err != nil { + b.Fatal(err) + } + + if n != int64(i) { + b.Fatalf("expected %d, got %d", i, n) + } + } +} + +func BenchmarkMinimalPreparedSelect(b *testing.B) { + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(b, err) + + config.AfterConnect = func(ctx context.Context, c *pgx.Conn) error { + _, err := c.Prepare(ctx, "ps1", "select $1::int8") + return err + } + + db, err := pgxpool.NewWithConfig(context.Background(), config) + require.NoError(b, err) + + var n int64 + + for i := 0; b.Loop(); i++ { + err = db.QueryRow(context.Background(), "ps1", i).Scan(&n) + if err != nil { + b.Fatal(err) + } + + if n != int64(i) { + b.Fatalf("expected %d, got %d", i, n) + } + } +} diff --git a/pgxpool/common_test.go b/pgxpool/common_test.go new file mode 100644 index 0000000..20ce808 --- /dev/null +++ b/pgxpool/common_test.go @@ -0,0 +1,203 @@ +package pgxpool_test + +import ( + "context" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// Conn.Release is an asynchronous process that returns immediately. There is no signal when the actual work is +// completed. To test something that relies on the actual work for Conn.Release being completed we must simply wait. +// This function wraps the sleep so there is more meaning for the callers. +func waitForReleaseToComplete() { + time.Sleep(500 * time.Millisecond) +} + +type execer interface { + Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) +} + +func testExec(t *testing.T, ctx context.Context, db execer) { + results, err := db.Exec(ctx, "set time zone 'America/Chicago'") + require.NoError(t, err) + assert.EqualValues(t, "SET", results.String()) +} + +type queryer interface { + Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) +} + +func testQuery(t *testing.T, ctx context.Context, db queryer) { + var sum, rowCount int32 + + rows, err := db.Query(ctx, "select generate_series(1,$1)", 10) + require.NoError(t, err) + + for rows.Next() { + var n int32 + rows.Scan(&n) + sum += n + rowCount++ + } + + assert.NoError(t, rows.Err()) + assert.Equal(t, int32(10), rowCount) + assert.Equal(t, int32(55), sum) +} + +type queryRower interface { + QueryRow(ctx context.Context, sql string, args ...any) pgx.Row +} + +func testQueryRow(t *testing.T, ctx context.Context, db queryRower) { + var what, who string + err := db.QueryRow(ctx, "select 'hello', $1::text", "world").Scan(&what, &who) + assert.NoError(t, err) + assert.Equal(t, "hello", what) + assert.Equal(t, "world", who) +} + +type sendBatcher interface { + SendBatch(context.Context, *pgx.Batch) pgx.BatchResults +} + +func testSendBatch(t *testing.T, ctx context.Context, db sendBatcher) { + batch := &pgx.Batch{} + batch.Queue("select 1") + batch.Queue("select 2") + + br := db.SendBatch(ctx, batch) + + var err error + var n int32 + err = br.QueryRow().Scan(&n) + assert.NoError(t, err) + assert.EqualValues(t, 1, n) + + err = br.QueryRow().Scan(&n) + assert.NoError(t, err) + assert.EqualValues(t, 2, n) + + err = br.Close() + assert.NoError(t, err) +} + +type copyFromer interface { + CopyFrom(context.Context, pgx.Identifier, []string, pgx.CopyFromSource) (int64, error) +} + +func testCopyFrom(t *testing.T, ctx context.Context, db interface { + execer + queryer + copyFromer +}, +) { + _, err := db.Exec(ctx, `create temporary table foo(a int2, b int4, c int8, d varchar, e text, f date, g timestamptz)`) + require.NoError(t, err) + + tzedTime := time.Date(2010, 2, 3, 4, 5, 6, 0, time.Local) + + inputRows := [][]any{ + {int16(0), int32(1), int64(2), "abc", "efg", time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC), tzedTime}, + {nil, nil, nil, nil, nil, nil, nil}, + } + + copyCount, err := db.CopyFrom(ctx, pgx.Identifier{"foo"}, []string{"a", "b", "c", "d", "e", "f", "g"}, pgx.CopyFromRows(inputRows)) + assert.NoError(t, err) + assert.EqualValues(t, len(inputRows), copyCount) + + rows, err := db.Query(ctx, "select * from foo") + assert.NoError(t, err) + + var outputRows [][]any + for rows.Next() { + row, err := rows.Values() + if err != nil { + t.Errorf("Unexpected error for rows.Values(): %v", err) + } + outputRows = append(outputRows, row) + } + + assert.NoError(t, rows.Err()) + assert.Equal(t, inputRows, outputRows) +} + +func assertConfigsEqual(t *testing.T, expected, actual *pgxpool.Config, testName string) { + if !assert.NotNil(t, expected) { + return + } + if !assert.NotNil(t, actual) { + return + } + + assert.Equalf(t, expected.ConnString(), actual.ConnString(), "%s - ConnString", testName) + + // Can't test function equality, so just test that they are set or not. + assert.Equalf(t, expected.AfterConnect == nil, actual.AfterConnect == nil, "%s - AfterConnect", testName) + assert.Equalf(t, expected.BeforeAcquire == nil, actual.BeforeAcquire == nil, "%s - BeforeAcquire", testName) + assert.Equalf(t, expected.PrepareConn == nil, actual.PrepareConn == nil, "%s - PrepareConn", testName) + assert.Equalf(t, expected.AfterRelease == nil, actual.AfterRelease == nil, "%s - AfterRelease", testName) + + assert.Equalf(t, expected.MaxConnLifetime, actual.MaxConnLifetime, "%s - MaxConnLifetime", testName) + assert.Equalf(t, expected.MaxConnIdleTime, actual.MaxConnIdleTime, "%s - MaxConnIdleTime", testName) + assert.Equalf(t, expected.MaxConns, actual.MaxConns, "%s - MaxConns", testName) + assert.Equalf(t, expected.MinConns, actual.MinConns, "%s - MinConns", testName) + assert.Equalf(t, expected.MinIdleConns, actual.MinIdleConns, "%s - MinIdleConns", testName) + assert.Equalf(t, expected.HealthCheckPeriod, actual.HealthCheckPeriod, "%s - HealthCheckPeriod", testName) + + assertConnConfigsEqual(t, expected.ConnConfig, actual.ConnConfig, testName) +} + +func assertConnConfigsEqual(t *testing.T, expected, actual *pgx.ConnConfig, testName string) { + if !assert.NotNil(t, expected) { + return + } + if !assert.NotNil(t, actual) { + return + } + + assert.Equalf(t, expected.Tracer, actual.Tracer, "%s - Tracer", testName) + assert.Equalf(t, expected.ConnString(), actual.ConnString(), "%s - ConnString", testName) + assert.Equalf(t, expected.StatementCacheCapacity, actual.StatementCacheCapacity, "%s - StatementCacheCapacity", testName) + assert.Equalf(t, expected.DescriptionCacheCapacity, actual.DescriptionCacheCapacity, "%s - DescriptionCacheCapacity", testName) + assert.Equalf(t, expected.DefaultQueryExecMode, actual.DefaultQueryExecMode, "%s - DefaultQueryExecMode", testName) + assert.Equalf(t, expected.Host, actual.Host, "%s - Host", testName) + assert.Equalf(t, expected.Database, actual.Database, "%s - Database", testName) + assert.Equalf(t, expected.Port, actual.Port, "%s - Port", testName) + assert.Equalf(t, expected.User, actual.User, "%s - User", testName) + assert.Equalf(t, expected.Password, actual.Password, "%s - Password", testName) + assert.Equalf(t, expected.ConnectTimeout, actual.ConnectTimeout, "%s - ConnectTimeout", testName) + assert.Equalf(t, expected.RuntimeParams, actual.RuntimeParams, "%s - RuntimeParams", testName) + + // Can't test function equality, so just test that they are set or not. + assert.Equalf(t, expected.ValidateConnect == nil, actual.ValidateConnect == nil, "%s - ValidateConnect", testName) + assert.Equalf(t, expected.AfterConnect == nil, actual.AfterConnect == nil, "%s - AfterConnect", testName) + + if assert.Equalf(t, expected.TLSConfig == nil, actual.TLSConfig == nil, "%s - TLSConfig", testName) { + if expected.TLSConfig != nil { + assert.Equalf(t, expected.TLSConfig.InsecureSkipVerify, actual.TLSConfig.InsecureSkipVerify, "%s - TLSConfig InsecureSkipVerify", testName) + assert.Equalf(t, expected.TLSConfig.ServerName, actual.TLSConfig.ServerName, "%s - TLSConfig ServerName", testName) + } + } + + if assert.Equalf(t, len(expected.Fallbacks), len(actual.Fallbacks), "%s - Fallbacks", testName) { + for i := range expected.Fallbacks { + assert.Equalf(t, expected.Fallbacks[i].Host, actual.Fallbacks[i].Host, "%s - Fallback %d - Host", testName, i) + assert.Equalf(t, expected.Fallbacks[i].Port, actual.Fallbacks[i].Port, "%s - Fallback %d - Port", testName, i) + + if assert.Equalf(t, expected.Fallbacks[i].TLSConfig == nil, actual.Fallbacks[i].TLSConfig == nil, "%s - Fallback %d - TLSConfig", testName, i) { + if expected.Fallbacks[i].TLSConfig != nil { + assert.Equalf(t, expected.Fallbacks[i].TLSConfig.InsecureSkipVerify, actual.Fallbacks[i].TLSConfig.InsecureSkipVerify, "%s - Fallback %d - TLSConfig InsecureSkipVerify", testName) + assert.Equalf(t, expected.Fallbacks[i].TLSConfig.ServerName, actual.Fallbacks[i].TLSConfig.ServerName, "%s - Fallback %d - TLSConfig ServerName", testName) + } + } + } + } +} diff --git a/pgxpool/conn.go b/pgxpool/conn.go new file mode 100644 index 0000000..b4f9060 --- /dev/null +++ b/pgxpool/conn.go @@ -0,0 +1,133 @@ +package pgxpool + +import ( + "context" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/puddle/v2" +) + +// Conn is an acquired *pgx.Conn from a Pool. +type Conn struct { + res *puddle.Resource[*connResource] + p *Pool +} + +// Release returns c to the pool it was acquired from. Once Release has been called, other methods must not be called. +// However, it is safe to call Release multiple times. Subsequent calls after the first will be ignored. +func (c *Conn) Release() { + if c.res == nil { + return + } + + conn := c.Conn() + res := c.res + c.res = nil + + if c.p.releaseTracer != nil { + c.p.releaseTracer.TraceRelease(c.p, TraceReleaseData{Conn: conn}) + } + + if conn.IsClosed() || conn.PgConn().IsBusy() || conn.PgConn().TxStatus() != 'I' { + res.Destroy() + // Signal to the health check to run since we just destroyed a connections + // and we might be below minConns now + c.p.triggerHealthCheck() + return + } + + // If the pool is consistently being used, we might never get to check the + // lifetime of a connection since we only check idle connections in checkConnsHealth + // so we also check the lifetime here and force a health check + if c.p.isExpired(res) { + c.p.lifetimeDestroyCount.Add(1) + res.Destroy() + // Signal to the health check to run since we just destroyed a connections + // and we might be below minConns now + c.p.triggerHealthCheck() + return + } + + if c.p.afterRelease == nil { + res.Release() + return + } + + go func() { + if c.p.afterRelease(conn) { + res.Release() + } else { + res.Destroy() + // Signal to the health check to run since we just destroyed a connections + // and we might be below minConns now + c.p.triggerHealthCheck() + } + }() +} + +// Hijack assumes ownership of the connection from the pool. Caller is responsible for closing the connection. Hijack +// will panic if called on an already released or hijacked connection. +func (c *Conn) Hijack() *pgx.Conn { + if c.res == nil { + panic("cannot hijack already released or hijacked connection") + } + + conn := c.Conn() + res := c.res + c.res = nil + + res.Hijack() + + return conn +} + +func (c *Conn) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) { + return c.Conn().Exec(ctx, sql, arguments...) +} + +func (c *Conn) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) { + return c.Conn().Query(ctx, sql, args...) +} + +func (c *Conn) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row { + return c.Conn().QueryRow(ctx, sql, args...) +} + +func (c *Conn) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults { + return c.Conn().SendBatch(ctx, b) +} + +func (c *Conn) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) { + return c.Conn().CopyFrom(ctx, tableName, columnNames, rowSrc) +} + +// Begin starts a transaction block from the *Conn without explicitly setting a transaction mode (see BeginTx with TxOptions if transaction mode is required). +func (c *Conn) Begin(ctx context.Context) (pgx.Tx, error) { + return c.Conn().Begin(ctx) +} + +// BeginTx starts a transaction block from the *Conn with txOptions determining the transaction mode. +func (c *Conn) BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) { + return c.Conn().BeginTx(ctx, txOptions) +} + +func (c *Conn) Ping(ctx context.Context) error { + return c.Conn().Ping(ctx) +} + +func (c *Conn) Conn() *pgx.Conn { + return c.connResource().conn +} + +func (c *Conn) connResource() *connResource { + return c.res.Value() +} + +func (c *Conn) getPoolRow(r pgx.Row) *poolRow { + return c.connResource().getPoolRow(c, r) +} + +func (c *Conn) getPoolRows(r pgx.Rows) *poolRows { + return c.connResource().getPoolRows(c, r) +} diff --git a/pgxpool/conn_test.go b/pgxpool/conn_test.go new file mode 100644 index 0000000..ce35c49 --- /dev/null +++ b/pgxpool/conn_test.go @@ -0,0 +1,96 @@ +package pgxpool_test + +import ( + "context" + "os" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +func TestConnExec(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + + testExec(t, ctx, c) +} + +func TestConnQuery(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + + testQuery(t, ctx, c) +} + +func TestConnQueryRow(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + + testQueryRow(t, ctx, c) +} + +func TestConnSendBatch(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + + testSendBatch(t, ctx, c) +} + +func TestConnCopyFrom(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + + testCopyFrom(t, ctx, c) +} diff --git a/pgxpool/doc.go b/pgxpool/doc.go new file mode 100644 index 0000000..099443b --- /dev/null +++ b/pgxpool/doc.go @@ -0,0 +1,27 @@ +// Package pgxpool is a concurrency-safe connection pool for pgx. +/* +pgxpool implements a nearly identical interface to pgx connections. + +Creating a Pool + +The primary way of creating a pool is with [pgxpool.New]: + + pool, err := pgxpool.New(context.Background(), os.Getenv("DATABASE_URL")) + +The database connection string can be in URL or keyword/value format. PostgreSQL settings, pgx settings, and pool settings can be +specified here. In addition, a config struct can be created by [ParseConfig]. + + config, err := pgxpool.ParseConfig(os.Getenv("DATABASE_URL")) + if err != nil { + // ... + } + config.AfterConnect = func(ctx context.Context, conn *pgx.Conn) error { + // do something with every new connection + } + + pool, err := pgxpool.NewWithConfig(context.Background(), config) + +A pool returns without waiting for any connections to be established. Acquire a connection immediately after creating +the pool to check if a connection can successfully be established. +*/ +package pgxpool diff --git a/pgxpool/helper_test.go b/pgxpool/helper_test.go new file mode 100644 index 0000000..7d63732 --- /dev/null +++ b/pgxpool/helper_test.go @@ -0,0 +1,39 @@ +package pgxpool_test + +import ( + "context" + "net" + "time" + + "github.com/jackc/pgx/v5/pgconn" +) + +// delayProxy is a that introduces a configurable delay on reads from the database connection. +type delayProxy struct { + net.Conn + readDelay time.Duration +} + +func newDelayProxy(conn net.Conn, readDelay time.Duration) *delayProxy { + p := &delayProxy{ + Conn: conn, + readDelay: readDelay, + } + + return p +} + +func (dp *delayProxy) Read(b []byte) (int, error) { + if dp.readDelay > 0 { + time.Sleep(dp.readDelay) + } + + return dp.Conn.Read(b) +} + +func newDelayProxyDialFunc(readDelay time.Duration) pgconn.DialFunc { + return func(ctx context.Context, network, addr string) (net.Conn, error) { + conn, err := net.Dial(network, addr) + return newDelayProxy(conn, readDelay), err + } +} diff --git a/pgxpool/pool.go b/pgxpool/pool.go new file mode 100644 index 0000000..34f46f6 --- /dev/null +++ b/pgxpool/pool.go @@ -0,0 +1,852 @@ +package pgxpool + +import ( + "context" + "errors" + "math/rand/v2" + "runtime" + "strconv" + "sync" + "sync/atomic" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/puddle/v2" +) + +var ( + defaultMaxConns = int32(4) + defaultMinConns = int32(0) + defaultMinIdleConns = int32(0) + defaultMaxConnLifetime = time.Hour + defaultMaxConnIdleTime = time.Minute * 30 + defaultHealthCheckPeriod = time.Minute +) + +type connResource struct { + conn *pgx.Conn + conns []Conn + poolRows []poolRow + poolRowss []poolRows + maxAgeTime time.Time +} + +func (cr *connResource) getConn(p *Pool, res *puddle.Resource[*connResource]) *Conn { + if len(cr.conns) == 0 { + cr.conns = make([]Conn, 128) + } + + c := &cr.conns[len(cr.conns)-1] + cr.conns = cr.conns[0 : len(cr.conns)-1] + + c.res = res + c.p = p + + return c +} + +func (cr *connResource) getPoolRow(c *Conn, r pgx.Row) *poolRow { + if len(cr.poolRows) == 0 { + cr.poolRows = make([]poolRow, 128) + } + + pr := &cr.poolRows[len(cr.poolRows)-1] + cr.poolRows = cr.poolRows[0 : len(cr.poolRows)-1] + + pr.c = c + pr.r = r + + return pr +} + +func (cr *connResource) getPoolRows(c *Conn, r pgx.Rows) *poolRows { + if len(cr.poolRowss) == 0 { + cr.poolRowss = make([]poolRows, 128) + } + + pr := &cr.poolRowss[len(cr.poolRowss)-1] + cr.poolRowss = cr.poolRowss[0 : len(cr.poolRowss)-1] + + pr.c = c + pr.r = r + + return pr +} + +// Pool allows for connection reuse. +type Pool struct { + newConnsCount atomic.Int64 + lifetimeDestroyCount atomic.Int64 + idleDestroyCount atomic.Int64 + + p *puddle.Pool[*connResource] + config *Config + beforeConnect func(context.Context, *pgx.ConnConfig) error + afterConnect func(context.Context, *pgx.Conn) error + prepareConn func(context.Context, *pgx.Conn) (bool, error) + afterRelease func(*pgx.Conn) bool + beforeClose func(*pgx.Conn) + shouldPing func(context.Context, ShouldPingParams) bool + minConns int32 + minIdleConns int32 + maxConns int32 + maxConnLifetime time.Duration + maxConnLifetimeJitter time.Duration + maxConnIdleTime time.Duration + healthCheckPeriod time.Duration + pingTimeout time.Duration + + healthCheckMu sync.Mutex + healthCheckTimer *time.Timer + + healthCheckChan chan struct{} + + acquireTracer AcquireTracer + releaseTracer ReleaseTracer + + closeOnce sync.Once + closeChan chan struct{} +} + +// ShouldPingParams are the parameters passed to ShouldPing. +type ShouldPingParams struct { + Conn *pgx.Conn + IdleDuration time.Duration +} + +// Config is the configuration struct for creating a pool. It must be created by [ParseConfig] and then it can be +// modified. +type Config struct { + ConnConfig *pgx.ConnConfig + + // BeforeConnect is called before a new connection is made. It is passed a copy of the underlying [pgx.ConnConfig] and + // will not impact any existing open connections. + BeforeConnect func(context.Context, *pgx.ConnConfig) error + + // AfterConnect is called after a connection is established, but before it is added to the pool. + AfterConnect func(context.Context, *pgx.Conn) error + + // BeforeAcquire is called before a connection is acquired from the pool. It must return true to allow the + // acquisition or false to indicate that the connection should be destroyed and a different connection should be + // acquired. + // + // Deprecated: Use PrepareConn instead. If both PrepareConn and BeforeAcquire are set, PrepareConn will take + // precedence, ignoring BeforeAcquire. + BeforeAcquire func(context.Context, *pgx.Conn) bool + + // PrepareConn is called before a connection is acquired from the pool. If this function returns true, the connection + // is considered valid, otherwise the connection is destroyed. If the function returns a non-nil error, the instigating + // query will fail with the returned error. + // + // Specifically, this means that: + // + // - If it returns true and a nil error, the query proceeds as normal. + // - If it returns true and an error, the connection will be returned to the pool, and the instigating query will fail with the returned error. + // - If it returns false, and an error, the connection will be destroyed, and the query will fail with the returned error. + // - If it returns false and a nil error, the connection will be destroyed, and the instigating query will be retried on a new connection. + PrepareConn func(context.Context, *pgx.Conn) (bool, error) + + // AfterRelease is called after a connection is released, but before it is returned to the pool. It must return true to + // return the connection to the pool or false to destroy the connection. + AfterRelease func(*pgx.Conn) bool + + // BeforeClose is called right before a connection is closed and removed from the pool. + BeforeClose func(*pgx.Conn) + + // ShouldPing is called after a connection is acquired from the pool. If it returns true, the connection is pinged to check for liveness. + // If this func is not set, the default behavior is to ping connections that have been idle for at least 1 second. + ShouldPing func(context.Context, ShouldPingParams) bool + + // MaxConnLifetime is the duration since creation after which a connection will be automatically closed. + MaxConnLifetime time.Duration + + // MaxConnLifetimeJitter is the duration after MaxConnLifetime to randomly decide to close a connection. + // This helps prevent all connections from being closed at the exact same time, starving the pool. + MaxConnLifetimeJitter time.Duration + + // MaxConnIdleTime is the duration after which an idle connection will be automatically closed by the health check. + MaxConnIdleTime time.Duration + + // PingTimeout is the maximum amount of time to wait for a connection to pong before considering it as unhealthy and + // destroying it. If zero, the default is no timeout. + PingTimeout time.Duration + + // MaxConns is the maximum size of the pool. The default is the greater of 4 or runtime.NumCPU(). + MaxConns int32 + + // MinConns is the minimum size of the pool. After connection closes, the pool might dip below MinConns. A low + // number of MinConns might mean the pool is empty after MaxConnLifetime until the health check has a chance + // to create new connections. + MinConns int32 + + // MinIdleConns is the minimum number of idle connections in the pool. You can increase this to ensure that + // there are always idle connections available. This can help reduce tail latencies during request processing, + // as you can avoid the latency of establishing a new connection while handling requests. It is superior + // to MinConns for this purpose. + // Similar to MinConns, the pool might temporarily dip below MinIdleConns after connection closes. + MinIdleConns int32 + + // HealthCheckPeriod is the duration between checks of the health of idle connections. + HealthCheckPeriod time.Duration + + createdByParseConfig bool // Used to enforce created by ParseConfig rule. +} + +// Copy returns a deep copy of the config that is safe to use and modify. +// The only exception is the tls.Config: +// according to the tls.Config docs it must not be modified after creation. +func (c *Config) Copy() *Config { + newConfig := new(Config) + *newConfig = *c + newConfig.ConnConfig = c.ConnConfig.Copy() + return newConfig +} + +// ConnString returns the connection string as parsed by pgxpool.ParseConfig into pgxpool.Config. +func (c *Config) ConnString() string { return c.ConnConfig.ConnString() } + +// New creates a new Pool. See [ParseConfig] for information on connString format. +func New(ctx context.Context, connString string) (*Pool, error) { + config, err := ParseConfig(connString) + if err != nil { + return nil, err + } + + return NewWithConfig(ctx, config) +} + +// NewWithConfig creates a new [Pool]. config must have been created by [ParseConfig]. +func NewWithConfig(ctx context.Context, config *Config) (*Pool, error) { + // Default values are set in ParseConfig. Enforce initial creation by ParseConfig rather than setting defaults from + // zero values. + if !config.createdByParseConfig { + panic("config must be created by ParseConfig") + } + + prepareConn := config.PrepareConn + if prepareConn == nil && config.BeforeAcquire != nil { + prepareConn = func(ctx context.Context, conn *pgx.Conn) (bool, error) { + return config.BeforeAcquire(ctx, conn), nil + } + } + + p := &Pool{ + config: config, + beforeConnect: config.BeforeConnect, + afterConnect: config.AfterConnect, + prepareConn: prepareConn, + afterRelease: config.AfterRelease, + beforeClose: config.BeforeClose, + minConns: config.MinConns, + minIdleConns: config.MinIdleConns, + maxConns: config.MaxConns, + maxConnLifetime: config.MaxConnLifetime, + maxConnLifetimeJitter: config.MaxConnLifetimeJitter, + maxConnIdleTime: config.MaxConnIdleTime, + pingTimeout: config.PingTimeout, + healthCheckPeriod: config.HealthCheckPeriod, + healthCheckChan: make(chan struct{}, 1), + closeChan: make(chan struct{}), + } + + if t, ok := config.ConnConfig.Tracer.(AcquireTracer); ok { + p.acquireTracer = t + } + + if t, ok := config.ConnConfig.Tracer.(ReleaseTracer); ok { + p.releaseTracer = t + } + + if config.ShouldPing != nil { + p.shouldPing = config.ShouldPing + } else { + p.shouldPing = func(ctx context.Context, params ShouldPingParams) bool { + return params.IdleDuration > time.Second + } + } + + var err error + p.p, err = puddle.NewPool( + &puddle.Config[*connResource]{ + Constructor: func(ctx context.Context) (*connResource, error) { + p.newConnsCount.Add(1) + connConfig := p.config.ConnConfig.Copy() + + // Connection will continue in background even if Acquire is canceled. Ensure that a connect won't hang forever. + if connConfig.ConnectTimeout <= 0 { + connConfig.ConnectTimeout = 2 * time.Minute + } + + if p.beforeConnect != nil { + if err := p.beforeConnect(ctx, connConfig); err != nil { + return nil, err + } + } + + conn, err := pgx.ConnectConfig(ctx, connConfig) + if err != nil { + return nil, err + } + + if p.afterConnect != nil { + err = p.afterConnect(ctx, conn) + if err != nil { + conn.Close(ctx) + return nil, err + } + } + + jitterSecs := rand.Float64() * config.MaxConnLifetimeJitter.Seconds() + maxAgeTime := time.Now().Add(config.MaxConnLifetime).Add(time.Duration(jitterSecs) * time.Second) + + cr := &connResource{ + conn: conn, + conns: make([]Conn, 64), + poolRows: make([]poolRow, 64), + poolRowss: make([]poolRows, 64), + maxAgeTime: maxAgeTime, + } + + return cr, nil + }, + Destructor: func(value *connResource) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + conn := value.conn + if p.beforeClose != nil { + p.beforeClose(conn) + } + conn.Close(ctx) + select { + case <-conn.PgConn().CleanupDone(): + case <-ctx.Done(): + } + cancel() + }, + MaxSize: config.MaxConns, + }, + ) + if err != nil { + return nil, err + } + + go func() { + targetIdleResources := max(int(p.minConns), int(p.minIdleConns)) + p.createIdleResources(ctx, targetIdleResources) + p.backgroundHealthCheck() + }() + + return p, nil +} + +// ParseConfig builds a Config from connString. It parses connString with the same behavior as [pgx.ParseConfig] with the +// addition of the following variables: +// +// - pool_max_conns: integer greater than 0 (default 4) +// - pool_min_conns: integer 0 or greater (default 0) +// - pool_max_conn_lifetime: duration string (default 1 hour) +// - pool_max_conn_idle_time: duration string (default 30 minutes) +// - pool_health_check_period: duration string (default 1 minute) +// - pool_max_conn_lifetime_jitter: duration string (default 0) +// +// See Config for definitions of these arguments. +// +// # Example Keyword/Value +// user=jack password=secret host=pg.example.com port=5432 dbname=mydb sslmode=verify-ca pool_max_conns=10 pool_max_conn_lifetime=1h30m +// +// # Example URL +// postgres://jack:secret@pg.example.com:5432/mydb?sslmode=verify-ca&pool_max_conns=10&pool_max_conn_lifetime=1h30m +func ParseConfig(connString string) (*Config, error) { + connConfig, err := pgx.ParseConfig(connString) + if err != nil { + return nil, err + } + + config := &Config{ + ConnConfig: connConfig, + createdByParseConfig: true, + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conns"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_max_conns") + n, err := strconv.ParseInt(s, 10, 32) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conns", err) + } + if n < 1 { + return nil, pgconn.NewParseConfigError(connString, "pool_max_conns too small", err) + } + config.MaxConns = int32(n) + } else { + config.MaxConns = defaultMaxConns + if numCPU := int32(runtime.NumCPU()); numCPU > config.MaxConns { + config.MaxConns = numCPU + } + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_min_conns"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_min_conns") + n, err := strconv.ParseInt(s, 10, 32) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_min_conns", err) + } + config.MinConns = int32(n) + } else { + config.MinConns = defaultMinConns + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_min_idle_conns"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_min_idle_conns") + n, err := strconv.ParseInt(s, 10, 32) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_min_idle_conns", err) + } + config.MinIdleConns = int32(n) + } else { + config.MinIdleConns = defaultMinIdleConns + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_lifetime"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_max_conn_lifetime") + d, err := time.ParseDuration(s) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_lifetime", err) + } + config.MaxConnLifetime = d + } else { + config.MaxConnLifetime = defaultMaxConnLifetime + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_idle_time"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_max_conn_idle_time") + d, err := time.ParseDuration(s) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_idle_time", err) + } + config.MaxConnIdleTime = d + } else { + config.MaxConnIdleTime = defaultMaxConnIdleTime + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_health_check_period"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_health_check_period") + d, err := time.ParseDuration(s) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_health_check_period", err) + } + config.HealthCheckPeriod = d + } else { + config.HealthCheckPeriod = defaultHealthCheckPeriod + } + + if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_lifetime_jitter"]; ok { + delete(connConfig.Config.RuntimeParams, "pool_max_conn_lifetime_jitter") + d, err := time.ParseDuration(s) + if err != nil { + return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_lifetime_jitter", err) + } + config.MaxConnLifetimeJitter = d + } + + return config, nil +} + +// Close closes all connections in the pool and rejects future [Pool.Acquire] calls. Blocks until all connections are returned +// to pool and closed. +func (p *Pool) Close() { + p.closeOnce.Do(func() { + close(p.closeChan) + p.p.Close() + }) +} + +func (p *Pool) isExpired(res *puddle.Resource[*connResource]) bool { + if p.maxConnLifetime <= 0 { + return false + } + return time.Now().After(res.Value().maxAgeTime) +} + +func (p *Pool) triggerHealthCheck() { + const healthCheckDelay = 500 * time.Millisecond + + p.healthCheckMu.Lock() + defer p.healthCheckMu.Unlock() + + if p.healthCheckTimer == nil { + // Destroy is asynchronous so we give it time to actually remove itself from + // the pool otherwise we might try to check the pool size too soon + p.healthCheckTimer = time.AfterFunc(healthCheckDelay, func() { + select { + case <-p.closeChan: + case p.healthCheckChan <- struct{}{}: + default: + } + }) + return + } + + p.healthCheckTimer.Reset(healthCheckDelay) +} + +func (p *Pool) backgroundHealthCheck() { + ticker := time.NewTicker(p.healthCheckPeriod) + defer ticker.Stop() + for { + select { + case <-p.closeChan: + return + case <-p.healthCheckChan: + p.checkHealth() + case <-ticker.C: + p.checkHealth() + } + } +} + +func (p *Pool) checkHealth() { + for { + // If checkMinConns failed we don't destroy any connections since we couldn't + // even get to minConns + if err := p.checkMinConns(); err != nil { + // Should we log this error somewhere? + break + } + if !p.checkConnsHealth() { + // Since we didn't destroy any connections we can stop looping + break + } + // Technically Destroy is asynchronous but 500ms should be enough for it to + // remove it from the underlying pool + select { + case <-p.closeChan: + return + case <-time.After(500 * time.Millisecond): + } + } +} + +// checkConnsHealth will check all idle connections, destroy a connection if +// it's idle or too old, and returns true if any were destroyed +func (p *Pool) checkConnsHealth() bool { + var destroyed bool + totalConns := p.Stat().TotalConns() + resources := p.p.AcquireAllIdle() + for _, res := range resources { + switch { + // We're okay going under minConns if the lifetime is up + case p.isExpired(res) && totalConns >= p.minConns: + p.lifetimeDestroyCount.Add(1) + res.Destroy() + destroyed = true + // Since Destroy is async we manually decrement totalConns. + totalConns-- + case res.IdleDuration() > p.maxConnIdleTime && totalConns > p.minConns: + p.idleDestroyCount.Add(1) + res.Destroy() + destroyed = true + // Since Destroy is async we manually decrement totalConns. + totalConns-- + default: + res.ReleaseUnused() + } + } + return destroyed +} + +func (p *Pool) checkMinConns() error { + // TotalConns can include ones that are being destroyed but we should have + // sleep(500ms) around all of the destroys to help prevent that from throwing + // off this check + + // Create the number of connections needed to get to both minConns and minIdleConns + stat := p.Stat() + toCreate := max(p.minConns-stat.TotalConns(), p.minIdleConns-stat.IdleConns()) + if toCreate > 0 { + return p.createIdleResources(context.Background(), int(toCreate)) + } + return nil +} + +func (p *Pool) createIdleResources(parentCtx context.Context, targetResources int) error { + ctx, cancel := context.WithCancel(parentCtx) + defer cancel() + + errs := make(chan error, targetResources) + + for range targetResources { + go func() { + err := p.p.CreateResource(ctx) + // Ignore ErrNotAvailable since it means that the pool has become full since we started creating resource. + if err == puddle.ErrNotAvailable { + err = nil + } + errs <- err + }() + } + + var firstError error + for range targetResources { + err := <-errs + if err != nil && firstError == nil { + cancel() + firstError = err + } + } + + return firstError +} + +// Acquire returns a connection ([Conn]) from the [Pool]. +func (p *Pool) Acquire(ctx context.Context) (c *Conn, err error) { + if p.acquireTracer != nil { + ctx = p.acquireTracer.TraceAcquireStart(ctx, p, TraceAcquireStartData{}) + defer func() { + var conn *pgx.Conn + if c != nil { + conn = c.Conn() + } + p.acquireTracer.TraceAcquireEnd(ctx, p, TraceAcquireEndData{Conn: conn, Err: err}) + }() + } + + // Try to acquire from the connection pool up to maxConns + 1 times, so that + // any that fatal errors would empty the pool and still at least try 1 fresh + // connection. + for range int(p.maxConns) + 1 { + res, err := p.p.Acquire(ctx) + if err != nil { + return nil, err + } + + cr := res.Value() + + // Destroy expired connections before doing any further work (such as + // pinging) on them. This enforces MaxConnLifetime at acquire time so that + // a connection that expired while idle on a busy pool is not handed out. + if p.isExpired(res) { + p.lifetimeDestroyCount.Add(1) + res.Destroy() + if ctx.Err() != nil { + return nil, ctx.Err() + } + continue + } + + shouldPingParams := ShouldPingParams{Conn: cr.conn, IdleDuration: res.IdleDuration()} + if p.shouldPing(ctx, shouldPingParams) { + err := func() error { + pingCtx := ctx + if p.pingTimeout > 0 { + var cancel context.CancelFunc + pingCtx, cancel = context.WithTimeout(ctx, p.pingTimeout) + defer cancel() + } + return cr.conn.Ping(pingCtx) + }() + if err != nil { + res.Destroy() + if ctx.Err() != nil { + return nil, ctx.Err() + } + continue + } + } + + if p.prepareConn != nil { + ok, err := p.prepareConn(ctx, cr.conn) + if !ok { + res.Destroy() + } + if err != nil { + if ok { + res.Release() + } + return nil, err + } + if !ok { + if ctx.Err() != nil { + return nil, ctx.Err() + } + continue + } + } + + return cr.getConn(p, res), nil + } + return nil, errors.New("pgxpool: too many failed attempts acquiring connection; likely bug in PrepareConn, BeforeAcquire, or ShouldPing hook") +} + +// AcquireFunc acquires a [Conn] and calls f with that [Conn]. ctx will only affect the [Pool.Acquire]. It has no effect on the +// call of f. The return value is either an error acquiring the [Conn] or the return value of f. The [Conn] is +// automatically released after the call of f. +func (p *Pool) AcquireFunc(ctx context.Context, f func(*Conn) error) error { + conn, err := p.Acquire(ctx) + if err != nil { + return err + } + defer conn.Release() + + return f(conn) +} + +// AcquireAllIdle atomically acquires all currently idle connections. Its intended use is for health check and +// keep-alive functionality. It does not update pool statistics. +func (p *Pool) AcquireAllIdle(ctx context.Context) []*Conn { + resources := p.p.AcquireAllIdle() + conns := make([]*Conn, 0, len(resources)) + for _, res := range resources { + cr := res.Value() + if p.prepareConn != nil { + ok, err := p.prepareConn(ctx, cr.conn) + if !ok || err != nil { + res.Destroy() + continue + } + } + conns = append(conns, cr.getConn(p, res)) + } + + return conns +} + +// Reset closes all connections, but leaves the pool open. It is intended for use when an error is detected that would +// disrupt all connections (such as a network interruption or a server state change). +// +// It is safe to reset a pool while connections are checked out. Those connections will be closed when they are returned +// to the pool. +func (p *Pool) Reset() { + p.p.Reset() +} + +// Config returns a copy of config that was used to initialize this [Pool]. +func (p *Pool) Config() *Config { return p.config.Copy() } + +// Stat returns a pgxpool.Stat struct with a snapshot of Pool statistics. +func (p *Pool) Stat() *Stat { + return &Stat{ + s: p.p.Stat(), + newConnsCount: p.newConnsCount.Load(), + lifetimeDestroyCount: p.lifetimeDestroyCount.Load(), + idleDestroyCount: p.idleDestroyCount.Load(), + } +} + +// Exec acquires a connection from the [Pool] and executes the given SQL. +// SQL can be either a prepared statement name or an SQL string. +// Arguments should be referenced positionally from the SQL string as $1, $2, etc. +// The acquired connection is returned to the pool when the [Pool.Exec] function returns. +func (p *Pool) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) { + c, err := p.Acquire(ctx) + if err != nil { + return pgconn.CommandTag{}, err + } + defer c.Release() + + return c.Exec(ctx, sql, arguments...) +} + +// Query acquires a connection and executes a query that returns [pgx.Rows]. +// Arguments should be referenced positionally from the SQL string as $1, $2, etc. +// See [pgx.Rows] documentation to close the returned [pgx.Rows] and return the acquired connection to the [Pool]. +// +// If there is an error, the returned [pgx.Rows] will be returned in an error state. +// If preferred, ignore the error returned from [Pool.Query] and handle errors using the returned [pgx.Rows]. +// +// For extra control over how the query is executed, the types [pgx.QueryExecMode], [pgx.QueryResultFormats], and +// [pgx.QueryResultFormatsByOID] may be used as the first args to control exactly how the query is executed. This is rarely +// needed. See the documentation for those types for details. +func (p *Pool) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) { + c, err := p.Acquire(ctx) + if err != nil { + return errRows{err: err}, err + } + + rows, err := c.Query(ctx, sql, args...) + if err != nil { + c.Release() + return errRows{err: err}, err + } + + return c.getPoolRows(rows), nil +} + +// QueryRow acquires a connection and executes a query that is expected +// to return at most one row ([pgx.Row]). Errors are deferred until [pgx.Row]'s +// Scan method is called. If the query selects no rows, [pgx.Row]'s Scan will +// return [pgx.ErrNoRows]. Otherwise, [pgx.Row]'s Scan scans the first selected row +// and discards the rest. The acquired connection is returned to the [Pool] when +// [pgx.Row]'s Scan method is called. +// +// Arguments should be referenced positionally from the SQL string as $1, $2, etc. +// +// For extra control over how the query is executed, the types [pgx.QueryExecMode], [pgx.QueryResultFormats], and +// [pgx.QueryResultFormatsByOID] may be used as the first args to control exactly how the query is executed. This is rarely +// needed. See the documentation for those types for details. +func (p *Pool) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row { + c, err := p.Acquire(ctx) + if err != nil { + return errRow{err: err} + } + + row := c.QueryRow(ctx, sql, args...) + return c.getPoolRow(row) +} + +func (p *Pool) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults { + c, err := p.Acquire(ctx) + if err != nil { + return errBatchResults{err: err} + } + + br := c.SendBatch(ctx, b) + return &poolBatchResults{br: br, c: c} +} + +// Begin acquires a connection from the [Pool] and starts a transaction. Unlike [database/sql], the context only affects the begin command. i.e. there is no +// auto-rollback on context cancellation. Begin initiates a transaction block without explicitly setting a transaction mode for the block (see [Pool.BeginTx] with [pgx.TxOptions] if transaction mode is required). +// [*Tx] is returned, which implements the [pgx.Tx] interface. +// [Tx.Commit] or [Tx.Rollback] must be called on the returned transaction to finalize the transaction block. +func (p *Pool) Begin(ctx context.Context) (pgx.Tx, error) { + return p.BeginTx(ctx, pgx.TxOptions{}) +} + +// BeginTx acquires a connection from the [Pool] and starts a transaction with [pgx.TxOptions] determining the transaction mode. +// Unlike [database/sql], the context only affects the begin command. i.e. there is no auto-rollback on context cancellation. +// [*Tx] is returned, which implements the [pgx.Tx] interface. +// [Tx.Commit] or [Tx.Rollback] must be called on the returned transaction to finalize the transaction block. +func (p *Pool) BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) { + c, err := p.Acquire(ctx) + if err != nil { + return nil, err + } + + t, err := c.BeginTx(ctx, txOptions) + if err != nil { + c.Release() + return nil, err + } + + return &Tx{t: t, c: c}, nil +} + +func (p *Pool) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) { + c, err := p.Acquire(ctx) + if err != nil { + return 0, err + } + defer c.Release() + + return c.Conn().CopyFrom(ctx, tableName, columnNames, rowSrc) +} + +// Ping acquires a connection from the [Pool] and executes an empty sql statement against it. +// If the sql returns without error, the database [Pool.Ping] is considered successful, otherwise, the error is returned. +func (p *Pool) Ping(ctx context.Context) error { + c, err := p.Acquire(ctx) + if err != nil { + return err + } + defer c.Release() + return c.Ping(ctx) +} diff --git a/pgxpool/pool_test.go b/pgxpool/pool_test.go new file mode 100644 index 0000000..2bae829 --- /dev/null +++ b/pgxpool/pool_test.go @@ -0,0 +1,1445 @@ +package pgxpool_test + +import ( + "context" + "errors" + "fmt" + "math" + "os" + "sync/atomic" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/jackc/pgx/v5/pgxtest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestConnect(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + connString := os.Getenv("PGX_TEST_DATABASE") + pool, err := pgxpool.New(ctx, connString) + require.NoError(t, err) + assert.Equal(t, connString, pool.Config().ConnString()) + pool.Close() +} + +func TestConnectConfig(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + connString := os.Getenv("PGX_TEST_DATABASE") + config, err := pgxpool.ParseConfig(connString) + require.NoError(t, err) + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + assertConfigsEqual(t, config, pool.Config(), "Pool.Config() returns original config") + pool.Close() +} + +func TestParseConfigExtractsPoolArguments(t *testing.T) { + t.Parallel() + + config, err := pgxpool.ParseConfig("pool_max_conns=42 pool_min_conns=1 pool_min_idle_conns=2") + assert.NoError(t, err) + assert.EqualValues(t, 42, config.MaxConns) + assert.EqualValues(t, 1, config.MinConns) + assert.EqualValues(t, 2, config.MinIdleConns) + assert.NotContains(t, config.ConnConfig.Config.RuntimeParams, "pool_max_conns") + assert.NotContains(t, config.ConnConfig.Config.RuntimeParams, "pool_min_conns") +} + +func TestConstructorIgnoresContext(t *testing.T) { + t.Parallel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + assert.NoError(t, err) + var cancel func() + config.BeforeConnect = func(context.Context, *pgx.ConnConfig) error { + // cancel the query's context before we actually Dial to ensure the Dial's + // context isn't cancelled + cancel() + return nil + } + + pool, err := pgxpool.NewWithConfig(context.Background(), config) + require.NoError(t, err) + + assert.EqualValues(t, 0, pool.Stat().TotalConns()) + + var ctx context.Context + ctx, cancel = context.WithCancel(context.Background()) + defer cancel() + _, err = pool.Exec(ctx, "SELECT 1") + assert.ErrorIs(t, err, context.Canceled) + assert.EqualValues(t, 1, pool.Stat().TotalConns()) +} + +func TestConnectConfigRequiresConnConfigFromParseConfig(t *testing.T) { + t.Parallel() + + config := &pgxpool.Config{} + + require.PanicsWithValue(t, "config must be created by ParseConfig", func() { pgxpool.NewWithConfig(context.Background(), config) }) +} + +func TestConfigCopyReturnsEqualConfig(t *testing.T) { + connString := "postgres://jack:secret@localhost:5432/mydb?application_name=pgxtest&search_path=myschema&connect_timeout=5" + original, err := pgxpool.ParseConfig(connString) + require.NoError(t, err) + + copied := original.Copy() + + assertConfigsEqual(t, original, copied, t.Name()) +} + +func TestConfigCopyCanBeUsedToConnect(t *testing.T) { + connString := os.Getenv("PGX_TEST_DATABASE") + original, err := pgxpool.ParseConfig(connString) + require.NoError(t, err) + + copied := original.Copy() + assert.NotPanics(t, func() { + _, err = pgxpool.NewWithConfig(context.Background(), copied) + }) + assert.NoError(t, err) +} + +func TestPoolAcquireAndConnRelease(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + c.Release() +} + +func TestPoolAcquireAndConnHijack(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + + connsBeforeHijack := pool.Stat().TotalConns() + + conn := c.Hijack() + defer conn.Close(ctx) + + connsAfterHijack := pool.Stat().TotalConns() + require.Equal(t, connsBeforeHijack-1, connsAfterHijack) + + var n int32 + err = conn.QueryRow(ctx, `select 1`).Scan(&n) + require.NoError(t, err) + require.Equal(t, int32(1), n) +} + +func TestPoolAcquireChecksIdleConns(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + controllerConn, err := pgx.Connect(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer controllerConn.Close(ctx) + pgxtest.SkipCockroachDB(t, controllerConn, "Server does not support pg_terminate_backend() (https://github.com/cockroachdb/cockroach/issues/35897)") + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + var conns []*pgxpool.Conn + for range 3 { + c, err := pool.Acquire(ctx) + require.NoError(t, err) + conns = append(conns, c) + } + + require.EqualValues(t, 3, pool.Stat().TotalConns()) + + var pids []uint32 + for _, c := range conns { + pids = append(pids, c.Conn().PgConn().PID()) + c.Release() + } + + _, err = controllerConn.Exec(ctx, `select pg_terminate_backend(n) from unnest($1::int[]) n`, pids) + require.NoError(t, err) + + // All conns are dead they don't know it and neither does the pool. + require.EqualValues(t, 3, pool.Stat().TotalConns()) + + // Wait long enough so the pool will realize it needs to check the connections. + time.Sleep(time.Second) + + // Pool should try all existing connections and find them dead, then create a new connection which should successfully ping. + err = pool.Ping(ctx) + require.NoError(t, err) + + // The original 3 conns should have been terminated and the a new conn established for the ping. + require.EqualValues(t, 1, pool.Stat().TotalConns()) + c, err := pool.Acquire(ctx) + require.NoError(t, err) + + cPID := c.Conn().PgConn().PID() + c.Release() + + require.NotContains(t, pids, cPID) +} + +func TestPoolAcquireChecksIdleConnsWithShouldPing(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + controllerConn, err := pgx.Connect(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer controllerConn.Close(ctx) + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + // Replace the default ShouldPing func + var shouldPingLastCalledWith *pgxpool.ShouldPingParams + config.ShouldPing = func(ctx context.Context, params pgxpool.ShouldPingParams) bool { + shouldPingLastCalledWith = ¶ms + return false + } + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + c.Release() + + time.Sleep(time.Millisecond * 200) + + c, err = pool.Acquire(ctx) + require.NoError(t, err) + conn := c.Conn() + + require.NotNil(t, shouldPingLastCalledWith) + assert.Equal(t, conn, shouldPingLastCalledWith.Conn) + assert.InDelta(t, time.Millisecond*200, shouldPingLastCalledWith.IdleDuration, float64(time.Millisecond*100)) + + c.Release() +} + +// https://github.com/jackc/pgx/issues/2379 +func TestPoolAcquireWithMaxConnsEqualsMaxInt32(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MaxConns = math.MaxInt32 + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + c.Release() +} + +func TestPoolAcquireFunc(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + var n int32 + err = pool.AcquireFunc(ctx, func(c *pgxpool.Conn) error { + return c.QueryRow(ctx, "select 1").Scan(&n) + }) + require.NoError(t, err) + require.EqualValues(t, 1, n) +} + +func TestPoolAcquireFuncReturnsFnError(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + err = pool.AcquireFunc(ctx, func(c *pgxpool.Conn) error { + return fmt.Errorf("some error") + }) + require.EqualError(t, err, "some error") +} + +func TestPoolBeforeConnect(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.BeforeConnect = func(ctx context.Context, cfg *pgx.ConnConfig) error { + cfg.Config.RuntimeParams["application_name"] = "pgx" + return nil + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + var str string + err = db.QueryRow(ctx, "SHOW application_name").Scan(&str) + require.NoError(t, err) + assert.EqualValues(t, "pgx", str) +} + +func TestPoolAfterConnect(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.AfterConnect = func(ctx context.Context, c *pgx.Conn) error { + _, err := c.Prepare(ctx, "ps1", "select 1") + return err + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + var n int32 + err = db.QueryRow(ctx, "ps1").Scan(&n) + require.NoError(t, err) + assert.EqualValues(t, 1, n) +} + +func TestPoolBeforeAcquire(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + acquireAttempts := 0 + + config.BeforeAcquire = func(ctx context.Context, c *pgx.Conn) bool { + acquireAttempts++ + return acquireAttempts%2 == 0 + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + conns := make([]*pgxpool.Conn, 4) + for i := range conns { + conns[i], err = db.Acquire(ctx) + assert.NoError(t, err) + } + + for _, c := range conns { + c.Release() + } + waitForReleaseToComplete() + + assert.EqualValues(t, 8, acquireAttempts) + + conns = db.AcquireAllIdle(ctx) + assert.Len(t, conns, 2) + + for _, c := range conns { + c.Release() + } + waitForReleaseToComplete() + + assert.EqualValues(t, 12, acquireAttempts) +} + +func TestPoolPrepareConn(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + acquireAttempts := 0 + + config.PrepareConn = func(context.Context, *pgx.Conn) (bool, error) { + acquireAttempts++ + var err error + if acquireAttempts%3 == 0 { + err = errors.New("PrepareConn error") + } + return acquireAttempts%2 == 0, err + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + t.Cleanup(db.Close) + + var errorCount int + conns := make([]*pgxpool.Conn, 0, 4) + for { + conn, err := db.Acquire(ctx) + if err != nil { + errorCount++ + continue + } + conns = append(conns, conn) + if len(conns) == 4 { + break + } + } + const wantErrorCount = 3 + assert.Equal(t, wantErrorCount, errorCount, "Acquire() should have failed %d times", wantErrorCount) + + for _, c := range conns { + c.Release() + } + waitForReleaseToComplete() + + assert.EqualValues(t, len(conns)*2+wantErrorCount-1, acquireAttempts) + + conns = db.AcquireAllIdle(ctx) + assert.Len(t, conns, 1) + + for _, c := range conns { + c.Release() + } + waitForReleaseToComplete() + + assert.EqualValues(t, 14, acquireAttempts) +} + +func TestPoolAfterRelease(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + func() { + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + }() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + afterReleaseCount := 0 + + config.AfterRelease = func(c *pgx.Conn) bool { + afterReleaseCount++ + return afterReleaseCount%2 == 1 + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + connPIDs := map[uint32]struct{}{} + + for range 10 { + conn, err := db.Acquire(ctx) + assert.NoError(t, err) + connPIDs[conn.Conn().PgConn().PID()] = struct{}{} + conn.Release() + waitForReleaseToComplete() + } + + assert.EqualValues(t, 5, len(connPIDs)) +} + +func TestPoolBeforeClose(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + func() { + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + }() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + connPIDs := make(chan uint32, 5) + config.BeforeClose = func(c *pgx.Conn) { + connPIDs <- c.PgConn().PID() + } + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + acquiredPIDs := make([]uint32, 0, 5) + closedPIDs := make([]uint32, 0, 5) + for range 5 { + conn, err := db.Acquire(ctx) + assert.NoError(t, err) + acquiredPIDs = append(acquiredPIDs, conn.Conn().PgConn().PID()) + conn.Release() + db.Reset() + closedPIDs = append(closedPIDs, <-connPIDs) + } + + assert.ElementsMatch(t, acquiredPIDs, closedPIDs) +} + +func TestPoolAcquireAllIdle(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + conns := make([]*pgxpool.Conn, 3) + for i := range conns { + conns[i], err = db.Acquire(ctx) + assert.NoError(t, err) + } + + for _, c := range conns { + if c != nil { + c.Release() + } + } + waitForReleaseToComplete() + + conns = db.AcquireAllIdle(ctx) + assert.Len(t, conns, 3) + + for _, c := range conns { + c.Release() + } +} + +func TestPoolReset(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + conns := make([]*pgxpool.Conn, 3) + for i := range conns { + conns[i], err = db.Acquire(ctx) + assert.NoError(t, err) + } + + db.Reset() + + for _, c := range conns { + if c != nil { + c.Release() + } + } + waitForReleaseToComplete() + + require.EqualValues(t, 0, db.Stat().TotalConns()) +} + +func TestConnReleaseChecksMaxConnLifetime(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MaxConnLifetime = 250 * time.Millisecond + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + c, err := db.Acquire(ctx) + require.NoError(t, err) + + time.Sleep(config.MaxConnLifetime) + + c.Release() + waitForReleaseToComplete() + + stats := db.Stat() + assert.EqualValues(t, 0, stats.TotalConns()) +} + +func TestConnReleaseClosesBusyConn(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + c, err := db.Acquire(ctx) + require.NoError(t, err) + + _, err = c.Query(ctx, "select generate_series(1,10)") + require.NoError(t, err) + + c.Release() + waitForReleaseToComplete() + + // wait for the connection to actually be destroyed + for range 1000 { + if db.Stat().TotalConns() == 0 { + break + } + time.Sleep(time.Millisecond) + } + + stats := db.Stat() + assert.EqualValues(t, 0, stats.TotalConns()) +} + +func TestPoolBackgroundChecksMaxConnLifetime(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MaxConnLifetime = 100 * time.Millisecond + config.HealthCheckPeriod = 100 * time.Millisecond + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + c, err := db.Acquire(ctx) + require.NoError(t, err) + c.Release() + time.Sleep(config.MaxConnLifetime + 500*time.Millisecond) + + stats := db.Stat() + assert.EqualValues(t, 0, stats.TotalConns()) + assert.EqualValues(t, 0, stats.MaxIdleDestroyCount()) + assert.EqualValues(t, 1, stats.MaxLifetimeDestroyCount()) + assert.EqualValues(t, 1, stats.NewConnsCount()) +} + +func TestPoolBackgroundChecksMaxConnIdleTime(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MaxConnLifetime = 1 * time.Minute + config.MaxConnIdleTime = 100 * time.Millisecond + config.HealthCheckPeriod = 150 * time.Millisecond + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + c, err := db.Acquire(ctx) + require.NoError(t, err) + c.Release() + time.Sleep(config.HealthCheckPeriod) + + for range 1000 { + if db.Stat().TotalConns() == 0 { + break + } + time.Sleep(time.Millisecond) + } + + stats := db.Stat() + assert.EqualValues(t, 0, stats.TotalConns()) + assert.EqualValues(t, 1, stats.MaxIdleDestroyCount()) + assert.EqualValues(t, 0, stats.MaxLifetimeDestroyCount()) + assert.EqualValues(t, 1, stats.NewConnsCount()) +} + +func TestPoolBackgroundChecksMinConns(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.HealthCheckPeriod = 100 * time.Millisecond + config.MinConns = 2 + + db, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer db.Close() + + stats := db.Stat() + for !(stats.IdleConns() == 2 && stats.MaxLifetimeDestroyCount() == 0 && stats.NewConnsCount() == 2) && ctx.Err() == nil { + time.Sleep(50 * time.Millisecond) + stats = db.Stat() + } + require.EqualValues(t, 2, stats.IdleConns()) + require.EqualValues(t, 0, stats.MaxLifetimeDestroyCount()) + require.EqualValues(t, 2, stats.NewConnsCount()) + + c, err := db.Acquire(ctx) + require.NoError(t, err) + + stats = db.Stat() + require.EqualValues(t, 1, stats.IdleConns()) + require.EqualValues(t, 0, stats.MaxLifetimeDestroyCount()) + require.EqualValues(t, 2, stats.NewConnsCount()) + + err = c.Conn().Close(ctx) + require.NoError(t, err) + c.Release() + + stats = db.Stat() + for !(stats.IdleConns() == 2 && stats.MaxIdleDestroyCount() == 0 && stats.NewConnsCount() == 3) && ctx.Err() == nil { + time.Sleep(50 * time.Millisecond) + stats = db.Stat() + } + require.EqualValues(t, 2, stats.TotalConns()) + require.EqualValues(t, 0, stats.MaxIdleDestroyCount()) + require.EqualValues(t, 3, stats.NewConnsCount()) +} + +func TestPoolExec(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + testExec(t, ctx, pool) +} + +func TestPoolQuery(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + // Test common usage + testQuery(t, ctx, pool) + waitForReleaseToComplete() + + // Test expected pool behavior + rows, err := pool.Query(ctx, "select generate_series(1,$1)", 10) + require.NoError(t, err) + + stats := pool.Stat() + assert.EqualValues(t, 1, stats.AcquiredConns()) + assert.EqualValues(t, 1, stats.TotalConns()) + + rows.Close() + assert.NoError(t, rows.Err()) + waitForReleaseToComplete() + + stats = pool.Stat() + assert.EqualValues(t, 0, stats.AcquiredConns()) + assert.EqualValues(t, 1, stats.TotalConns()) +} + +func TestPoolQueryRow(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + testQueryRow(t, ctx, pool) + waitForReleaseToComplete() + + stats := pool.Stat() + assert.EqualValues(t, 0, stats.AcquiredConns()) + assert.EqualValues(t, 1, stats.TotalConns()) +} + +// https://github.com/jackc/pgx/issues/677 +func TestPoolQueryRowErrNoRows(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + err = pool.QueryRow(ctx, "select n from generate_series(1,10) n where n=0").Scan(nil) + require.Equal(t, pgx.ErrNoRows, err) +} + +// https://github.com/jackc/pgx/issues/1628 +func TestPoolQueryRowScanPanicReleasesConnection(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + require.Panics(t, func() { + var greeting *string + pool.QueryRow(ctx, "select 'Hello, world!'").Scan(greeting) // Note lack of &. This means that a typed nil is passed to Scan. + }) + + // If the connection is not released this will block forever in the defer pool.Close(). +} + +func TestPoolSendBatch(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + testSendBatch(t, ctx, pool) + waitForReleaseToComplete() + + stats := pool.Stat() + assert.EqualValues(t, 0, stats.AcquiredConns()) + assert.EqualValues(t, 1, stats.TotalConns()) +} + +func TestPoolCopyFrom(t *testing.T) { + // Not able to use testCopyFrom because it relies on temporary tables and the pool may run subsequent calls under + // different connections. + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + _, err = pool.Exec(ctx, `drop table if exists poolcopyfromtest`) + require.NoError(t, err) + + _, err = pool.Exec(ctx, `create table poolcopyfromtest(a int2, b int4, c int8, d varchar, e text, f date, g timestamptz)`) + require.NoError(t, err) + defer pool.Exec(ctx, `drop table poolcopyfromtest`) + + tzedTime := time.Date(2010, 2, 3, 4, 5, 6, 0, time.Local) + + inputRows := [][]any{ + {int16(0), int32(1), int64(2), "abc", "efg", time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC), tzedTime}, + {nil, nil, nil, nil, nil, nil, nil}, + } + + copyCount, err := pool.CopyFrom(ctx, pgx.Identifier{"poolcopyfromtest"}, []string{"a", "b", "c", "d", "e", "f", "g"}, pgx.CopyFromRows(inputRows)) + assert.NoError(t, err) + assert.EqualValues(t, len(inputRows), copyCount) + + rows, err := pool.Query(ctx, "select * from poolcopyfromtest") + assert.NoError(t, err) + + var outputRows [][]any + for rows.Next() { + row, err := rows.Values() + if err != nil { + t.Errorf("Unexpected error for rows.Values(): %v", err) + } + outputRows = append(outputRows, row) + } + + assert.NoError(t, rows.Err()) + assert.Equal(t, inputRows, outputRows) +} + +func TestConnReleaseClosesConnInFailedTransaction(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + + pid := c.Conn().PgConn().PID() + + assert.Equal(t, byte('I'), c.Conn().PgConn().TxStatus()) + + _, err = c.Exec(ctx, "begin") + assert.NoError(t, err) + + assert.Equal(t, byte('T'), c.Conn().PgConn().TxStatus()) + + _, err = c.Exec(ctx, "selct") + assert.Error(t, err) + + assert.Equal(t, byte('E'), c.Conn().PgConn().TxStatus()) + + c.Release() + waitForReleaseToComplete() + + c, err = pool.Acquire(ctx) + require.NoError(t, err) + + assert.NotEqual(t, pid, c.Conn().PgConn().PID()) + assert.Equal(t, byte('I'), c.Conn().PgConn().TxStatus()) + + c.Release() +} + +func TestConnReleaseClosesConnInTransaction(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + + pid := c.Conn().PgConn().PID() + + assert.Equal(t, byte('I'), c.Conn().PgConn().TxStatus()) + + _, err = c.Exec(ctx, "begin") + assert.NoError(t, err) + + assert.Equal(t, byte('T'), c.Conn().PgConn().TxStatus()) + + c.Release() + waitForReleaseToComplete() + + c, err = pool.Acquire(ctx) + require.NoError(t, err) + + assert.NotEqual(t, pid, c.Conn().PgConn().PID()) + assert.Equal(t, byte('I'), c.Conn().PgConn().TxStatus()) + + c.Release() +} + +func TestConnReleaseDestroysClosedConn(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + + err = c.Conn().Close(ctx) + require.NoError(t, err) + + assert.EqualValues(t, 1, pool.Stat().TotalConns()) + + c.Release() + waitForReleaseToComplete() + + // wait for the connection to actually be destroyed + for range 1000 { + if pool.Stat().TotalConns() == 0 { + break + } + time.Sleep(time.Millisecond) + } + + assert.EqualValues(t, 0, pool.Stat().TotalConns()) +} + +func TestConnPoolQueryConcurrentLoad(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + n := 100 + done := make(chan bool) + + for range n { + go func() { + defer func() { done <- true }() + testQuery(t, ctx, pool) + testQueryRow(t, ctx, pool) + }() + } + + for range n { + <-done + } +} + +func TestConnReleaseWhenBeginFail(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + tx, err := db.BeginTx(ctx, pgx.TxOptions{ + IsoLevel: pgx.TxIsoLevel("foo"), + }) + require.Error(t, err) + require.Zero(t, tx) + + require.EqualValues(t, 1, db.Stat().TotalConns()) + + var n int + require.NoError(t, db.QueryRow(ctx, "select 1").Scan(&n)) + require.EqualValues(t, 1, n) +} + +func TestConnDestroyedWhenBeginFailsFatally(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + controllerConn, err := pgx.Connect(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer controllerConn.Close(ctx) + pgxtest.SkipCockroachDB(t, controllerConn, "Server does not support pg_terminate_backend() (https://github.com/cockroachdb/cockroach/issues/35897)") + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + tx, err := db.BeginTx(ctx, pgx.TxOptions{BeginQuery: "select pg_terminate_backend(pg_backend_pid())"}) + require.Error(t, err) + require.Zero(t, tx) + + for range 1000 { + if db.Stat().TotalConns() == 0 { + break + } + time.Sleep(time.Millisecond) + } + + require.EqualValues(t, 0, db.Stat().TotalConns()) +} + +func TestTxBeginFuncNestedTransactionCommit(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + createSql := ` + drop table if exists pgxpooltx; + create temporary table pgxpooltx( + id integer, + unique (id) + ); + ` + + _, err = db.Exec(ctx, createSql) + require.NoError(t, err) + + defer func() { + db.Exec(ctx, "drop table pgxpooltx") + }() + + err = pgx.BeginFunc(ctx, db, func(db pgx.Tx) error { + _, err := db.Exec(ctx, "insert into pgxpooltx(id) values (1)") + require.NoError(t, err) + + err = pgx.BeginFunc(ctx, db, func(db pgx.Tx) error { + _, err := db.Exec(ctx, "insert into pgxpooltx(id) values (2)") + require.NoError(t, err) + + err = pgx.BeginFunc(ctx, db, func(db pgx.Tx) error { + _, err := db.Exec(ctx, "insert into pgxpooltx(id) values (3)") + require.NoError(t, err) + return nil + }) + require.NoError(t, err) + return nil + }) + require.NoError(t, err) + return nil + }) + require.NoError(t, err) + + var n int64 + err = db.QueryRow(ctx, "select count(*) from pgxpooltx").Scan(&n) + require.NoError(t, err) + require.EqualValues(t, 3, n) +} + +func TestTxBeginFuncNestedTransactionRollback(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + db, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer db.Close() + + createSql := ` + drop table if exists pgxpooltx; + create temporary table pgxpooltx( + id integer, + unique (id) + ); + ` + + _, err = db.Exec(ctx, createSql) + require.NoError(t, err) + + defer func() { + db.Exec(ctx, "drop table pgxpooltx") + }() + + err = pgx.BeginFunc(ctx, db, func(db pgx.Tx) error { + _, err := db.Exec(ctx, "insert into pgxpooltx(id) values (1)") + require.NoError(t, err) + + err = pgx.BeginFunc(ctx, db, func(db pgx.Tx) error { + _, err := db.Exec(ctx, "insert into pgxpooltx(id) values (2)") + require.NoError(t, err) + return errors.New("do a rollback") + }) + require.EqualError(t, err, "do a rollback") + + _, err = db.Exec(ctx, "insert into pgxpooltx(id) values (3)") + require.NoError(t, err) + + return nil + }) + require.NoError(t, err) + + var n int64 + err = db.QueryRow(ctx, "select count(*) from pgxpooltx").Scan(&n) + require.NoError(t, err) + require.EqualValues(t, 2, n) +} + +func TestIdempotentPoolClose(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + // Close the open pool. + require.NotPanics(t, func() { pool.Close() }) + + // Close the already closed pool. + require.NotPanics(t, func() { pool.Close() }) +} + +func TestConnectEagerlyReachesMinPoolSize(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MinConns = int32(12) + config.MaxConns = int32(15) + + var acquireAttempts atomic.Int64 + var connectAttempts atomic.Int64 + + config.PrepareConn = func(ctx context.Context, conn *pgx.Conn) (bool, error) { + acquireAttempts.Add(1) + return true, nil + } + config.BeforeConnect = func(ctx context.Context, cfg *pgx.ConnConfig) error { + connectAttempts.Add(1) + return nil + } + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + for range 500 { + time.Sleep(10 * time.Millisecond) + + stat := pool.Stat() + if stat.IdleConns() == 12 && stat.AcquireCount() == 0 && stat.TotalConns() == 12 && acquireAttempts.Load() == 0 && connectAttempts.Load() == 12 { + return + } + } + + t.Fatal("did not reach min pool size") +} + +func TestPoolSendBatchBatchCloseTwice(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + errChan := make(chan error) + testCount := 5000 + + for range testCount { + go func() { + batch := &pgx.Batch{} + batch.Queue("select 1") + batch.Queue("select 2") + + br := pool.SendBatch(ctx, batch) + defer br.Close() + + var err error + var n int32 + err = br.QueryRow().Scan(&n) + if err != nil { + errChan <- err + return + } + if n != 1 { + errChan <- fmt.Errorf("expected 1 got %v", n) + return + } + + err = br.QueryRow().Scan(&n) + if err != nil { + errChan <- err + return + } + if n != 2 { + errChan <- fmt.Errorf("expected 2 got %v", n) + return + } + + err = br.Close() + errChan <- err + }() + } + + for range testCount { + err := <-errChan + assert.NoError(t, err) + } +} + +func TestPoolAcquireDestroysExpiredIdleConn(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.MaxConnLifetime = 250 * time.Millisecond + config.HealthCheckPeriod = 10 * time.Second + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + // Acquire and release before expiry so the connection goes idle while still valid. + c, err := pool.Acquire(ctx) + require.NoError(t, err) + c.Release() + waitForReleaseToComplete() + + require.EqualValues(t, 1, pool.Stat().TotalConns()) + require.EqualValues(t, 0, pool.Stat().MaxLifetimeDestroyCount()) + + // Wait for the idle connection to expire. + time.Sleep(config.MaxConnLifetime + 100*time.Millisecond) + + // Acquire should pick up the expired idle conn, the new isExpired check in Acquire + // destroys it, and a fresh connection is created. + c, err = pool.Acquire(ctx) + require.NoError(t, err) + c.Release() + + // Give destroy time to settle. + time.Sleep(500 * time.Millisecond) + + stats := pool.Stat() + require.EqualValues(t, 1, stats.MaxLifetimeDestroyCount()) + require.EqualValues(t, 1, stats.TotalConns()) +} + +func TestPoolAcquireUnlimitedLifetimeDoesNotExpire(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + // A MaxConnLifetime of zero means connections never expire due to age. The + // acquire-time expiry check must not treat such connections as expired, + // otherwise Acquire destroys and recreates them in a loop and ultimately fails. + config.MaxConnLifetime = 0 + config.HealthCheckPeriod = 100 * time.Millisecond + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + // Establish a single connection and remember its identity. + c, err := pool.Acquire(ctx) + require.NoError(t, err) + firstConn := c.Conn() + c.Release() + waitForReleaseToComplete() + + require.EqualValues(t, 1, pool.Stat().TotalConns()) + + // Give the background health check several chances to (wrongly) reap the conn. + time.Sleep(500 * time.Millisecond) + + // Re-acquiring must succeed and hand back the very same connection: it was + // neither expired at acquire time nor destroyed by the health check. + for range 5 { + c, err = pool.Acquire(ctx) + require.NoError(t, err) + require.Same(t, firstConn, c.Conn()) + c.Release() + waitForReleaseToComplete() + } + + stats := pool.Stat() + require.EqualValues(t, 0, stats.MaxLifetimeDestroyCount()) + require.EqualValues(t, 1, stats.TotalConns()) +} + +func TestPoolAcquirePingTimeout(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + config.PingTimeout = 200 * time.Millisecond + config.ConnConfig.DialFunc = newDelayProxyDialFunc(500 * time.Millisecond) + + var conID *uint32 + // Only ping the connection with the original PID to force creation of a new connection + config.ShouldPing = func(_ context.Context, params pgxpool.ShouldPingParams) bool { + if conID != nil && params.Conn.PgConn().PID() == *conID { + return true + } + return false + } + + // Limit to a single connection to ensure the same connection is reused + config.MinConns = 1 + config.MaxConns = 1 + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + require.EqualValues(t, 1, pool.Stat().TotalConns()) + originalPID := c.Conn().PgConn().PID() + conID = &originalPID + + c.Release() + require.EqualValues(t, 1, pool.Stat().TotalConns()) + + c, err = pool.Acquire(ctx) + require.NoError(t, err) + require.EqualValues(t, 1, pool.Stat().TotalConns()) + newPID := c.Conn().PgConn().PID() + + c.Release() + + require.EqualValues(t, 1, pool.Stat().TotalConns()) + assert.Nil(t, ctx.Err()) + assert.NotEqualValues(t, originalPID, newPID, + "Expected new connection due to ping timeout, but got same connection") +} diff --git a/pgxpool/rows.go b/pgxpool/rows.go new file mode 100644 index 0000000..f834b7e --- /dev/null +++ b/pgxpool/rows.go @@ -0,0 +1,116 @@ +package pgxpool + +import ( + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" +) + +type errRows struct { + err error +} + +func (errRows) Close() {} +func (e errRows) Err() error { return e.err } +func (errRows) CommandTag() pgconn.CommandTag { return pgconn.CommandTag{} } +func (errRows) FieldDescriptions() []pgconn.FieldDescription { return nil } +func (errRows) Next() bool { return false } +func (e errRows) Scan(dest ...any) error { return e.err } +func (e errRows) Values() ([]any, error) { return nil, e.err } +func (e errRows) RawValues() [][]byte { return nil } +func (e errRows) Conn() *pgx.Conn { return nil } + +type errRow struct { + err error +} + +func (e errRow) Scan(dest ...any) error { return e.err } + +type poolRows struct { + r pgx.Rows + c *Conn + err error +} + +func (rows *poolRows) Close() { + rows.r.Close() + if rows.c != nil { + rows.c.Release() + rows.c = nil + } +} + +func (rows *poolRows) Err() error { + if rows.err != nil { + return rows.err + } + return rows.r.Err() +} + +func (rows *poolRows) CommandTag() pgconn.CommandTag { + return rows.r.CommandTag() +} + +func (rows *poolRows) FieldDescriptions() []pgconn.FieldDescription { + return rows.r.FieldDescriptions() +} + +func (rows *poolRows) Next() bool { + if rows.err != nil { + return false + } + + n := rows.r.Next() + if !n { + rows.Close() + } + return n +} + +func (rows *poolRows) Scan(dest ...any) error { + err := rows.r.Scan(dest...) + if err != nil { + rows.Close() + } + return err +} + +func (rows *poolRows) Values() ([]any, error) { + values, err := rows.r.Values() + if err != nil { + rows.Close() + } + return values, err +} + +func (rows *poolRows) RawValues() [][]byte { + return rows.r.RawValues() +} + +func (rows *poolRows) Conn() *pgx.Conn { + return rows.r.Conn() +} + +type poolRow struct { + r pgx.Row + c *Conn + err error +} + +func (row *poolRow) Scan(dest ...any) error { + if row.err != nil { + return row.err + } + + panicked := true + defer func() { + if panicked && row.c != nil { + row.c.Release() + } + }() + err := row.r.Scan(dest...) + panicked = false + if row.c != nil { + row.c.Release() + } + return err +} diff --git a/pgxpool/stat.go b/pgxpool/stat.go new file mode 100644 index 0000000..e02b6ac --- /dev/null +++ b/pgxpool/stat.go @@ -0,0 +1,91 @@ +package pgxpool + +import ( + "time" + + "github.com/jackc/puddle/v2" +) + +// Stat is a snapshot of Pool statistics. +type Stat struct { + s *puddle.Stat + newConnsCount int64 + lifetimeDestroyCount int64 + idleDestroyCount int64 +} + +// AcquireCount returns the cumulative count of successful acquires from the pool. +func (s *Stat) AcquireCount() int64 { + return s.s.AcquireCount() +} + +// AcquireDuration returns the total duration of all successful acquires from +// the pool. +func (s *Stat) AcquireDuration() time.Duration { + return s.s.AcquireDuration() +} + +// AcquiredConns returns the number of currently acquired connections in the pool. +func (s *Stat) AcquiredConns() int32 { + return s.s.AcquiredResources() +} + +// CanceledAcquireCount returns the cumulative count of acquires from the pool +// that were canceled by a context. +func (s *Stat) CanceledAcquireCount() int64 { + return s.s.CanceledAcquireCount() +} + +// ConstructingConns returns the number of conns with construction in progress in +// the pool. +func (s *Stat) ConstructingConns() int32 { + return s.s.ConstructingResources() +} + +// EmptyAcquireCount returns the cumulative count of successful acquires from the pool +// that waited for a resource to be released or constructed because the pool was +// empty. +func (s *Stat) EmptyAcquireCount() int64 { + return s.s.EmptyAcquireCount() +} + +// IdleConns returns the number of currently idle conns in the pool. +func (s *Stat) IdleConns() int32 { + return s.s.IdleResources() +} + +// MaxConns returns the maximum size of the pool. +func (s *Stat) MaxConns() int32 { + return s.s.MaxResources() +} + +// TotalConns returns the total number of resources currently in the pool. +// The value is the sum of ConstructingConns, AcquiredConns, and +// IdleConns. +func (s *Stat) TotalConns() int32 { + return s.s.TotalResources() +} + +// NewConnsCount returns the cumulative count of new connections opened. +func (s *Stat) NewConnsCount() int64 { + return s.newConnsCount +} + +// MaxLifetimeDestroyCount returns the cumulative count of connections destroyed +// because they exceeded MaxConnLifetime. +func (s *Stat) MaxLifetimeDestroyCount() int64 { + return s.lifetimeDestroyCount +} + +// MaxIdleDestroyCount returns the cumulative count of connections destroyed because +// they exceeded MaxConnIdleTime. +func (s *Stat) MaxIdleDestroyCount() int64 { + return s.idleDestroyCount +} + +// EmptyAcquireWaitTime returns the cumulative time waited for successful acquires +// from the pool for a resource to be released or constructed because the pool was +// empty. +func (s *Stat) EmptyAcquireWaitTime() time.Duration { + return s.s.EmptyAcquireWaitTime() +} diff --git a/pgxpool/tracer.go b/pgxpool/tracer.go new file mode 100644 index 0000000..78b9d15 --- /dev/null +++ b/pgxpool/tracer.go @@ -0,0 +1,33 @@ +package pgxpool + +import ( + "context" + + "github.com/jackc/pgx/v5" +) + +// AcquireTracer traces Acquire. +type AcquireTracer interface { + // TraceAcquireStart is called at the beginning of Acquire. + // The returned context is used for the rest of the call and will be passed to the TraceAcquireEnd. + TraceAcquireStart(ctx context.Context, pool *Pool, data TraceAcquireStartData) context.Context + // TraceAcquireEnd is called when a connection has been acquired. + TraceAcquireEnd(ctx context.Context, pool *Pool, data TraceAcquireEndData) +} + +type TraceAcquireStartData struct{} + +type TraceAcquireEndData struct { + Conn *pgx.Conn + Err error +} + +// ReleaseTracer traces Release. +type ReleaseTracer interface { + // TraceRelease is called at the beginning of Release. + TraceRelease(pool *Pool, data TraceReleaseData) +} + +type TraceReleaseData struct { + Conn *pgx.Conn +} diff --git a/pgxpool/tracer_test.go b/pgxpool/tracer_test.go new file mode 100644 index 0000000..10724d9 --- /dev/null +++ b/pgxpool/tracer_test.go @@ -0,0 +1,130 @@ +package pgxpool_test + +import ( + "context" + "os" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +type testTracer struct { + traceAcquireStart func(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireStartData) context.Context + traceAcquireEnd func(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireEndData) + traceRelease func(pool *pgxpool.Pool, data pgxpool.TraceReleaseData) +} + +type ctxKey string + +func (tt *testTracer) TraceAcquireStart(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireStartData) context.Context { + if tt.traceAcquireStart != nil { + return tt.traceAcquireStart(ctx, pool, data) + } + return ctx +} + +func (tt *testTracer) TraceAcquireEnd(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireEndData) { + if tt.traceAcquireEnd != nil { + tt.traceAcquireEnd(ctx, pool, data) + } +} + +func (tt *testTracer) TraceRelease(pool *pgxpool.Pool, data pgxpool.TraceReleaseData) { + if tt.traceRelease != nil { + tt.traceRelease(pool, data) + } +} + +func (tt *testTracer) TraceQueryStart(ctx context.Context, conn *pgx.Conn, data pgx.TraceQueryStartData) context.Context { + return ctx +} + +func (tt *testTracer) TraceQueryEnd(ctx context.Context, conn *pgx.Conn, data pgx.TraceQueryEndData) { +} + +func TestTraceAcquire(t *testing.T) { + t.Parallel() + + tracer := &testTracer{} + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + config.ConnConfig.Tracer = tracer + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + traceAcquireStartCalled := false + tracer.traceAcquireStart = func(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireStartData) context.Context { + traceAcquireStartCalled = true + require.NotNil(t, pool) + return context.WithValue(ctx, ctxKey("fromTraceAcquireStart"), "foo") + } + + traceAcquireEndCalled := false + tracer.traceAcquireEnd = func(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireEndData) { + traceAcquireEndCalled = true + require.Equal(t, "foo", ctx.Value(ctxKey("fromTraceAcquireStart"))) + require.NotNil(t, pool) + require.NotNil(t, data.Conn) + require.NoError(t, data.Err) + } + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + defer c.Release() + require.True(t, traceAcquireStartCalled) + require.True(t, traceAcquireEndCalled) + + traceAcquireStartCalled = false + traceAcquireEndCalled = false + tracer.traceAcquireEnd = func(ctx context.Context, pool *pgxpool.Pool, data pgxpool.TraceAcquireEndData) { + traceAcquireEndCalled = true + require.NotNil(t, pool) + require.Nil(t, data.Conn) + require.Error(t, data.Err) + } + + ctx, cancel = context.WithCancel(ctx) + cancel() + _, err = pool.Acquire(ctx) + require.ErrorIs(t, err, context.Canceled) + require.True(t, traceAcquireStartCalled) + require.True(t, traceAcquireEndCalled) +} + +func TestTraceRelease(t *testing.T) { + t.Parallel() + + tracer := &testTracer{} + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + config, err := pgxpool.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + config.ConnConfig.Tracer = tracer + + pool, err := pgxpool.NewWithConfig(ctx, config) + require.NoError(t, err) + defer pool.Close() + + traceReleaseCalled := false + tracer.traceRelease = func(pool *pgxpool.Pool, data pgxpool.TraceReleaseData) { + traceReleaseCalled = true + require.NotNil(t, pool) + require.NotNil(t, data.Conn) + } + + c, err := pool.Acquire(ctx) + require.NoError(t, err) + c.Release() + require.True(t, traceReleaseCalled) +} diff --git a/pgxpool/tx.go b/pgxpool/tx.go new file mode 100644 index 0000000..b49e7f4 --- /dev/null +++ b/pgxpool/tx.go @@ -0,0 +1,83 @@ +package pgxpool + +import ( + "context" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" +) + +// Tx represents a database transaction acquired from a Pool. +type Tx struct { + t pgx.Tx + c *Conn +} + +// Begin starts a pseudo nested transaction implemented with a savepoint. +func (tx *Tx) Begin(ctx context.Context) (pgx.Tx, error) { + return tx.t.Begin(ctx) +} + +// Commit commits the transaction and returns the associated connection back to the Pool. Commit will return an error +// where errors.Is(ErrTxClosed) is true if the Tx is already closed, but is otherwise safe to call multiple times. If +// the commit fails with a rollback status (e.g. the transaction was already in a broken state) then ErrTxCommitRollback +// will be returned. +func (tx *Tx) Commit(ctx context.Context) error { + err := tx.t.Commit(ctx) + if tx.c != nil { + tx.c.Release() + tx.c = nil + } + return err +} + +// Rollback rolls back the transaction and returns the associated connection back to the Pool. Rollback will return +// where an error where errors.Is(ErrTxClosed) is true if the Tx is already closed, but is otherwise safe to call +// multiple times. Hence, defer tx.Rollback() is safe even if tx.Commit() will be called first in a non-error condition. +func (tx *Tx) Rollback(ctx context.Context) error { + err := tx.t.Rollback(ctx) + if tx.c != nil { + tx.c.Release() + tx.c = nil + } + return err +} + +func (tx *Tx) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) { + return tx.t.CopyFrom(ctx, tableName, columnNames, rowSrc) +} + +func (tx *Tx) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults { + return tx.t.SendBatch(ctx, b) +} + +func (tx *Tx) LargeObjects() pgx.LargeObjects { + return tx.t.LargeObjects() +} + +// Prepare creates a prepared statement with name and sql. If the name is empty, +// an anonymous prepared statement will be used. sql can contain placeholders +// for bound parameters. These placeholders are referenced positionally as $1, $2, etc. +// +// Prepare is idempotent; i.e. it is safe to call Prepare multiple times with the same +// name and sql arguments. This allows a code path to Prepare and Query/Exec without +// needing to first check whether the statement has already been prepared. +func (tx *Tx) Prepare(ctx context.Context, name, sql string) (*pgconn.StatementDescription, error) { + return tx.t.Prepare(ctx, name, sql) +} + +func (tx *Tx) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) { + return tx.t.Exec(ctx, sql, arguments...) +} + +func (tx *Tx) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) { + return tx.t.Query(ctx, sql, args...) +} + +func (tx *Tx) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row { + return tx.t.QueryRow(ctx, sql, args...) +} + +func (tx *Tx) Conn() *pgx.Conn { + return tx.t.Conn() +} diff --git a/pgxpool/tx_test.go b/pgxpool/tx_test.go new file mode 100644 index 0000000..e1611e6 --- /dev/null +++ b/pgxpool/tx_test.go @@ -0,0 +1,96 @@ +package pgxpool_test + +import ( + "context" + "os" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +func TestTxExec(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + defer tx.Rollback(ctx) + + testExec(t, ctx, tx) +} + +func TestTxQuery(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + defer tx.Rollback(ctx) + + testQuery(t, ctx, tx) +} + +func TestTxQueryRow(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + defer tx.Rollback(ctx) + + testQueryRow(t, ctx, tx) +} + +func TestTxSendBatch(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + defer tx.Rollback(ctx) + + testSendBatch(t, ctx, tx) +} + +func TestTxCopyFrom(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + pool, err := pgxpool.New(ctx, os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + defer pool.Close() + + tx, err := pool.Begin(ctx) + require.NoError(t, err) + defer tx.Rollback(ctx) + + testCopyFrom(t, ctx, tx) +} diff --git a/puddle/.devcontainer/devcontainer.json b/puddle/.devcontainer/devcontainer.json new file mode 100644 index 0000000..18f90b9 --- /dev/null +++ b/puddle/.devcontainer/devcontainer.json @@ -0,0 +1,22 @@ +// For format details, see https://aka.ms/devcontainer.json. For config options, see the +// README at: https://github.com/devcontainers/templates/tree/main/src/go . +{ + "name": "puddle", + // Or use a Dockerfile or Docker Compose file. More info: https://containers.dev/guide/dockerfile + "image": "mcr.microsoft.com/devcontainers/go:2-1.25-trixie" + + // Features to add to the dev container. More info: https://containers.dev/features. + // "features": {}, + + // Use 'forwardPorts' to make a list of ports inside the container available locally. + // "forwardPorts": [], + + // Use 'postCreateCommand' to run commands after the container is created. + // "postCreateCommand": "go version", + + // Configure tool-specific properties. + // "customizations": {}, + + // Uncomment to connect as root instead. More info: https://aka.ms/dev-containers-non-root. + // "remoteUser": "root" +} diff --git a/puddle/CHANGELOG.md b/puddle/CHANGELOG.md new file mode 100644 index 0000000..d0d202c --- /dev/null +++ b/puddle/CHANGELOG.md @@ -0,0 +1,79 @@ +# 2.2.2 (September 10, 2024) + +* Add empty acquire time to stats (Maxim Ivanov) +* Stop importing nanotime from runtime via linkname (maypok86) + +# 2.2.1 (July 15, 2023) + +* Fix: CreateResource cannot overflow pool. This changes documented behavior of CreateResource. Previously, + CreateResource could create a resource even if the pool was full. This could cause the pool to overflow. While this + was documented, it was documenting incorrect behavior. CreateResource now returns an error if the pool is full. + +# 2.2.0 (February 11, 2023) + +* Use Go 1.19 atomics and drop go.uber.org/atomic dependency + +# 2.1.2 (November 12, 2022) + +* Restore support to Go 1.18 via go.uber.org/atomic + +# 2.1.1 (November 11, 2022) + +* Fix create resource concurrently with Stat call race + +# 2.1.0 (October 28, 2022) + +* Concurrency control is now implemented with a semaphore. This simplifies some internal logic, resolves a few error conditions (including a deadlock), and improves performance. (Jan Dubsky) +* Go 1.19 is now required for the improved atomic support. + +# 2.0.1 (October 28, 2022) + +* Fix race condition when Close is called concurrently with multiple constructors + +# 2.0.0 (September 17, 2022) + +* Use generics instead of interface{} (Столяров Владимир Алексеевич) +* Add Reset +* Do not cancel resource construction when Acquire is canceled +* NewPool takes Config + +# 1.3.0 (August 27, 2022) + +* Acquire creates resources in background to allow creation to continue after Acquire is canceled (James Hartig) + +# 1.2.1 (December 2, 2021) + +* TryAcquire now does not block when background constructing resource + +# 1.2.0 (November 20, 2021) + +* Add TryAcquire (A. Jensen) +* Fix: remove memory leak / unintentionally pinned memory when shrinking slices (Alexander Staubo) +* Fix: Do not leave pool locked after panic from nil context + +# 1.1.4 (September 11, 2021) + +* Fix: Deadlock in CreateResource if pool was closed during resource acquisition (Dmitriy Matrenichev) + +# 1.1.3 (December 3, 2020) + +* Fix: Failed resource creation could cause concurrent Acquire to hang. (Evgeny Vanslov) + +# 1.1.2 (September 26, 2020) + +* Fix: Resource.Destroy no longer removes itself from the pool before its destructor has completed. +* Fix: Prevent crash when pool is closed while resource is being created. + +# 1.1.1 (April 2, 2020) + +* Pool.Close can be safely called multiple times +* AcquireAllIDle immediately returns nil if pool is closed +* CreateResource checks if pool is closed before taking any action +* Fix potential race condition when CreateResource and Close are called concurrently. CreateResource now checks if pool is closed before adding newly created resource to pool. + +# 1.1.0 (February 5, 2020) + +* Use runtime.nanotime for faster tracking of acquire time and last usage time. +* Track resource idle time to enable client health check logic. (Patrick Ellul) +* Add CreateResource to construct a new resource without acquiring it. (Patrick Ellul) +* Fix deadlock race when acquire is cancelled. (Michael Tharp) diff --git a/puddle/LICENSE b/puddle/LICENSE new file mode 100644 index 0000000..bcc286c --- /dev/null +++ b/puddle/LICENSE @@ -0,0 +1,22 @@ +Copyright (c) 2018 Jack Christensen + +MIT License + +Permission is hereby granted, free of charge, to any person obtaining +a copy of this software and associated documentation files (the +"Software"), to deal in the Software without restriction, including +without limitation the rights to use, copy, modify, merge, publish, +distribute, sublicense, and/or sell copies of the Software, and to +permit persons to whom the Software is furnished to do so, subject to +the following conditions: + +The above copyright notice and this permission notice shall be +included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, +EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF +MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND +NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE +LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION +OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION +WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. diff --git a/puddle/README.md b/puddle/README.md new file mode 100644 index 0000000..402e6c8 --- /dev/null +++ b/puddle/README.md @@ -0,0 +1,96 @@ +[![Go Reference](https://pkg.go.dev/badge/github.com/jackc/puddle/v2.svg)](https://pkg.go.dev/github.com/jackc/puddle/v2) +![Build Status](https://github.com/jackc/puddle/actions/workflows/ci.yml/badge.svg) + +# Puddle + +Puddle is a tiny generic resource pool library for Go that uses the standard +context library to signal cancellation of acquires. It is designed to contain +the minimum functionality required for a resource pool. It can be used directly +or it can be used as the base for a domain specific resource pool. For example, +a database connection pool may use puddle internally and implement health checks +and keep-alive behavior without needing to implement any concurrent code of its +own. + +## Features + +* Acquire cancellation via context standard library +* Statistics API for monitoring pool pressure +* No dependencies outside of standard library and golang.org/x/sync +* High performance +* 100% test coverage of reachable code + +## Example Usage + +```go +package main + +import ( + "context" + "log" + "net" + + "github.com/jackc/puddle/v2" +) + +func main() { + constructor := func(context.Context) (net.Conn, error) { + return net.Dial("tcp", "127.0.0.1:8080") + } + destructor := func(value net.Conn) { + value.Close() + } + maxPoolSize := int32(10) + + pool, err := puddle.NewPool(&puddle.Config[net.Conn]{Constructor: constructor, Destructor: destructor, MaxSize: maxPoolSize}) + if err != nil { + log.Fatal(err) + } + + // Acquire resource from the pool. + res, err := pool.Acquire(context.Background()) + if err != nil { + log.Fatal(err) + } + + // Use resource. + _, err = res.Value().Write([]byte{1}) + if err != nil { + log.Fatal(err) + } + + // Release when done. + res.Release() +} +``` + +## Status + +Puddle is stable and feature complete. + +* Bug reports and fixes are welcome. +* New features will usually not be accepted if they can be feasibly implemented in a wrapper. +* Performance optimizations will usually not be accepted unless the performance issue rises to the level of a bug. + +## Supported Go Versions + +puddle supports the same versions of Go that are supported by the Go project. For [Go](https://golang.org/doc/devel/release.html#policy) that is the two most recent major releases. This means puddle supports Go 1.19 and higher. + +## Differences with Go sync.Pool + +They are intended for entirely different types of resources: + +* [sync.Pool](https://pkg.go.dev/sync#Pool) would generally be used for in memory objects. +* Puddle would generally be used for handles to external objects such as connections, file handles, etc. + +Specific differences: + +* sync.Pool does not have a way to limit max resources in pool. +* sync.Pool does not have a way to ensure at least min resources in pool. +* sync.Pool can drop resources in pool at any time - not ideal for expensive to create connections. +* sync.Pool does not have a cleanup / release function. Resources are GCed without a chance to close cleanly. +* sync.Pool does not have a way to handle errors creating the resource +* sync.Pool does not support context to limit time to wait creating a resource + +## License + +MIT diff --git a/puddle/context.go b/puddle/context.go new file mode 100644 index 0000000..e19d2a6 --- /dev/null +++ b/puddle/context.go @@ -0,0 +1,24 @@ +package puddle + +import ( + "context" + "time" +) + +// valueCancelCtx combines two contexts into one. One context is used for values and the other is used for cancellation. +type valueCancelCtx struct { + valueCtx context.Context + cancelCtx context.Context +} + +func (ctx *valueCancelCtx) Deadline() (time.Time, bool) { return ctx.cancelCtx.Deadline() } +func (ctx *valueCancelCtx) Done() <-chan struct{} { return ctx.cancelCtx.Done() } +func (ctx *valueCancelCtx) Err() error { return ctx.cancelCtx.Err() } +func (ctx *valueCancelCtx) Value(key any) any { return ctx.valueCtx.Value(key) } + +func newValueCancelCtx(valueCtx, cancelContext context.Context) context.Context { + return &valueCancelCtx{ + valueCtx: valueCtx, + cancelCtx: cancelContext, + } +} diff --git a/puddle/doc.go b/puddle/doc.go new file mode 100644 index 0000000..818e4a6 --- /dev/null +++ b/puddle/doc.go @@ -0,0 +1,11 @@ +// Package puddle is a generic resource pool with type-parametrized api. +/* + +Puddle is a tiny generic resource pool library for Go that uses the standard +context library to signal cancellation of acquires. It is designed to contain +the minimum functionality a resource pool needs that cannot be implemented +without concurrency concerns. For example, a database connection pool may use +puddle internally and implement health checks and keep-alive behavior without +needing to implement any concurrent code of its own. +*/ +package puddle diff --git a/puddle/export_test.go b/puddle/export_test.go new file mode 100644 index 0000000..36e8df6 --- /dev/null +++ b/puddle/export_test.go @@ -0,0 +1,9 @@ +package puddle + +import "context" + +func (p *Pool[T]) AcquireRaw(ctx context.Context) (*Resource[T], error) { + return p.acquire(ctx) +} + +var AcquireSemAll = acquireSemAll diff --git a/puddle/go.mod b/puddle/go.mod new file mode 100644 index 0000000..b66cc9c --- /dev/null +++ b/puddle/go.mod @@ -0,0 +1,14 @@ +module github.com/jackc/puddle/v2 + +go 1.19 + +require ( + github.com/stretchr/testify v1.8.1 + golang.org/x/sync v0.1.0 +) + +require ( + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/puddle/go.sum b/puddle/go.sum new file mode 100644 index 0000000..96e82f8 --- /dev/null +++ b/puddle/go.sum @@ -0,0 +1,19 @@ +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/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/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= +github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk= +github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= +golang.org/x/sync v0.1.0 h1:wsuoTGHzEhffawBOhz5CYhcrV4IdKZbEyZjBMuTp12o= +golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +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/puddle/internal/genstack/gen_stack.go b/puddle/internal/genstack/gen_stack.go new file mode 100644 index 0000000..7e4660c --- /dev/null +++ b/puddle/internal/genstack/gen_stack.go @@ -0,0 +1,85 @@ +package genstack + +// GenStack implements a generational stack. +// +// GenStack works as common stack except for the fact that all elements in the +// older generation are guaranteed to be popped before any element in the newer +// generation. New elements are always pushed to the current (newest) +// generation. +// +// We could also say that GenStack behaves as a stack in case of a single +// generation, but it behaves as a queue of individual generation stacks. +type GenStack[T any] struct { + // We can represent arbitrary number of generations using 2 stacks. The + // new stack stores all new pushes and the old stack serves all reads. + // Old stack can represent multiple generations. If old == new, then all + // elements pushed in previous (not current) generations have already + // been popped. + + old *stack[T] + new *stack[T] +} + +// NewGenStack creates a new empty GenStack. +func NewGenStack[T any]() *GenStack[T] { + s := &stack[T]{} + return &GenStack[T]{ + old: s, + new: s, + } +} + +func (s *GenStack[T]) Pop() (T, bool) { + // Pushes always append to the new stack, so if the old once becomes + // empty, it will remail empty forever. + if s.old.len() == 0 && s.old != s.new { + s.old = s.new + } + + if s.old.len() == 0 { + var zero T + return zero, false + } + + return s.old.pop(), true +} + +// Push pushes a new element at the top of the stack. +func (s *GenStack[T]) Push(v T) { s.new.push(v) } + +// NextGen starts a new stack generation. +func (s *GenStack[T]) NextGen() { + if s.old == s.new { + s.new = &stack[T]{} + return + } + + // We need to pop from the old stack to the top of the new stack. Let's + // have an example: + // + // Old: 4 3 2 1 + // New: 8 7 6 5 + // PopOrder: 1 2 3 4 5 6 7 8 + // + // + // To preserve pop order, we have to take all elements from the old + // stack and push them to the top of new stack: + // + // New: 8 7 6 5 4 3 2 1 + // + s.new.push(s.old.takeAll()...) + + // We have the old stack allocated and empty, so why not to reuse it as + // new new stack. + s.old, s.new = s.new, s.old +} + +// Len returns number of elements in the stack. +func (s *GenStack[T]) Len() int { + l := s.old.len() + if s.old != s.new { + l += s.new.len() + } + + return l +} diff --git a/puddle/internal/genstack/gen_stack_test.go b/puddle/internal/genstack/gen_stack_test.go new file mode 100644 index 0000000..519bd3b --- /dev/null +++ b/puddle/internal/genstack/gen_stack_test.go @@ -0,0 +1,90 @@ +package genstack + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func requirePopEmpty[T any](t testing.TB, s *GenStack[T]) { + v, ok := s.Pop() + require.False(t, ok) + require.Zero(t, v) +} + +func requirePop[T any](t testing.TB, s *GenStack[T], expected T) { + v, ok := s.Pop() + require.True(t, ok) + require.Equal(t, expected, v) +} + +func TestGenStack_Empty(t *testing.T) { + s := NewGenStack[int]() + requirePopEmpty(t, s) +} + +func TestGenStack_SingleGen(t *testing.T) { + r := require.New(t) + s := NewGenStack[int]() + + s.Push(1) + s.Push(2) + r.Equal(2, s.Len()) + + requirePop(t, s, 2) + requirePop(t, s, 1) + requirePopEmpty(t, s) +} + +func TestGenStack_TwoGen(t *testing.T) { + r := require.New(t) + s := NewGenStack[int]() + + s.Push(3) + s.Push(4) + s.Push(5) + r.Equal(3, s.Len()) + s.NextGen() + r.Equal(3, s.Len()) + s.Push(6) + s.Push(7) + r.Equal(5, s.Len()) + + requirePop(t, s, 5) + requirePop(t, s, 4) + requirePop(t, s, 3) + requirePop(t, s, 7) + requirePop(t, s, 6) + requirePopEmpty(t, s) +} + +func TestGenStack_MuptiGen(t *testing.T) { + r := require.New(t) + s := NewGenStack[int]() + + s.Push(10) + s.Push(11) + s.Push(12) + r.Equal(3, s.Len()) + s.NextGen() + r.Equal(3, s.Len()) + s.Push(13) + s.Push(14) + r.Equal(5, s.Len()) + s.NextGen() + r.Equal(5, s.Len()) + s.Push(15) + s.Push(16) + s.Push(17) + r.Equal(8, s.Len()) + + requirePop(t, s, 12) + requirePop(t, s, 11) + requirePop(t, s, 10) + requirePop(t, s, 14) + requirePop(t, s, 13) + requirePop(t, s, 17) + requirePop(t, s, 16) + requirePop(t, s, 15) + requirePopEmpty(t, s) +} diff --git a/puddle/internal/genstack/stack.go b/puddle/internal/genstack/stack.go new file mode 100644 index 0000000..dbced0c --- /dev/null +++ b/puddle/internal/genstack/stack.go @@ -0,0 +1,39 @@ +package genstack + +// stack is a wrapper around an array implementing a stack. +// +// We cannot use slice to represent the stack because append might change the +// pointer value of the slice. That would be an issue in GenStack +// implementation. +type stack[T any] struct { + arr []T +} + +// push pushes a new element at the top of a stack. +func (s *stack[T]) push(vs ...T) { s.arr = append(s.arr, vs...) } + +// pop pops the stack top-most element. +// +// If stack length is zero, this method panics. +func (s *stack[T]) pop() T { + idx := s.len() - 1 + val := s.arr[idx] + + // Avoid memory leak + var zero T + s.arr[idx] = zero + + s.arr = s.arr[:idx] + return val +} + +// takeAll returns all elements in the stack in order as they are stored - i.e. +// the top-most stack element is the last one. +func (s *stack[T]) takeAll() []T { + arr := s.arr + s.arr = nil + return arr +} + +// len returns number of elements in the stack. +func (s *stack[T]) len() int { return len(s.arr) } diff --git a/puddle/nanotime.go b/puddle/nanotime.go new file mode 100644 index 0000000..8a5351a --- /dev/null +++ b/puddle/nanotime.go @@ -0,0 +1,16 @@ +package puddle + +import "time" + +// nanotime returns the time in nanoseconds since process start. +// +// This approach, described at +// https://github.com/golang/go/issues/61765#issuecomment-1672090302, +// is fast, monotonic, and portable, and avoids the previous +// dependence on runtime.nanotime using the (unsafe) linkname hack. +// In particular, time.Since does less work than time.Now. +func nanotime() int64 { + return time.Since(globalStart).Nanoseconds() +} + +var globalStart = time.Now() diff --git a/puddle/pool.go b/puddle/pool.go new file mode 100644 index 0000000..722edb1 --- /dev/null +++ b/puddle/pool.go @@ -0,0 +1,728 @@ +package puddle + +import ( + "context" + "errors" + "math/bits" + "sync" + "sync/atomic" + "time" + + "github.com/jackc/puddle/v2/internal/genstack" + "golang.org/x/sync/semaphore" +) + +const ( + resourceStatusConstructing = 0 + resourceStatusIdle = iota + resourceStatusAcquired = iota + resourceStatusHijacked = iota +) + +// ErrClosedPool occurs on an attempt to acquire a connection from a closed pool +// or a pool that is closed while the acquire is waiting. +var ErrClosedPool = errors.New("closed pool") + +// ErrNotAvailable occurs on an attempt to acquire a resource from a pool +// that is at maximum capacity and has no available resources. +var ErrNotAvailable = errors.New("resource not available") + +// Constructor is a function called by the pool to construct a resource. +type Constructor[T any] func(ctx context.Context) (res T, err error) + +// Destructor is a function called by the pool to destroy a resource. +type Destructor[T any] func(res T) + +// Resource is the resource handle returned by acquiring from the pool. +type Resource[T any] struct { + value T + pool *Pool[T] + creationTime time.Time + lastUsedNano int64 + poolResetCount int + status byte +} + +// Value returns the resource value. +func (res *Resource[T]) Value() T { + if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) { + panic("tried to access resource that is not acquired or hijacked") + } + return res.value +} + +// Release returns the resource to the pool. res must not be subsequently used. +func (res *Resource[T]) Release() { + if res.status != resourceStatusAcquired { + panic("tried to release resource that is not acquired") + } + res.pool.releaseAcquiredResource(res, nanotime()) +} + +// ReleaseUnused returns the resource to the pool without updating when it was last used used. i.e. LastUsedNanotime +// will not change. res must not be subsequently used. +func (res *Resource[T]) ReleaseUnused() { + if res.status != resourceStatusAcquired { + panic("tried to release resource that is not acquired") + } + res.pool.releaseAcquiredResource(res, res.lastUsedNano) +} + +// Destroy returns the resource to the pool for destruction. res must not be +// subsequently used. +func (res *Resource[T]) Destroy() { + if res.status != resourceStatusAcquired { + panic("tried to destroy resource that is not acquired") + } + go res.pool.destroyAcquiredResource(res) +} + +// Hijack assumes ownership of the resource from the pool. Caller is responsible +// for cleanup of resource value. +func (res *Resource[T]) Hijack() { + if res.status != resourceStatusAcquired { + panic("tried to hijack resource that is not acquired") + } + res.pool.hijackAcquiredResource(res) +} + +// CreationTime returns when the resource was created by the pool. +func (res *Resource[T]) CreationTime() time.Time { + if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) { + panic("tried to access resource that is not acquired or hijacked") + } + return res.creationTime +} + +// LastUsedNanotime returns when Release was last called on the resource measured in nanoseconds from an arbitrary time +// (a monotonic time). Returns creation time if Release has never been called. This is only useful to compare with +// other calls to LastUsedNanotime. In almost all cases, IdleDuration should be used instead. +func (res *Resource[T]) LastUsedNanotime() int64 { + if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) { + panic("tried to access resource that is not acquired or hijacked") + } + + return res.lastUsedNano +} + +// IdleDuration returns the duration since Release was last called on the resource. This is equivalent to subtracting +// LastUsedNanotime to the current nanotime. +func (res *Resource[T]) IdleDuration() time.Duration { + if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) { + panic("tried to access resource that is not acquired or hijacked") + } + + return time.Duration(nanotime() - res.lastUsedNano) +} + +// Pool is a concurrency-safe resource pool. +type Pool[T any] struct { + // mux is the pool internal lock. Any modification of shared state of + // the pool (but Acquires of acquireSem) must be performed only by + // holder of the lock. Long running operations are not allowed when mux + // is held. + mux sync.Mutex + // acquireSem provides an allowance to acquire a resource. + // + // Releases are allowed only when caller holds mux. Acquires have to + // happen before mux is locked (doesn't apply to semaphore.TryAcquire in + // AcquireAllIdle). + acquireSem *semaphore.Weighted + destructWG sync.WaitGroup + + allResources resList[T] + idleResources *genstack.GenStack[*Resource[T]] + + constructor Constructor[T] + destructor Destructor[T] + maxSize int32 + + acquireCount int64 + acquireDuration time.Duration + emptyAcquireCount int64 + emptyAcquireWaitTime time.Duration + canceledAcquireCount atomic.Int64 + + resetCount int + + baseAcquireCtx context.Context + cancelBaseAcquireCtx context.CancelFunc + closed bool +} + +type Config[T any] struct { + Constructor Constructor[T] + Destructor Destructor[T] + MaxSize int32 +} + +// NewPool creates a new pool. Returns an error iff MaxSize is less than 1. +func NewPool[T any](config *Config[T]) (*Pool[T], error) { + if config.MaxSize < 1 { + return nil, errors.New("MaxSize must be >= 1") + } + + baseAcquireCtx, cancelBaseAcquireCtx := context.WithCancel(context.Background()) + + return &Pool[T]{ + acquireSem: semaphore.NewWeighted(int64(config.MaxSize)), + idleResources: genstack.NewGenStack[*Resource[T]](), + maxSize: config.MaxSize, + constructor: config.Constructor, + destructor: config.Destructor, + baseAcquireCtx: baseAcquireCtx, + cancelBaseAcquireCtx: cancelBaseAcquireCtx, + }, nil +} + +// Close destroys all resources in the pool and rejects future Acquire calls. +// Blocks until all resources are returned to pool and destroyed. +func (p *Pool[T]) Close() { + defer p.destructWG.Wait() + + p.mux.Lock() + defer p.mux.Unlock() + + if p.closed { + return + } + p.closed = true + p.cancelBaseAcquireCtx() + + for res, ok := p.idleResources.Pop(); ok; res, ok = p.idleResources.Pop() { + p.allResources.remove(res) + go p.destructResourceValue(res.value) + } +} + +// Stat is a snapshot of Pool statistics. +type Stat struct { + constructingResources int32 + acquiredResources int32 + idleResources int32 + maxResources int32 + acquireCount int64 + acquireDuration time.Duration + emptyAcquireCount int64 + emptyAcquireWaitTime time.Duration + canceledAcquireCount int64 +} + +// TotalResources returns the total number of resources currently in the pool. +// The value is the sum of ConstructingResources, AcquiredResources, and +// IdleResources. +func (s *Stat) TotalResources() int32 { + return s.constructingResources + s.acquiredResources + s.idleResources +} + +// ConstructingResources returns the number of resources with construction in progress in +// the pool. +func (s *Stat) ConstructingResources() int32 { + return s.constructingResources +} + +// AcquiredResources returns the number of currently acquired resources in the pool. +func (s *Stat) AcquiredResources() int32 { + return s.acquiredResources +} + +// IdleResources returns the number of currently idle resources in the pool. +func (s *Stat) IdleResources() int32 { + return s.idleResources +} + +// MaxResources returns the maximum size of the pool. +func (s *Stat) MaxResources() int32 { + return s.maxResources +} + +// AcquireCount returns the cumulative count of successful acquires from the pool. +func (s *Stat) AcquireCount() int64 { + return s.acquireCount +} + +// AcquireDuration returns the total duration of all successful acquires from +// the pool. +func (s *Stat) AcquireDuration() time.Duration { + return s.acquireDuration +} + +// EmptyAcquireCount returns the cumulative count of successful acquires from the pool +// that waited for a resource to be released or constructed because the pool was +// empty. +func (s *Stat) EmptyAcquireCount() int64 { + return s.emptyAcquireCount +} + +// EmptyAcquireWaitTime returns the cumulative time waited for successful acquires +// from the pool for a resource to be released or constructed because the pool was +// empty. +func (s *Stat) EmptyAcquireWaitTime() time.Duration { + return s.emptyAcquireWaitTime +} + +// CanceledAcquireCount returns the cumulative count of acquires from the pool +// that were canceled by a context. +func (s *Stat) CanceledAcquireCount() int64 { + return s.canceledAcquireCount +} + +// Stat returns the current pool statistics. +func (p *Pool[T]) Stat() *Stat { + p.mux.Lock() + defer p.mux.Unlock() + + s := &Stat{ + maxResources: p.maxSize, + acquireCount: p.acquireCount, + emptyAcquireCount: p.emptyAcquireCount, + emptyAcquireWaitTime: p.emptyAcquireWaitTime, + canceledAcquireCount: p.canceledAcquireCount.Load(), + acquireDuration: p.acquireDuration, + } + + for _, res := range p.allResources { + switch res.status { + case resourceStatusConstructing: + s.constructingResources += 1 + case resourceStatusIdle: + s.idleResources += 1 + case resourceStatusAcquired: + s.acquiredResources += 1 + } + } + + return s +} + +// tryAcquireIdleResource checks if there is any idle resource. If there is +// some, this method removes it from idle list and returns it. If the idle pool +// is empty, this method returns nil and doesn't modify the idleResources slice. +// +// WARNING: Caller of this method must hold the pool mutex! +func (p *Pool[T]) tryAcquireIdleResource() *Resource[T] { + res, ok := p.idleResources.Pop() + if !ok { + return nil + } + + res.status = resourceStatusAcquired + return res +} + +// createNewResource creates a new resource and inserts it into list of pool +// resources. +// +// WARNING: Caller of this method must hold the pool mutex! +func (p *Pool[T]) createNewResource() *Resource[T] { + res := &Resource[T]{ + pool: p, + creationTime: time.Now(), + lastUsedNano: nanotime(), + poolResetCount: p.resetCount, + status: resourceStatusConstructing, + } + + p.allResources.append(res) + p.destructWG.Add(1) + + return res +} + +// Acquire gets a resource from the pool. If no resources are available and the pool is not at maximum capacity it will +// create a new resource. If the pool is at maximum capacity it will block until a resource is available. ctx can be +// used to cancel the Acquire. +// +// If Acquire creates a new resource the resource constructor function will receive a context that delegates Value() to +// ctx. Canceling ctx will cause Acquire to return immediately but it will not cancel the resource creation. This avoids +// the problem of it being impossible to create resources when the time to create a resource is greater than any one +// caller of Acquire is willing to wait. +// +// Acquire guarantees that even if ctx is canceled concurrently with a successful resource creation, the resource will +// be properly returned to the idle pool (or destroyed) and the semaphore count will remain consistent. +func (p *Pool[T]) Acquire(ctx context.Context) (*Resource[T], error) { + select { + case <-ctx.Done(): + p.canceledAcquireCount.Add(1) + return nil, ctx.Err() + default: + } + + return p.acquire(ctx) +} + +// acquire is a continuation of Acquire function that doesn't check context +// validity. +// +// This function exists solely only for benchmarking purposes. +func (p *Pool[T]) acquire(ctx context.Context) (*Resource[T], error) { + startNano := nanotime() + + var waitedForLock bool + if !p.acquireSem.TryAcquire(1) { + waitedForLock = true + err := p.acquireSem.Acquire(ctx, 1) + if err != nil { + p.canceledAcquireCount.Add(1) + return nil, err + } + } + + // If the context has been cancelled between acquiring the semaphore and + // attempting to use it, release the semaphore and return immediately. + // This prevents goroutine leaks when the caller's context is cancelled + // concurrently with a TryAcquire success. + select { + case <-ctx.Done(): + p.acquireSem.Release(1) + p.canceledAcquireCount.Add(1) + return nil, ctx.Err() + default: + } + + p.mux.Lock() + if p.closed { + p.acquireSem.Release(1) + p.mux.Unlock() + return nil, ErrClosedPool + } + + // If a resource is available in the pool. + if res := p.tryAcquireIdleResource(); res != nil { + waitTime := time.Duration(nanotime() - startNano) + if waitedForLock { + p.emptyAcquireCount += 1 + p.emptyAcquireWaitTime += waitTime + } + p.acquireCount += 1 + p.acquireDuration += waitTime + p.mux.Unlock() + return res, nil + } + + if len(p.allResources) >= int(p.maxSize) { + // Unreachable code. + panic("bug: semaphore allowed more acquires than pool allows") + } + + // The resource is not idle, but there is enough space to create one. + res := p.createNewResource() + p.mux.Unlock() + + res, err := p.initResourceValue(ctx, res) + if err != nil { + return nil, err + } + + p.mux.Lock() + defer p.mux.Unlock() + + p.emptyAcquireCount += 1 + p.acquireCount += 1 + waitTime := time.Duration(nanotime() - startNano) + p.acquireDuration += waitTime + p.emptyAcquireWaitTime += waitTime + + return res, nil +} + +func (p *Pool[T]) initResourceValue(ctx context.Context, res *Resource[T]) (*Resource[T], error) { + // Create the resource in a goroutine to immediately return from Acquire + // if ctx is canceled without also canceling the constructor. + // + // See: + // - https://github.com/jackc/pgx/issues/1287 + // - https://github.com/jackc/pgx/issues/1259 + constructErrChan := make(chan error) + go func() { + constructorCtx := newValueCancelCtx(ctx, p.baseAcquireCtx) + value, err := p.constructor(constructorCtx) + if err != nil { + p.mux.Lock() + p.allResources.remove(res) + p.destructWG.Done() + + // The resource won't be acquired because its + // construction failed. We have to allow someone else to + // take that resouce. + p.acquireSem.Release(1) + p.mux.Unlock() + + select { + case constructErrChan <- err: + case <-ctx.Done(): + // The caller is cancelled, so no-one awaits the + // error. This branch avoid goroutine leak. + } + return + } + + // The resource is already in p.allResources where it might be read. So we need to acquire the lock to update its + // status. + p.mux.Lock() + res.value = value + res.status = resourceStatusAcquired + p.mux.Unlock() + + // This select works because the channel is unbuffered. + select { + case constructErrChan <- nil: + case <-ctx.Done(): + p.releaseAcquiredResource(res, res.lastUsedNano) + } + }() + + select { + case <-ctx.Done(): + p.canceledAcquireCount.Add(1) + return nil, ctx.Err() + case err := <-constructErrChan: + if err != nil { + return nil, err + } + return res, nil + } +} + +// TryAcquire gets a resource from the pool if one is immediately available. If not, it returns ErrNotAvailable. If no +// resources are available but the pool has room to grow, a resource will be created in the background. ctx is only +// used to cancel the background creation. +func (p *Pool[T]) TryAcquire(ctx context.Context) (*Resource[T], error) { + if !p.acquireSem.TryAcquire(1) { + return nil, ErrNotAvailable + } + + p.mux.Lock() + defer p.mux.Unlock() + + if p.closed { + p.acquireSem.Release(1) + return nil, ErrClosedPool + } + + // If a resource is available now + if res := p.tryAcquireIdleResource(); res != nil { + p.acquireCount += 1 + return res, nil + } + + if len(p.allResources) >= int(p.maxSize) { + // Unreachable code. + panic("bug: semaphore allowed more acquires than pool allows") + } + + res := p.createNewResource() + go func() { + value, err := p.constructor(ctx) + + p.mux.Lock() + defer p.mux.Unlock() + // We have to create the resource and only then release the + // semaphore - For the time being there is no resource that + // someone could acquire. + defer p.acquireSem.Release(1) + + if err != nil { + p.allResources.remove(res) + p.destructWG.Done() + return + } + + res.value = value + res.status = resourceStatusIdle + p.idleResources.Push(res) + }() + + return nil, ErrNotAvailable +} + +// acquireSemAll tries to acquire num free tokens from sem. This function is +// guaranteed to acquire at least the lowest number of tokens that has been +// available in the semaphore during runtime of this function. +// +// For the time being, semaphore doesn't allow to acquire all tokens atomically +// (see https://github.com/golang/sync/pull/19). We simulate this by trying all +// powers of 2 that are less or equal to num. +// +// For example, let's immagine we have 19 free tokens in the semaphore which in +// total has 24 tokens (i.e. the maxSize of the pool is 24 resources). Then if +// num is 24, the log2Uint(24) is 4 and we try to acquire 16, 8, 4, 2 and 1 +// tokens. Out of those, the acquire of 16, 2 and 1 tokens will succeed. +// +// Naturally, Acquires and Releases of the semaphore might take place +// concurrently. For this reason, it's not guaranteed that absolutely all free +// tokens in the semaphore will be acquired. But it's guaranteed that at least +// the minimal number of tokens that has been present over the whole process +// will be acquired. This is sufficient for the use-case we have in this +// package. +// +// TODO: Replace this with acquireSem.TryAcquireAll() if it gets to +// upstream. https://github.com/golang/sync/pull/19 +func acquireSemAll(sem *semaphore.Weighted, num int) int { + if num <= 0 { + panic("aquireSemAll: num <= 0") + } + if sem.TryAcquire(int64(num)) { + return num + } + var acquired int + for i := bits.Len64(uint64(num)) - 1; i >= 0; i-- { + val := 1 << i + if sem.TryAcquire(int64(val)) { + acquired += val + } + } + + return acquired +} + +// AcquireAllIdle acquires all currently idle resources. Its intended use is for +// health check and keep-alive functionality. It does not update pool +// statistics. +func (p *Pool[T]) AcquireAllIdle() []*Resource[T] { + p.mux.Lock() + defer p.mux.Unlock() + + if p.closed { + return nil + } + + numIdle := p.idleResources.Len() + if numIdle == 0 { + return nil + } + + // In acquireSemAll we use only TryAcquire and not Acquire. Because + // TryAcquire cannot block, the fact that we hold mutex locked and try + // to acquire semaphore cannot result in dead-lock. + // + // Because the mutex is locked, no parallel Release can run. This + // implies that the number of tokens can only decrease because some + // Acquire/TryAcquire call can consume the semaphore token. Consequently + // acquired is always less or equal to numIdle. Moreover if acquired < + // numIdle, then there are some parallel Acquire/TryAcquire calls that + // will take the remaining idle connections. + acquired := acquireSemAll(p.acquireSem, numIdle) + + idle := make([]*Resource[T], acquired) + for i := range idle { + res, _ := p.idleResources.Pop() + res.status = resourceStatusAcquired + idle[i] = res + } + + // We have to bump the generation to ensure that Acquire/TryAcquire + // calls running in parallel (those which caused acquired < numIdle) + // will consume old connections and not freshly released connections + // instead. + p.idleResources.NextGen() + + return idle +} + +// CreateResource constructs a new resource without acquiring it. It goes straight in the IdlePool. If the pool is full +// it returns an error. It can be useful to maintain warm resources under little load. +func (p *Pool[T]) CreateResource(ctx context.Context) error { + if !p.acquireSem.TryAcquire(1) { + return ErrNotAvailable + } + + p.mux.Lock() + if p.closed { + p.acquireSem.Release(1) + p.mux.Unlock() + return ErrClosedPool + } + + if len(p.allResources) >= int(p.maxSize) { + p.acquireSem.Release(1) + p.mux.Unlock() + return ErrNotAvailable + } + + res := p.createNewResource() + p.mux.Unlock() + + value, err := p.constructor(ctx) + p.mux.Lock() + defer p.mux.Unlock() + defer p.acquireSem.Release(1) + if err != nil { + p.allResources.remove(res) + p.destructWG.Done() + return err + } + + res.value = value + res.status = resourceStatusIdle + + // If closed while constructing resource then destroy it and return an error + if p.closed { + go p.destructResourceValue(res.value) + return ErrClosedPool + } + + p.idleResources.Push(res) + + return nil +} + +// Reset destroys all resources, but leaves the pool open. It is intended for use when an error is detected that would +// disrupt all resources (such as a network interruption or a server state change). +// +// It is safe to reset a pool while resources are checked out. Those resources will be destroyed when they are returned +// to the pool. +func (p *Pool[T]) Reset() { + p.mux.Lock() + defer p.mux.Unlock() + + p.resetCount++ + + for res, ok := p.idleResources.Pop(); ok; res, ok = p.idleResources.Pop() { + p.allResources.remove(res) + go p.destructResourceValue(res.value) + } +} + +// releaseAcquiredResource returns res to the the pool. +func (p *Pool[T]) releaseAcquiredResource(res *Resource[T], lastUsedNano int64) { + p.mux.Lock() + defer p.mux.Unlock() + defer p.acquireSem.Release(1) + + if p.closed || res.poolResetCount != p.resetCount { + p.allResources.remove(res) + go p.destructResourceValue(res.value) + } else { + res.lastUsedNano = lastUsedNano + res.status = resourceStatusIdle + p.idleResources.Push(res) + } +} + +// Remove removes res from the pool and closes it. If res is not part of the +// pool Remove will panic. +func (p *Pool[T]) destroyAcquiredResource(res *Resource[T]) { + p.destructResourceValue(res.value) + + p.mux.Lock() + defer p.mux.Unlock() + defer p.acquireSem.Release(1) + + p.allResources.remove(res) +} + +func (p *Pool[T]) hijackAcquiredResource(res *Resource[T]) { + p.mux.Lock() + defer p.mux.Unlock() + defer p.acquireSem.Release(1) + + p.allResources.remove(res) + res.status = resourceStatusHijacked + p.destructWG.Done() // not responsible for destructing hijacked resources +} + +func (p *Pool[T]) destructResourceValue(value T) { + p.destructor(value) + p.destructWG.Done() +} diff --git a/puddle/pool_test.go b/puddle/pool_test.go new file mode 100644 index 0000000..351853c --- /dev/null +++ b/puddle/pool_test.go @@ -0,0 +1,1576 @@ +package puddle_test + +import ( + "context" + "errors" + "fmt" + "log" + "math/rand" + "net" + "os" + "runtime" + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/jackc/puddle/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sync/semaphore" +) + +type Counter struct { + mutex sync.Mutex + n int +} + +// Next increments the counter and returns the value +func (c *Counter) Next() int { + c.mutex.Lock() + defer c.mutex.Unlock() + + c.n += 1 + return c.n +} + +// Value returns the counter +func (c *Counter) Value() int { + c.mutex.Lock() + defer c.mutex.Unlock() + + return c.n +} + +func createConstructor() (puddle.Constructor[int], *Counter) { + var c Counter + f := func(ctx context.Context) (int, error) { + return c.Next(), nil + } + return f, &c +} + +func stubDestructor(int) {} + +func TestNewPoolRequiresMaxSizeGreaterThan0(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: -1}) + assert.Nil(t, pool) + assert.Error(t, err) + + pool, err = puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 0}) + assert.Nil(t, pool) + assert.Error(t, err) +} + +func TestPoolAcquireCreatesResourceWhenNoneIdle(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + assert.WithinDuration(t, time.Now(), res.CreationTime(), time.Second) + res.Release() +} + +func TestPoolAcquireCallsConstructorWithAcquireContextValuesButNotDeadline(t *testing.T) { + constructor := func(ctx context.Context) (int, error) { + if ctx.Value("test") != "from Acquire" { + return 0, errors.New("did not get value from Acquire") + } + if _, ok := ctx.Deadline(); ok { + return 0, errors.New("should not have gotten deadline from Acquire") + } + + return 1, nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + ctx := context.WithValue(context.Background(), "test", "from Acquire") + ctx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + res, err := pool.Acquire(ctx) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + assert.WithinDuration(t, time.Now(), res.CreationTime(), time.Second) + res.Release() +} + +func TestPoolAcquireCalledConstructorIsNotCanceledByAcquireCancellation(t *testing.T) { + constructor := func(ctx context.Context) (int, error) { + time.Sleep(100 * time.Millisecond) + return 1, nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond) + defer cancel() + res, err := pool.Acquire(ctx) + assert.Nil(t, res) + assert.Equal(t, context.DeadlineExceeded, err) + + time.Sleep(200 * time.Millisecond) + + assert.EqualValues(t, 1, pool.Stat().TotalResources()) + assert.EqualValues(t, 1, pool.Stat().CanceledAcquireCount()) +} + +func TestPoolAcquireDoesNotCreatesResourceWhenItWouldExceedMaxSize(t *testing.T) { + constructor, createCounter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + + wg := &sync.WaitGroup{} + + for i := 0; i < 100; i++ { + wg.Add(1) + go func() { + for j := 0; j < 100; j++ { + res, err := pool.Acquire(context.Background()) + assert.NoError(t, err) + assert.Equal(t, 1, res.Value()) + res.Release() + } + wg.Done() + }() + } + + wg.Wait() + + assert.EqualValues(t, 1, createCounter.Value()) + assert.EqualValues(t, 1, pool.Stat().TotalResources()) +} + +func TestPoolAcquireWithCancellableContext(t *testing.T) { + constructor, createCounter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + + wg := &sync.WaitGroup{} + + for i := 0; i < 100; i++ { + wg.Add(1) + go func() { + for j := 0; j < 100; j++ { + ctx, cancel := context.WithCancel(context.Background()) + res, err := pool.Acquire(ctx) + assert.NoError(t, err) + assert.Equal(t, 1, res.Value()) + res.Release() + cancel() + } + wg.Done() + }() + } + + wg.Wait() + + assert.EqualValues(t, 1, createCounter.Value()) + assert.EqualValues(t, 1, pool.Stat().TotalResources()) +} + +func TestPoolAcquireReturnsErrorFromFailedResourceCreate(t *testing.T) { + errCreateFailed := errors.New("create failed") + constructor := func(ctx context.Context) (int, error) { + return 0, errCreateFailed + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + assert.Equal(t, errCreateFailed, err) + assert.Nil(t, res) +} + +func TestPoolAcquireCreatesResourceRespectingContext(t *testing.T) { + var cancel func() + constructor := func(ctx context.Context) (int, error) { + cancel() + // sleep to give a chance for the acquire to recognize it's cancelled + time.Sleep(10 * time.Millisecond) + return 1, nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + defer pool.Close() + + var ctx context.Context + ctx, cancel = context.WithCancel(context.Background()) + defer cancel() + _, err = pool.Acquire(ctx) + assert.ErrorIs(t, err, context.Canceled) + + // wait for the constructor to sleep and then for the resource to be added back + // to the idle pool + time.Sleep(100 * time.Millisecond) + + stat := pool.Stat() + assert.EqualValues(t, 1, stat.IdleResources()) + assert.EqualValues(t, 1, stat.TotalResources()) +} + +func TestPoolAcquireReusesResources(t *testing.T) { + constructor, createCounter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + + res.Release() + + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + + res.Release() + + assert.Equal(t, 1, createCounter.Value()) +} + +func TestPoolTryAcquire(t *testing.T) { + constructor, createCounter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + + // Pool is initially empty so TryAcquire fails but starts construction of resource in the background. + res, err := pool.TryAcquire(context.Background()) + require.EqualError(t, err, puddle.ErrNotAvailable.Error()) + assert.Nil(t, res) + + // Wait for background creation to complete. + time.Sleep(100 * time.Millisecond) + + res, err = pool.TryAcquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + defer res.Release() + + res, err = pool.TryAcquire(context.Background()) + require.EqualError(t, err, puddle.ErrNotAvailable.Error()) + assert.Nil(t, res) + + assert.Equal(t, 1, createCounter.Value()) +} + +func TestPoolTryAcquireReturnsErrorWhenPoolIsClosed(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + pool.Close() + + res, err := pool.TryAcquire(context.Background()) + assert.Equal(t, puddle.ErrClosedPool, err) + assert.Nil(t, res) +} + +func TestPoolTryAcquireWithFailedResourceCreate(t *testing.T) { + errCreateFailed := errors.New("create failed") + constructor := func(ctx context.Context) (int, error) { + return 0, errCreateFailed + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.TryAcquire(context.Background()) + require.EqualError(t, err, puddle.ErrNotAvailable.Error()) + assert.Nil(t, res) +} + +func TestPoolAcquireNilContextDoesNotLeavePoolLocked(t *testing.T) { + constructor, createCounter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + assert.Panics(t, func() { pool.Acquire(nil) }) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + res.Release() + + assert.Equal(t, 1, createCounter.Value()) +} + +func TestPoolAcquireContextAlreadyCanceled(t *testing.T) { + constructor := func(ctx context.Context) (int, error) { + panic("should never be called") + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + res, err := pool.Acquire(ctx) + assert.Equal(t, context.Canceled, err) + assert.Nil(t, res) +} + +func TestPoolAcquireContextCanceledDuringCreate(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + time.AfterFunc(100*time.Millisecond, cancel) + timeoutChan := time.After(1 * time.Second) + + var constructorCalls Counter + constructor := func(ctx context.Context) (int, error) { + select { + case <-ctx.Done(): + return 0, ctx.Err() + case <-timeoutChan: + } + return constructorCalls.Next(), nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(ctx) + assert.Equal(t, context.Canceled, err) + assert.Nil(t, res) +} + +func TestPoolAcquireAllIdle(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + resources := make([]*puddle.Resource[int], 4) + + resources[0], err = pool.Acquire(context.Background()) + require.NoError(t, err) + resources[1], err = pool.Acquire(context.Background()) + require.NoError(t, err) + resources[2], err = pool.Acquire(context.Background()) + require.NoError(t, err) + resources[3], err = pool.Acquire(context.Background()) + require.NoError(t, err) + + assert.Len(t, pool.AcquireAllIdle(), 0) + + resources[0].Release() + resources[3].Release() + + assert.ElementsMatch(t, []*puddle.Resource[int]{resources[0], resources[3]}, pool.AcquireAllIdle()) + + resources[0].Release() + resources[3].Release() + resources[1].Release() + resources[2].Release() + + assert.ElementsMatch(t, resources, pool.AcquireAllIdle()) + + resources[0].Release() + resources[1].Release() + resources[2].Release() + resources[3].Release() +} + +func TestPoolAcquireAllIdleWhenClosedIsNil(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + pool.Close() + assert.Nil(t, pool.AcquireAllIdle()) +} + +func TestPoolCreateResource(t *testing.T) { + constructor, counter := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + err = pool.CreateResource(context.Background()) + require.NoError(t, err) + + stats := pool.Stat() + assert.EqualValues(t, 1, stats.IdleResources()) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, counter.Value(), res.Value()) + assert.True(t, res.LastUsedNanotime() > 0, "should set LastUsedNanotime so that idle calculations can still work") + assert.Equal(t, 1, res.Value()) + assert.WithinDuration(t, time.Now(), res.CreationTime(), time.Second) + res.Release() + + assert.EqualValues(t, 0, pool.Stat().EmptyAcquireCount(), "should have been a warm resource") +} + +func TestPoolCreateResourceReturnsErrorFromFailedResourceCreate(t *testing.T) { + errCreateFailed := errors.New("create failed") + constructor := func(ctx context.Context) (int, error) { + return 0, errCreateFailed + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + err = pool.CreateResource(context.Background()) + assert.Equal(t, errCreateFailed, err) +} + +func TestPoolCreateResourceReturnsErrorWhenAlreadyClosed(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + pool.Close() + err = pool.CreateResource(context.Background()) + assert.Equal(t, puddle.ErrClosedPool, err) +} + +func TestPoolCreateResourceReturnsErrorWhenClosedWhileCreatingResource(t *testing.T) { + // There is no way to guarantee the correct order of the pool being closed while the resource is being constructed. + // But these sleeps should make it extremely likely. (Ah, the lengths we go for 100% test coverage...) + constructor := func(ctx context.Context) (int, error) { + time.Sleep(500 * time.Millisecond) + return 123, nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + acquireErrChan := make(chan error) + go func() { + err := pool.CreateResource(context.Background()) + acquireErrChan <- err + }() + + time.Sleep(250 * time.Millisecond) + pool.Close() + + err = <-acquireErrChan + assert.Equal(t, puddle.ErrClosedPool, err) +} + +func TestPoolCreateResourceReturnsErrorWhenPoolFull(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 2}) + require.NoError(t, err) + defer pool.Close() + + err = pool.CreateResource(context.Background()) + require.NoError(t, err) + + stats := pool.Stat() + assert.EqualValues(t, 1, stats.IdleResources()) + + err = pool.CreateResource(context.Background()) + require.NoError(t, err) + + stats = pool.Stat() + assert.EqualValues(t, 2, stats.IdleResources()) + + err = pool.CreateResource(context.Background()) + require.Error(t, err) +} + +func TestPoolCloseClosesAllIdleResources(t *testing.T) { + constructor, _ := createConstructor() + + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + p, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + resources := make([]*puddle.Resource[int], 4) + for i := range resources { + var err error + resources[i], err = p.Acquire(context.Background()) + require.Nil(t, err) + } + + for _, res := range resources { + res.Release() + } + + p.Close() + + assert.Equal(t, len(resources), destructorCalls.Value()) +} + +func TestPoolCloseBlocksUntilAllResourcesReleasedAndClosed(t *testing.T) { + constructor, _ := createConstructor() + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + p, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + resources := make([]*puddle.Resource[int], 4) + for i := range resources { + var err error + resources[i], err = p.Acquire(context.Background()) + require.Nil(t, err) + } + + for _, res := range resources { + go func(res *puddle.Resource[int]) { + time.Sleep(100 * time.Millisecond) + res.Release() + }(res) + } + + p.Close() + assert.Equal(t, len(resources), destructorCalls.Value()) +} + +func TestPoolCloseIsSafeToCallMultipleTimes(t *testing.T) { + constructor, _ := createConstructor() + + p, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + p.Close() + p.Close() +} + +func TestPoolResetDestroysAllIdleResources(t *testing.T) { + constructor, _ := createConstructor() + + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + p, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + resources := make([]*puddle.Resource[int], 4) + for i := range resources { + var err error + resources[i], err = p.Acquire(context.Background()) + require.Nil(t, err) + } + + for _, res := range resources { + res.Release() + } + + require.EqualValues(t, 4, p.Stat().TotalResources()) + p.Reset() + require.EqualValues(t, 0, p.Stat().TotalResources()) + + // Destructors are called in the background. No way to know when they are all finished. + for i := 0; i < 100; i++ { + if destructorCalls.Value() == len(resources) { + break + } + time.Sleep(100 * time.Millisecond) + } + require.Equal(t, len(resources), destructorCalls.Value()) + + p.Close() +} + +func TestPoolResetDestroysCheckedOutResourcesOnReturn(t *testing.T) { + constructor, _ := createConstructor() + + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + p, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + resources := make([]*puddle.Resource[int], 4) + for i := range resources { + var err error + resources[i], err = p.Acquire(context.Background()) + require.Nil(t, err) + } + + require.EqualValues(t, 4, p.Stat().TotalResources()) + p.Reset() + require.EqualValues(t, 4, p.Stat().TotalResources()) + + for _, res := range resources { + res.Release() + } + + require.EqualValues(t, 0, p.Stat().TotalResources()) + + // Destructors are called in the background. No way to know when they are all finished. + for i := 0; i < 100; i++ { + if destructorCalls.Value() == len(resources) { + break + } + time.Sleep(100 * time.Millisecond) + } + require.Equal(t, len(resources), destructorCalls.Value()) + + p.Close() +} + +func TestPoolStatResources(t *testing.T) { + startWaitChan := make(chan struct{}) + waitingChan := make(chan struct{}) + endWaitChan := make(chan struct{}) + + var constructorCalls Counter + constructor := func(ctx context.Context) (int, error) { + select { + case <-startWaitChan: + close(waitingChan) + <-endWaitChan + default: + } + + return constructorCalls.Next(), nil + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + defer pool.Close() + + resAcquired, err := pool.Acquire(context.Background()) + require.Nil(t, err) + + close(startWaitChan) + go func() { + res, err := pool.Acquire(context.Background()) + require.Nil(t, err) + res.Release() + }() + <-waitingChan + stat := pool.Stat() + + assert.EqualValues(t, 2, stat.TotalResources()) + assert.EqualValues(t, 1, stat.ConstructingResources()) + assert.EqualValues(t, 1, stat.AcquiredResources()) + assert.EqualValues(t, 0, stat.IdleResources()) + assert.EqualValues(t, 10, stat.MaxResources()) + + resAcquired.Release() + + stat = pool.Stat() + assert.EqualValues(t, 2, stat.TotalResources()) + assert.EqualValues(t, 1, stat.ConstructingResources()) + assert.EqualValues(t, 0, stat.AcquiredResources()) + assert.EqualValues(t, 1, stat.IdleResources()) + assert.EqualValues(t, 10, stat.MaxResources()) + + close(endWaitChan) +} + +func TestPoolStatSuccessfulAcquireCounters(t *testing.T) { + constructor, _ := createConstructor() + sleepConstructor := func(ctx context.Context) (int, error) { + // sleep to make sure we don't fail the AcquireDuration test + time.Sleep(time.Nanosecond) + return constructor(ctx) + } + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: sleepConstructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + defer pool.Close() + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + res.Release() + + stat := pool.Stat() + assert.Equal(t, int64(1), stat.AcquireCount()) + assert.Equal(t, int64(1), stat.EmptyAcquireCount()) + assert.Positive(t, stat.AcquireDuration(), "expected stat.AcquireDuration() > 0 but %v", stat.AcquireDuration()) + assert.Equal(t, stat.EmptyAcquireWaitTime(), stat.AcquireDuration()) + lastAcquireDuration := stat.AcquireDuration() + + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + res.Release() + + stat = pool.Stat() + assert.Equal(t, int64(2), stat.AcquireCount()) + assert.Equal(t, int64(1), stat.EmptyAcquireCount()) + assert.Greater(t, stat.AcquireDuration(), lastAcquireDuration) + assert.Less(t, stat.EmptyAcquireWaitTime(), stat.AcquireDuration()) + lastAcquireDuration = stat.AcquireDuration() + + wg := &sync.WaitGroup{} + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + time.Sleep(50 * time.Millisecond) + res.Release() + wg.Done() + }() + } + + wg.Wait() + + stat = pool.Stat() + assert.Equal(t, int64(4), stat.AcquireCount()) + assert.Equal(t, int64(2), stat.EmptyAcquireCount()) + assert.Greater(t, stat.AcquireDuration(), lastAcquireDuration) + assert.Less(t, stat.EmptyAcquireWaitTime(), stat.AcquireDuration()) + lastAcquireDuration = stat.AcquireDuration() +} + +func TestPoolStatCanceledAcquireBeforeStart(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + defer pool.Close() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err = pool.Acquire(ctx) + require.Equal(t, context.Canceled, err) + + stat := pool.Stat() + assert.Equal(t, int64(0), stat.AcquireCount()) + assert.Equal(t, int64(1), stat.CanceledAcquireCount()) +} + +func TestPoolStatCanceledAcquireDuringCreate(t *testing.T) { + constructor := func(ctx context.Context) (int, error) { + <-ctx.Done() + return 0, ctx.Err() + } + + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + defer pool.Close() + + ctx, cancel := context.WithCancel(context.Background()) + time.AfterFunc(50*time.Millisecond, cancel) + _, err = pool.Acquire(ctx) + require.Equal(t, context.Canceled, err) + + // sleep to give the constructor goroutine time to mark cancelled + time.Sleep(10 * time.Millisecond) + + stat := pool.Stat() + assert.Equal(t, int64(0), stat.AcquireCount()) + assert.Equal(t, int64(1), stat.CanceledAcquireCount()) +} + +func TestPoolStatCanceledAcquireDuringWait(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + defer pool.Close() + + res, err := pool.Acquire(context.Background()) + require.Nil(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + time.AfterFunc(50*time.Millisecond, cancel) + _, err = pool.Acquire(ctx) + require.Equal(t, context.Canceled, err) + + res.Release() + + stat := pool.Stat() + assert.Equal(t, int64(1), stat.AcquireCount()) + assert.Equal(t, int64(1), stat.CanceledAcquireCount()) +} + +func TestResourceHijackRemovesResourceFromPoolButDoesNotDestroy(t *testing.T) { + constructor, _ := createConstructor() + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + + res.Hijack() + + assert.EqualValues(t, 0, pool.Stat().TotalResources()) + assert.EqualValues(t, 0, destructorCalls.Value()) + + // Can still call Value, CreationTime and IdleDuration + res.Value() + res.CreationTime() + res.IdleDuration() +} + +func TestResourceDestroyRemovesResourceFromPool(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, res.Value()) + + assert.EqualValues(t, 1, pool.Stat().TotalResources()) + res.Destroy() + for i := 0; i < 1000; i++ { + if pool.Stat().TotalResources() == 0 { + break + } + time.Sleep(time.Millisecond) + } + + assert.EqualValues(t, 0, pool.Stat().TotalResources()) +} + +func TestResourceLastUsageTimeTracking(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 1}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + t1 := res.LastUsedNanotime() + res.Release() + + // Greater than zero after initial usage + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + t2 := res.LastUsedNanotime() + d2 := res.IdleDuration() + assert.True(t, t2 > t1) + res.ReleaseUnused() + + // ReleaseUnused does not update usage tracking + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + t3 := res.LastUsedNanotime() + d3 := res.IdleDuration() + assert.EqualValues(t, t2, t3) + assert.True(t, d3 > d2) + res.Release() + + // Release does update usage tracking + res, err = pool.Acquire(context.Background()) + require.NoError(t, err) + t4 := res.LastUsedNanotime() + assert.True(t, t4 > t3) + res.Release() +} + +func TestResourcePanicsOnUsageWhenNotAcquired(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + + res, err := pool.Acquire(context.Background()) + require.NoError(t, err) + res.Release() + + assert.PanicsWithValue(t, "tried to release resource that is not acquired", res.Release) + assert.PanicsWithValue(t, "tried to release resource that is not acquired", res.ReleaseUnused) + assert.PanicsWithValue(t, "tried to destroy resource that is not acquired", res.Destroy) + assert.PanicsWithValue(t, "tried to hijack resource that is not acquired", res.Hijack) + assert.PanicsWithValue(t, "tried to access resource that is not acquired or hijacked", func() { res.Value() }) + assert.PanicsWithValue(t, "tried to access resource that is not acquired or hijacked", func() { res.CreationTime() }) + assert.PanicsWithValue(t, "tried to access resource that is not acquired or hijacked", func() { res.LastUsedNanotime() }) + assert.PanicsWithValue(t, "tried to access resource that is not acquired or hijacked", func() { res.IdleDuration() }) +} + +func TestPoolAcquireReturnsErrorWhenPoolIsClosed(t *testing.T) { + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 10}) + require.NoError(t, err) + pool.Close() + + res, err := pool.Acquire(context.Background()) + assert.Equal(t, puddle.ErrClosedPool, err) + assert.Nil(t, res) +} + +func TestSignalIsSentWhenResourceFailedToCreate(t *testing.T) { + var c Counter + constructor := func(context.Context) (a any, err error) { + if c.Next() == 2 { + return nil, errors.New("outage") + } + return 1, nil + } + destructor := func(value any) {} + + pool, err := puddle.NewPool(&puddle.Config[any]{Constructor: constructor, Destructor: destructor, MaxSize: 10}) + require.NoError(t, err) + + res1, err := pool.Acquire(context.Background()) + require.NoError(t, err) + + var wg sync.WaitGroup + for i := 0; i < 2; i++ { + wg.Add(1) + go func(name string) { + defer wg.Done() + _, _ = pool.Acquire(context.Background()) + }(strconv.Itoa(i)) + } + + // ensure that both goroutines above are waiting for condition variable signal + time.Sleep(500 * time.Millisecond) + res1.Destroy() + wg.Wait() +} + +func stressTestDur(t testing.TB) time.Duration { + s := os.Getenv("STRESS_TEST_DURATION") + if s == "" { + s = "1s" + } + + dur, err := time.ParseDuration(s) + require.Nil(t, err) + return dur +} + +func TestStress(t *testing.T) { + constructor, _ := createConstructor() + var destructorCalls Counter + destructor := func(int) { + destructorCalls.Next() + } + + poolSize := runtime.NumCPU() + if poolSize < 4 { + poolSize = 4 + } + + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: destructor, MaxSize: int32(poolSize)}) + require.NoError(t, err) + + finishChan := make(chan struct{}) + wg := &sync.WaitGroup{} + + releaseOrDestroyOrHijack := func(res *puddle.Resource[int]) { + n := rand.Intn(100) + if n < 5 { + res.Hijack() + destructor(res.Value()) + } else if n < 10 { + res.Destroy() + } else { + res.Release() + } + } + + actions := []func(){ + // Acquire + func() { + res, err := pool.Acquire(context.Background()) + if err != nil { + if err != puddle.ErrClosedPool { + assert.Failf(t, "stress acquire", "pool.Acquire returned unexpected err: %v", err) + } + return + } + + time.Sleep(time.Duration(rand.Int63n(100)) * time.Millisecond) + releaseOrDestroyOrHijack(res) + }, + // Acquire possibly canceled by context + func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Duration(rand.Int63n(2000))*time.Nanosecond) + defer cancel() + res, err := pool.Acquire(ctx) + if err != nil { + if err != puddle.ErrClosedPool && err != context.Canceled && err != context.DeadlineExceeded { + assert.Failf(t, "stress acquire possibly canceled by context", "pool.Acquire returned unexpected err: %v", err) + } + return + } + + time.Sleep(time.Duration(rand.Int63n(2000)) * time.Nanosecond) + releaseOrDestroyOrHijack(res) + }, + // TryAcquire + func() { + res, err := pool.TryAcquire(context.Background()) + if err != nil { + if err != puddle.ErrClosedPool && err != puddle.ErrNotAvailable { + assert.Failf(t, "stress TryAcquire", "pool.TryAcquire returned unexpected err: %v", err) + } + return + } + + time.Sleep(time.Duration(rand.Int63n(100)) * time.Millisecond) + releaseOrDestroyOrHijack(res) + }, + // AcquireAllIdle (though under heavy load this will almost certainly always get an empty slice) + func() { + resources := pool.AcquireAllIdle() + for _, res := range resources { + res.Release() + } + }, + // Stat + func() { + stat := pool.Stat() + assert.NotNil(t, stat) + }, + // CreateResource + func() { + err := pool.CreateResource(context.Background()) + if err != nil && !errors.Is(err, puddle.ErrClosedPool) && !errors.Is(err, puddle.ErrNotAvailable) { + t.Error(err) + } + }, + } + + workerCount := int(poolSize) * 2 + + for i := 0; i < workerCount; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-finishChan: + return + default: + } + + actions[rand.Intn(len(actions))]() + } + }() + } + + time.AfterFunc(stressTestDur(t), func() { close(finishChan) }) + wg.Wait() + pool.Close() +} + +func TestStress_AcquireAllIdle_TryAcquire(t *testing.T) { + r := require.New(t) + + pool := testPool[int32](t) + + var wg sync.WaitGroup + done := make(chan struct{}) + + wg.Add(1) + go func() { + defer wg.Done() + + for { + select { + case <-done: + return + default: + } + + idleRes := pool.AcquireAllIdle() + r.Less(len(idleRes), 2) + for _, res := range idleRes { + res.Release() + } + } + }() + + wg.Add(1) + go func() { + defer wg.Done() + + for { + select { + case <-done: + return + default: + } + + res, err := pool.TryAcquire(context.Background()) + if err != nil { + r.Equal(puddle.ErrNotAvailable, err) + } else { + r.NotNil(res) + res.Release() + } + } + }() + + time.AfterFunc(stressTestDur(t), func() { close(done) }) + wg.Wait() +} + +func TestStress_AcquireAllIdle_Acquire(t *testing.T) { + r := require.New(t) + + pool := testPool[int32](t) + + var wg sync.WaitGroup + done := make(chan struct{}) + + wg.Add(1) + go func() { + defer wg.Done() + + for { + select { + case <-done: + return + default: + } + + idleRes := pool.AcquireAllIdle() + r.Less(len(idleRes), 2) + for _, res := range idleRes { + r.NotNil(res) + res.Release() + } + } + }() + + wg.Add(1) + go func() { + defer wg.Done() + + for { + select { + case <-done: + return + default: + } + + res, err := pool.Acquire(context.Background()) + if err != nil { + r.Equal(puddle.ErrNotAvailable, err) + } else { + r.NotNil(res) + res.Release() + } + } + }() + + time.AfterFunc(stressTestDur(t), func() { close(done) }) + wg.Wait() +} + +func startAcceptOnceDummyServer(laddr string) { + ln, err := net.Listen("tcp", laddr) + if err != nil { + log.Fatalln("Listen:", err) + } + + // Listen one time + go func() { + conn, err := ln.Accept() + if err != nil { + log.Fatalln("Accept:", err) + } + + for { + buf := make([]byte, 1) + _, err := conn.Read(buf) + if err != nil { + return + } + } + }() + +} + +func ExamplePool() { + // Dummy server + laddr := "127.0.0.1:8080" + startAcceptOnceDummyServer(laddr) + + // Pool creation + constructor := func(context.Context) (any, error) { + return net.Dial("tcp", laddr) + } + destructor := func(value any) { + value.(net.Conn).Close() + } + maxPoolSize := int32(10) + + pool, err := puddle.NewPool(&puddle.Config[any]{Constructor: constructor, Destructor: destructor, MaxSize: int32(maxPoolSize)}) + if err != nil { + log.Fatalln("NewPool", err) + } + + // Use pool multiple times + for i := 0; i < 10; i++ { + // Acquire resource + res, err := pool.Acquire(context.Background()) + if err != nil { + log.Fatalln("Acquire", err) + } + + // Type-assert value and use + _, err = res.Value().(net.Conn).Write([]byte{1}) + if err != nil { + log.Fatalln("Write", err) + } + + // Release when done. + res.Release() + } + + stats := pool.Stat() + pool.Close() + + fmt.Println("Connections:", stats.TotalResources()) + fmt.Println("Acquires:", stats.AcquireCount()) + // Output: + // Connections: 1 + // Acquires: 10 +} + +func BenchmarkPoolAcquireAndRelease(b *testing.B) { + benchmarks := []struct { + poolSize int32 + clientCount int + cancellable bool + }{ + {8, 1, false}, + {8, 2, false}, + {8, 8, false}, + {8, 32, false}, + {8, 128, false}, + {8, 512, false}, + {8, 2048, false}, + {8, 8192, false}, + + {64, 2, false}, + {64, 8, false}, + {64, 32, false}, + {64, 128, false}, + {64, 512, false}, + {64, 2048, false}, + {64, 8192, false}, + + {512, 2, false}, + {512, 8, false}, + {512, 32, false}, + {512, 128, false}, + {512, 512, false}, + {512, 2048, false}, + {512, 8192, false}, + + {8, 2, true}, + {8, 8, true}, + {8, 32, true}, + {8, 128, true}, + {8, 512, true}, + {8, 2048, true}, + {8, 8192, true}, + + {64, 2, true}, + {64, 8, true}, + {64, 32, true}, + {64, 128, true}, + {64, 512, true}, + {64, 2048, true}, + {64, 8192, true}, + + {512, 2, true}, + {512, 8, true}, + {512, 32, true}, + {512, 128, true}, + {512, 512, true}, + {512, 2048, true}, + {512, 8192, true}, + } + + for _, bm := range benchmarks { + name := fmt.Sprintf("PoolSize=%d/ClientCount=%d/Cancellable=%v", bm.poolSize, bm.clientCount, bm.cancellable) + + b.Run(name, func(b *testing.B) { + ctx := context.Background() + cancel := func() {} + if bm.cancellable { + ctx, cancel = context.WithCancel(ctx) + } + + wg := &sync.WaitGroup{} + + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: bm.poolSize}) + if err != nil { + b.Fatal(err) + } + + for i := 0; i < bm.clientCount; i++ { + wg.Add(1) + go func() { + defer wg.Done() + + for j := 0; j < b.N; j++ { + res, err := pool.Acquire(ctx) + if err != nil { + b.Fatal(err) + } + res.Release() + } + }() + } + + wg.Wait() + cancel() + }) + } +} + +func TestAcquireAllSem(t *testing.T) { + r := require.New(t) + + sem := semaphore.NewWeighted(5) + r.Equal(4, puddle.AcquireSemAll(sem, 4)) + sem.Release(4) + + r.Equal(5, puddle.AcquireSemAll(sem, 5)) + sem.Release(5) + + r.Equal(5, puddle.AcquireSemAll(sem, 6)) + sem.Release(5) +} + +func testPool[T any](t testing.TB) *puddle.Pool[T] { + cfg := puddle.Config[T]{ + MaxSize: 1, + Constructor: func(ctx context.Context) (T, error) { + var zero T + return zero, nil + }, + Destructor: func(T) {}, + } + + pool, err := puddle.NewPool(&cfg) + require.NoError(t, err) + t.Cleanup(pool.Close) + + return pool +} + +func releaser[T any](t testing.TB) chan<- *puddle.Resource[T] { + startChan := make(chan struct{}) + workChan := make(chan *puddle.Resource[T], 1) + + go func() { + close(startChan) + + for r := range workChan { + r.Release() + } + }() + t.Cleanup(func() { close(workChan) }) + + // Wait for goroutine start. + <-startChan + return workChan +} + +func TestReleaseAfterAcquire(t *testing.T) { + const cnt = 100000 + + r := require.New(t) + ctx := context.Background() + pool := testPool[int32](t) + releaseChan := releaser[int32](t) + + res, err := pool.Acquire(ctx) + r.NoError(err) + // We need to release the last connection. Otherwise the pool.Close() + // method will block and this function will never return. + defer func() { res.Release() }() + + for i := 0; i < cnt; i++ { + releaseChan <- res + res, err = pool.Acquire(ctx) + r.NoError(err) + } +} + +func BenchmarkAcquire_ReleaseAfterAcquire(b *testing.B) { + r := require.New(b) + ctx := context.Background() + pool := testPool[int32](b) + releaseChan := releaser[int32](b) + + res, err := pool.Acquire(ctx) + r.NoError(err) + // We need to release the last connection. Otherwise the pool.Close() + // method will block and this function will never return. + defer func() { res.Release() }() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + releaseChan <- res + res, err = pool.Acquire(ctx) + r.NoError(err) + } +} + +func withCPULoad() { + // Multiply by 2 to similate overload of the system. + numGoroutines := runtime.NumCPU() * 2 + + var wg sync.WaitGroup + for i := 0; i < numGoroutines; i++ { + wg.Add(1) + go func() { + wg.Done() + + // Similate computationally intensive task. + for j := 0; true; j++ { + } + }() + } + + wg.Wait() +} + +func BenchmarkAcquire_ReleaseAfterAcquireWithCPULoad(b *testing.B) { + r := require.New(b) + ctx := context.Background() + pool := testPool[int32](b) + releaseChan := releaser[int32](b) + + withCPULoad() + + res, err := pool.Acquire(ctx) + r.NoError(err) + // We need to release the last connection. Otherwise the pool.Close() + // method will block and this function will never return. + defer func() { res.Release() }() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + releaseChan <- res + res, err = pool.Acquire(ctx) + r.NoError(err) + } +} + +func BenchmarkAcquire_MultipleCancelled(b *testing.B) { + const cancelCnt = 64 + + r := require.New(b) + ctx := context.Background() + pool := testPool[int32](b) + releaseChan := releaser[int32](b) + + cancelCtx, cancel := context.WithCancel(ctx) + cancel() + + res, err := pool.Acquire(ctx) + r.NoError(err) + // We need to release the last connection. Otherwise the pool.Close() + // method will block and this function will never return. + defer func() { res.Release() }() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + for j := 0; j < cancelCnt; j++ { + _, err = pool.AcquireRaw(cancelCtx) + r.Equal(context.Canceled, err) + } + + releaseChan <- res + res, err = pool.Acquire(ctx) + r.NoError(err) + } +} + +func BenchmarkAcquire_MultipleCancelledWithCPULoad(b *testing.B) { + const cancelCnt = 3 + + r := require.New(b) + ctx := context.Background() + pool := testPool[int32](b) + releaseChan := releaser[int32](b) + + cancelCtx, cancel := context.WithCancel(ctx) + cancel() + + withCPULoad() + + res, err := pool.Acquire(ctx) + r.NoError(err) + // We need to release the last connection. Otherwise the pool.Close() + // method will block and this function will never return. + defer func() { res.Release() }() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + for j := 0; j < cancelCnt; j++ { + _, err = pool.AcquireRaw(cancelCtx) + r.Equal(context.Canceled, err) + } + + releaseChan <- res + res, err = pool.Acquire(ctx) + r.NoError(err) + } +} + +func TestPoolAcquireStressCancelContention(t *testing.T) { + t.Parallel() + + constructor, _ := createConstructor() + pool, err := puddle.NewPool(&puddle.Config[int]{Constructor: constructor, Destructor: stubDestructor, MaxSize: 3}) + require.NoError(t, err) + defer pool.Close() + + res1, err := pool.Acquire(context.Background()) + require.NoError(t, err) + res2, err := pool.Acquire(context.Background()) + require.NoError(t, err) + res3, err := pool.Acquire(context.Background()) + require.NoError(t, err) + + var wg sync.WaitGroup + var cancelCount atomic.Int64 + + for i := 0; i < 30; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 10; j++ { + ms := time.Duration(rand.Intn(5)+1) * time.Millisecond + ctx, cancel := context.WithTimeout(context.Background(), ms) + res, err := pool.Acquire(ctx) + cancel() + if err == nil { + res.Release() + } else { + cancelCount.Add(1) + } + } + }() + } + + time.Sleep(20 * time.Millisecond) + res1.Release() + time.Sleep(100 * time.Millisecond) + res2.Release() + res3.Release() + + wg.Wait() + + cleanCtx, cleanCancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cleanCancel() + res, err := pool.Acquire(cleanCtx) + require.NoError(t, err, "pool should be functional after cancellation chaos") + res.Release() + + stat := pool.Stat() + t.Logf("Total: %d, Acquire: %d, Canceled: %d, Idle: %d", + stat.TotalResources(), stat.AcquireCount(), stat.CanceledAcquireCount(), stat.IdleResources()) + require.Equal(t, cancelCount.Load(), stat.CanceledAcquireCount(), + "canceled acquire count mismatch") + require.GreaterOrEqual(t, stat.IdleResources(), int32(0), + "idle resources should never be negative") + require.GreaterOrEqual(t, stat.TotalResources(), int32(1), + "pool should have resources available") + require.LessOrEqual(t, stat.TotalResources(), int32(3), + "pool should not exceed maxSize") +} diff --git a/puddle/resource_list.go b/puddle/resource_list.go new file mode 100644 index 0000000..b243095 --- /dev/null +++ b/puddle/resource_list.go @@ -0,0 +1,28 @@ +package puddle + +type resList[T any] []*Resource[T] + +func (l *resList[T]) append(val *Resource[T]) { *l = append(*l, val) } + +func (l *resList[T]) popBack() *Resource[T] { + idx := len(*l) - 1 + val := (*l)[idx] + (*l)[idx] = nil // Avoid memory leak + *l = (*l)[:idx] + + return val +} + +func (l *resList[T]) remove(val *Resource[T]) { + for i, elem := range *l { + if elem == val { + lastIdx := len(*l) - 1 + (*l)[i] = (*l)[lastIdx] + (*l)[lastIdx] = nil // Avoid memory leak + (*l) = (*l)[:lastIdx] + return + } + } + + panic("BUG: removeResource could not find res in slice") +} diff --git a/puddle/resource_list_test.go b/puddle/resource_list_test.go new file mode 100644 index 0000000..7104189 --- /dev/null +++ b/puddle/resource_list_test.go @@ -0,0 +1,62 @@ +package puddle + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestResList_Append(t *testing.T) { + r := require.New(t) + + arr := []*Resource[any]{ + new(Resource[any]), + new(Resource[any]), + new(Resource[any]), + } + + list := resList[any](arr) + + list.append(new(Resource[any])) + r.Len(list, 4) + list.append(new(Resource[any])) + r.Len(list, 5) + list.append(new(Resource[any])) + r.Len(list, 6) +} + +func TestResList_PopBack(t *testing.T) { + r := require.New(t) + + arr := []*Resource[any]{ + new(Resource[any]), + new(Resource[any]), + new(Resource[any]), + } + + list := resList[any](arr) + + list.popBack() + r.Len(list, 2) + list.popBack() + r.Len(list, 1) + list.popBack() + r.Len(list, 0) + + r.Panics(func() { list.popBack() }) +} + +func TestResList_PanicsWithBugReportIfResourceDoesNotExist(t *testing.T) { + arr := []*Resource[any]{ + new(Resource[any]), + new(Resource[any]), + new(Resource[any]), + } + + list := resList[any](arr) + + assert.PanicsWithValue(t, "BUG: removeResource could not find res in slice", func() { + list.remove(new(Resource[any])) + }) +}