From 7e69cd9f62245edb27b732b1a72215b545a207a3 Mon Sep 17 00:00:00 2001 From: prasadlohakpure Date: Tue, 29 Sep 2026 09:57:56 +0530 Subject: [PATCH] fix: reuse one database pool and keep queries on the session transaction Opening a new pool per session exhausted Postgres. One *sql.DB is shared per Database, and Query and QueryRow use the open transaction so FOR UPDATE locks last until commit. --- configs/local.yaml | 2 + internal/pkg/database/database.go | 63 ++++++++++++++++++++++---- internal/pkg/database/database_test.go | 52 +++++++++++++++++++++ 3 files changed, 108 insertions(+), 9 deletions(-) create mode 100644 internal/pkg/database/database_test.go diff --git a/configs/local.yaml b/configs/local.yaml index 11c535df..f0a5e696 100644 --- a/configs/local.yaml +++ b/configs/local.yaml @@ -2,6 +2,8 @@ # database settings database: connection_string: "postgres://heimdall:heimdall@postgres:5432/heimdall?sslmode=disable" + max_open_conns: 25 + max_idle_conns: 10 # workers pool to execute async jobs pool: diff --git a/internal/pkg/database/database.go b/internal/pkg/database/database.go index a932fa71..cd15e18f 100644 --- a/internal/pkg/database/database.go +++ b/internal/pkg/database/database.go @@ -3,6 +3,7 @@ package database import ( "context" "database/sql" + "sync" "time" "github.com/babourine/x/pkg/set" @@ -12,6 +13,8 @@ import ( const ( dbDriverName = `postgres` + maxOpenConns = 25 + maxIdleConns = 10 ) var ( @@ -21,6 +24,12 @@ var ( type Database struct { ConnectionString string `yaml:"connection_string,omitempty" json:"connection_string,omitempty"` + MaxOpenConns int `yaml:"max_open_conns,omitempty" json:"max_open_conns,omitempty"` + MaxIdleConns int `yaml:"max_idle_conns,omitempty" json:"max_idle_conns,omitempty"` + + once sync.Once + db *sql.DB + err error } type Session struct { @@ -29,6 +38,45 @@ type Session struct { committed bool } +type queryer interface { + Query(string, ...any) (*sql.Rows, error) + QueryRow(string, ...any) *sql.Row +} + +func (s *Session) queryer() queryer { + + if s.trx != nil { + return s.trx + } + + return s.db + +} + +func (d *Database) conn() (*sql.DB, error) { + + d.once.Do(func() { + d.db, d.err = sql.Open(dbDriverName, d.ConnectionString) + if d.err != nil { + return + } + maxOpen, maxIdle := d.MaxOpenConns, d.MaxIdleConns + if maxOpen <= 0 { + maxOpen = maxOpenConns + } + if maxIdle <= 0 { + maxIdle = maxIdleConns + } + d.db.SetMaxOpenConns(maxOpen) + d.db.SetMaxIdleConns(maxIdle) + d.db.SetConnMaxLifetime(30 * time.Minute) + d.db.SetConnMaxIdleTime(5 * time.Minute) + }) + + return d.db, d.err + +} + func (d *Database) NewSession(withTransaction bool) (*Session, error) { // Track session creation metrics @@ -49,7 +97,7 @@ func (d *Database) NewSession(withTransaction bool) (*Session, error) { newSessionMethod.CountRequest("with_transaction", transactionLabel) // open connection - if s.db, err = sql.Open(dbDriverName, d.ConnectionString); err != nil { + if s.db, err = d.conn(); err != nil { newSessionMethod.LogAndCountError(err, "with_transaction", transactionLabel) return nil, err } @@ -60,7 +108,6 @@ func (d *Database) NewSession(withTransaction bool) (*Session, error) { // start transaction if withTransaction { if s.trx, err = s.db.BeginTx(ctx, nil); err != nil { - s.db.Close() // Close the connection before returning error newSessionMethod.LogAndCountError(err, "with_transaction", transactionLabel) return nil, err } @@ -73,15 +120,10 @@ func (d *Database) NewSession(withTransaction bool) (*Session, error) { func (s *Session) Close() error { - // do we have an uncommitted transaction going? rollback! if s.trx != nil && !s.committed { s.trx.Rollback() } - if s.db != nil { - return s.db.Close() - } - return nil } @@ -157,7 +199,7 @@ func (s *Session) Exec(query string, args ...any) (int64, error) { func (s *Session) QueryRow(query string, args ...any) (*sql.Row, error) { - row := s.db.QueryRow(query, args...) + row := s.queryer().QueryRow(query, args...) if err := row.Err(); err != nil { return nil, err @@ -169,7 +211,7 @@ func (s *Session) QueryRow(query string, args ...any) (*sql.Row, error) { func (s *Session) Query(query string, args ...any) (*sql.Rows, error) { - return s.db.Query(query, args...) + return s.queryer().Query(query, args...) } @@ -192,6 +234,9 @@ func (s *Session) SelectSet(query string, args ...any) (*set.Set[string], error) } result.Add(item) } + if err := rows.Err(); err != nil { + return nil, err + } return result, nil diff --git a/internal/pkg/database/database_test.go b/internal/pkg/database/database_test.go new file mode 100644 index 00000000..0a7796f2 --- /dev/null +++ b/internal/pkg/database/database_test.go @@ -0,0 +1,52 @@ +package database + +import ( + "strings" + "testing" +) + +func TestSessionsSharePool(t *testing.T) { + + d := &Database{ + ConnectionString: "postgres://heimdall:heimdall@127.0.0.1:1/heimdall_pool_test?sslmode=disable", + MaxOpenConns: 4, + MaxIdleConns: 2, + } + + s1, err := d.NewSession(false) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = s1.Close() }) + + s2, err := d.NewSession(false) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = s2.Close() }) + + if s1.db != s2.db { + t.Fatal("sessions opened separate connection pools") + } + if s1.db.Stats().MaxOpenConnections != 4 { + t.Fatalf("max open connections = %d, want 4", s1.db.Stats().MaxOpenConnections) + } + + if err := s1.Close(); err != nil { + t.Fatal(err) + } + + // Closing a session must not close the shared pool. + if err := s2.db.Ping(); err != nil && strings.Contains(err.Error(), "database is closed") { + t.Fatal("closing a session closed the shared pool") + } + + // A failed transaction must not close the pool either. + if _, err := d.NewSession(true); err == nil { + t.Fatal("expected begin transaction to fail without a database") + } + if err := s2.db.Ping(); err != nil && strings.Contains(err.Error(), "database is closed") { + t.Fatal("failed NewSession closed the shared pool") + } + +}