Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions configs/local.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
63 changes: 54 additions & 9 deletions internal/pkg/database/database.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package database
import (
"context"
"database/sql"
"sync"
"time"

"github.com/babourine/x/pkg/set"
Expand All @@ -12,6 +13,8 @@ import (

const (
dbDriverName = `postgres`
maxOpenConns = 25
maxIdleConns = 10
)

var (
Expand All @@ -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 {
Expand All @@ -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
Expand All @@ -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
}
Expand All @@ -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
}
Expand All @@ -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

}
Expand Down Expand Up @@ -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
Expand All @@ -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...)

}

Expand All @@ -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

Expand Down
52 changes: 52 additions & 0 deletions internal/pkg/database/database_test.go
Original file line number Diff line number Diff line change
@@ -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")
}

}
Loading