fix(dbmanager): keep the pool alive across errors and restarts

Implements the fixes from audit/pkg/dbmanager.audit.md.

- Stop closing the shared *sql.DB to recover from errors. Adapter
  factories and the health checker no longer call Reconnect; Reconnect is
  atomic and operator-only.
- Postgres uses a custom driver.Connector: Reconnect retires pooled
  connections by generation without closing the pool, so held Bun/GORM
  handles keep working. Verified against a live server restart.
- Add TCP keepalive, TCP_USER_TIMEOUT, a bounded reuse ping and
  statement_timeout as a runtime parameter; drop the 2 min timeout floor.
- Health check pings without holding the connection lock.
- Listener: single goroutine pair, bounded Close without UNLISTEN, and
  serialised use of the pgx connection (fixes conn busy and a close race).
- Fix Connect/Close/Connect/Close panic, idempotent Connect, dial outside
  the manager lock, clean up on partial failure.
- SQLite: pin :memory: to one connection, pragmas via DSN.
- Escape credentials in Postgres/MSSQL/Mongo DSNs; sslmode defaults to
  prefer. Wire retry settings, publish metrics, fix logger calls.
- NewConnectionFromDB: Close is a no-op with a warning (caller owns the
  pool); Reconnect only pings.
- Document correct usage in the README; mark the audit with what was done.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
Hein
2026-09-30 12:40:14 +02:00
co-authored by Claude Sonnet 5.5
parent bc8bff7955
commit da1af1487e
25 changed files with 2089 additions and 719 deletions
+1 -3
View File
@@ -50,7 +50,6 @@ dbmanager:
# Health checks
health_check_interval: 30s
enable_auto_reconnect: true
connections:
# Primary PostgreSQL connection
@@ -256,7 +255,7 @@ db, _ := mgr.GetDefaultDatabase()
| `retry_delay` | duration | 1s | Initial retry delay |
| `retry_max_delay` | duration | 10s | Maximum retry delay |
| `health_check_interval` | duration | 30s | Interval between health checks |
| `enable_auto_reconnect` | bool | true | Auto-reconnect on health check failure |
| `enable_auto_reconnect` | bool | - | Deprecated and ignored: the manager never closes the pool to recover from errors |
### Connection Configuration
@@ -451,7 +450,6 @@ db.NewSelect().Model(&User{}).Scan(ctx)
3. **Enable Health Checks**: Catch connection issues early
```yaml
health_check_interval: 30s
enable_auto_reconnect: true
```
4. **Use Appropriate ORM**: Choose based on your needs
+106 -71
View File
@@ -2,6 +2,10 @@ package dbmanager
import (
"fmt"
"net"
"net/url"
"strconv"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/config"
@@ -57,9 +61,15 @@ type ManagerConfig struct {
RetryDelay time.Duration `mapstructure:"retry_delay"`
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
// Health checks
// Health checks. A zero HealthCheckInterval selects the default (15s); a
// negative value disables the background health checker.
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
// Deprecated: ignored. The manager never closes a pool to recover from an
// error because database/sql already replaces broken connections; closing
// it would invalidate every handle handed out. Use Connection.Reconnect for
// an explicit, handle-preserving refresh.
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
}
// ConnectionConfig defines configuration for a single database connection
@@ -103,6 +113,11 @@ type ConnectionConfig struct {
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
QueryTimeout time.Duration `mapstructure:"query_timeout"`
// Retry policy for the initial connect (inherited from the manager config)
RetryAttempts int `mapstructure:"retry_attempts"`
RetryDelay time.Duration `mapstructure:"retry_delay"`
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
// Features
EnableTracing bool `mapstructure:"enable_tracing"`
EnableMetrics bool `mapstructure:"enable_metrics"`
@@ -129,7 +144,6 @@ func DefaultManagerConfig() ManagerConfig {
RetryDelay: 1 * time.Second,
RetryMaxDelay: 10 * time.Second,
HealthCheckInterval: 15 * time.Second,
EnableAutoReconnect: true,
}
}
@@ -161,11 +175,6 @@ func (c *ManagerConfig) ApplyDefaults() {
if c.HealthCheckInterval == 0 {
c.HealthCheckInterval = defaults.HealthCheckInterval
}
// EnableAutoReconnect defaults to true - apply if not explicitly set
// Since this is a boolean, we apply the default unconditionally when it's false
if !c.EnableAutoReconnect {
c.EnableAutoReconnect = defaults.EnableAutoReconnect
}
}
// Validate validates the manager configuration
@@ -222,9 +231,18 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
}
if cc.QueryTimeout == 0 {
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
} else if cc.QueryTimeout < 2*time.Minute {
// Enforce minimum of 2 minutes
cc.QueryTimeout = 2 * time.Minute
}
if global != nil {
if cc.RetryAttempts == 0 {
cc.RetryAttempts = global.RetryAttempts
}
if cc.RetryDelay == 0 {
cc.RetryDelay = global.RetryDelay
}
if cc.RetryMaxDelay == 0 {
cc.RetryMaxDelay = global.RetryMaxDelay
}
}
// Default ORM
@@ -314,108 +332,122 @@ func (cc *ConnectionConfig) BuildDSN() (string, error) {
}
}
// buildPostgresDSN builds a postgres:// URL so credentials and other values are
// escaped rather than spliced into a key=value string. statement_timeout is
// applied by the provider as a runtime parameter.
func (cc *ConnectionConfig) buildPostgresDSN() string {
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s",
cc.Host, cc.Port, cc.User, cc.Password, cc.Database)
q := url.Values{}
if cc.SSLMode != "" {
dsn += fmt.Sprintf(" sslmode=%s", cc.SSLMode)
q.Set("sslmode", cc.SSLMode)
} else {
dsn += " sslmode=disable"
// prefer: use TLS when the server offers it, without failing on
// servers that do not.
q.Set("sslmode", "prefer")
}
if cc.Schema != "" {
dsn += fmt.Sprintf(" search_path=%s", cc.Schema)
q.Set("search_path", cc.Schema)
}
// Add statement_timeout for query execution timeout (in milliseconds)
if cc.QueryTimeout > 0 {
timeoutMs := int(cc.QueryTimeout.Milliseconds())
dsn += fmt.Sprintf(" statement_timeout=%d", timeoutMs)
u := url.URL{
Scheme: "postgres",
Host: hostPort(cc.Host, cc.Port),
Path: "/" + cc.Database,
RawQuery: q.Encode(),
}
return dsn
if cc.User != "" || cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
func hostPort(host string, port int) string {
if port == 0 {
return host
}
// JoinHostPort brackets IPv6 literals.
return net.JoinHostPort(host, strconv.Itoa(port))
}
// buildSQLiteDSN puts per-connection settings in the DSN as _pragma parameters
// so every pooled connection gets them, not just the one that ran an Exec.
func (cc *ConnectionConfig) buildSQLiteDSN() string {
filepath := cc.FilePath
if filepath == "" {
filepath = ":memory:"
}
// Add query parameters for timeouts
// Note: SQLite driver supports _timeout parameter (in milliseconds)
var pragmas []string
if cc.QueryTimeout > 0 {
timeoutMs := int(cc.QueryTimeout.Milliseconds())
filepath += fmt.Sprintf("?_timeout=%d", timeoutMs)
pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds()))
}
if filepath != ":memory:" {
pragmas = append(pragmas, "journal_mode(WAL)")
}
if len(pragmas) == 0 {
return filepath
}
return filepath
q := url.Values{}
for _, p := range pragmas {
q.Add("_pragma", p)
}
sep := "?"
if strings.Contains(filepath, "?") {
sep = "&"
}
return filepath + sep + q.Encode()
}
func (cc *ConnectionConfig) buildMSSQLDSN() string {
// Format: sqlserver://username:password@host:port?database=dbname
dsn := fmt.Sprintf("sqlserver://%s:%s@%s:%d?database=%s",
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
q := url.Values{}
q.Set("database", cc.Database)
if cc.Schema != "" {
dsn += fmt.Sprintf("&schema=%s", cc.Schema)
q.Set("schema", cc.Schema)
}
// Add connection timeout (in seconds)
if cc.ConnectTimeout > 0 {
timeoutSec := int(cc.ConnectTimeout.Seconds())
dsn += fmt.Sprintf("&connection timeout=%d", timeoutSec)
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
q.Set("connection timeout", sec)
q.Set("dial timeout", sec)
}
// Add dial timeout for TCP connection (in seconds)
if cc.ConnectTimeout > 0 {
dialTimeoutSec := int(cc.ConnectTimeout.Seconds())
dsn += fmt.Sprintf("&dial timeout=%d", dialTimeoutSec)
}
// Add read timeout (in seconds) - enforces timeout for reading data
if cc.QueryTimeout > 0 {
readTimeoutSec := int(cc.QueryTimeout.Seconds())
dsn += fmt.Sprintf("&read timeout=%d", readTimeoutSec)
q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds())))
}
return dsn
u := url.URL{
Scheme: "sqlserver",
Host: hostPort(cc.Host, cc.Port),
RawQuery: q.Encode(),
}
if cc.User != "" || cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
func (cc *ConnectionConfig) buildMongoDSN() string {
// Format: mongodb://username:password@host:port/database?authSource=admin
var dsn string
if cc.User != "" && cc.Password != "" {
dsn = fmt.Sprintf("mongodb://%s:%s@%s:%d/%s",
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
} else {
dsn = fmt.Sprintf("mongodb://%s:%d/%s", cc.Host, cc.Port, cc.Database)
}
params := ""
q := url.Values{}
if cc.AuthSource != "" {
params += fmt.Sprintf("authSource=%s", cc.AuthSource)
q.Set("authSource", cc.AuthSource)
}
if cc.ReplicaSet != "" {
if params != "" {
params += "&"
}
params += fmt.Sprintf("replicaSet=%s", cc.ReplicaSet)
q.Set("replicaSet", cc.ReplicaSet)
}
if cc.ReadPreference != "" {
if params != "" {
params += "&"
}
params += fmt.Sprintf("readPreference=%s", cc.ReadPreference)
q.Set("readPreference", cc.ReadPreference)
}
if params != "" {
dsn += "?" + params
u := url.URL{
Scheme: "mongodb",
Host: hostPort(cc.Host, cc.Port),
Path: "/" + cc.Database,
RawQuery: q.Encode(),
}
return dsn
if cc.User != "" && cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
@@ -487,3 +519,6 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
+95
View File
@@ -0,0 +1,95 @@
package dbmanager
import (
"net/url"
"strings"
"testing"
"time"
)
func TestPostgresDSNEscapesCredentials(t *testing.T) {
cc := ConnectionConfig{
Type: DatabaseTypePostgreSQL, Host: "db", Port: 5432, Database: "app",
User: "u@x", Password: "p w'd sslmode=disable&x=y/?#",
}
dsn := cc.buildPostgresDSN()
u, err := url.Parse(dsn)
if err != nil {
t.Fatalf("DSN is not a valid URL: %v", err)
}
if pw, _ := u.User.Password(); pw != cc.Password {
t.Errorf("password did not round-trip: %q", pw)
}
if u.User.Username() != cc.User {
t.Errorf("user did not round-trip: %q", u.User.Username())
}
if got := u.Query().Get("sslmode"); got != "prefer" {
t.Errorf("sslmode = %q, want prefer (password must not inject parameters)", got)
}
}
func TestMSSQLAndMongoDSNEscapeCredentials(t *testing.T) {
cc := ConnectionConfig{Host: "h", Port: 1, Database: "d", User: "u", Password: "a@b:c/d?e&f"}
for name, dsn := range map[string]string{"mssql": cc.buildMSSQLDSN(), "mongo": cc.buildMongoDSN()} {
u, err := url.Parse(dsn)
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if pw, _ := u.User.Password(); pw != cc.Password {
t.Errorf("%s: password did not round-trip: %q", name, pw)
}
if u.Host != "h:1" {
t.Errorf("%s: host = %q", name, u.Host)
}
}
}
func TestSQLiteDSNUsesPragmas(t *testing.T) {
cc := ConnectionConfig{FilePath: "/tmp/x.db", QueryTimeout: 3 * time.Second}
dsn := cc.buildSQLiteDSN()
if strings.Contains(dsn, "?_timeout=") {
t.Errorf("unsupported _timeout parameter present: %s", dsn)
}
if !strings.Contains(dsn, "busy_timeout%283000%29") || !strings.Contains(dsn, "journal_mode%28WAL%29") {
t.Errorf("expected busy_timeout and WAL pragmas in DSN: %s", dsn)
}
}
func TestQueryTimeoutHonoredWithoutFloor(t *testing.T) {
cc := ConnectionConfig{QueryTimeout: 30 * time.Second}
cc.ApplyDefaults(&ManagerConfig{})
if cc.QueryTimeout != 30*time.Second {
t.Errorf("QueryTimeout = %v, want 30s", cc.QueryTimeout)
}
}
func TestRetryPolicyInherited(t *testing.T) {
g := ManagerConfig{RetryAttempts: 5, RetryDelay: time.Second, RetryMaxDelay: time.Minute}
cc := ConnectionConfig{}
cc.ApplyDefaults(&g)
if cc.GetRetryAttempts() != 5 || cc.GetRetryMaxDelay() != time.Minute {
t.Errorf("retry policy not inherited: %+v", cc)
}
}
func TestSQLiteMemoryPoolPinned(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
if err := mgr.Connect(t.Context()); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
db, _ := conn.Native()
if got := db.Stats().MaxOpenConnections; got != 1 {
t.Fatalf("MaxOpenConnections = %d, want 1 for :memory:", got)
}
if _, err := db.Exec("CREATE TABLE t(a int)"); err != nil {
t.Fatal(err)
}
for i := 0; i < 5; i++ {
var n int
if err := db.QueryRow("SELECT count(*) FROM t").Scan(&n); err != nil {
t.Fatalf("table missing on later use: %v", err)
}
}
}
+143 -68
View File
@@ -3,6 +3,7 @@ package dbmanager
import (
"context"
"database/sql"
"errors"
"fmt"
"sync"
"time"
@@ -14,6 +15,7 @@ import (
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
)
// Connection represents a single named database connection
@@ -82,6 +84,9 @@ type sqlConnection struct {
// State
connected bool
mu sync.RWMutex
// lifecycleMu serialises Connect/Close/Reconnect against health-check pings.
// Lock order: lifecycleMu before mu.
lifecycleMu sync.RWMutex
// Health check
lastHealthCheck time.Time
@@ -110,9 +115,16 @@ func (c *sqlConnection) Type() DatabaseType {
// Connect establishes the database connection
func (c *sqlConnection) Connect(ctx context.Context) error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.connectLocked(ctx)
}
// connectLocked requires lifecycleMu and mu held for writing.
func (c *sqlConnection) connectLocked(ctx context.Context) error {
if c.connected {
return ErrAlreadyConnected
}
@@ -127,17 +139,29 @@ func (c *sqlConnection) Connect(ctx context.Context) error {
// Close closes the database connection and all ORM instances
func (c *sqlConnection) Close() error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.closeLocked()
}
// closeLocked requires lifecycleMu and mu held for writing. The connection is
// always marked disconnected and its cached handles dropped, even when closing
// fails, so accessors never hand out handles over a half-closed pool.
func (c *sqlConnection) closeLocked() error {
if !c.connected {
return nil
}
// Close Bun if initialized
if c.bunDB != nil {
var errs []error
// Close Bun if initialized. bun.DB.Close closes the underlying *sql.DB, so
// skip it when the pool belongs to the caller.
if o, ok := c.provider.(interface{ OwnsDB() bool }); c.bunDB != nil && (!ok || o.OwnsDB()) {
if err := c.bunDB.Close(); err != nil {
return NewConnectionError(c.name, "close bun", err)
errs = append(errs, NewConnectionError(c.name, "close bun", err))
}
}
@@ -145,7 +169,7 @@ func (c *sqlConnection) Close() error {
// Close the provider (which closes the underlying sql.DB)
if err := c.provider.Close(); err != nil {
return NewConnectionError(c.name, "close", err)
errs = append(errs, NewConnectionError(c.name, "close", err))
}
c.connected = false
@@ -156,39 +180,75 @@ func (c *sqlConnection) Close() error {
c.gormAdapter = nil
c.nativeAdapter = nil
return nil
return errors.Join(errs...)
}
// HealthCheck verifies the connection is alive
// HealthCheck verifies the connection is alive. The network ping runs without
// holding mu, so handle accessors are never blocked behind a slow ping.
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
if c == nil {
return fmt.Errorf("connection is nil")
}
c.mu.Lock()
defer c.mu.Unlock()
c.lastHealthCheck = time.Now()
// lifecycleMu (read) keeps Close/Reconnect from tearing the provider down
// mid-ping without blocking the accessors that only need mu.
c.lifecycleMu.RLock()
defer c.lifecycleMu.RUnlock()
if !c.connected {
c.healthCheckStatus = "disconnected"
c.mu.RLock()
connected := c.connected
provider := c.provider
c.mu.RUnlock()
if !connected {
c.setHealth("disconnected")
return ErrConnectionClosed
}
if err := c.provider.HealthCheck(ctx); err != nil {
c.healthCheckStatus = "unhealthy: " + err.Error()
if err := provider.HealthCheck(ctx); err != nil {
c.setHealth("unhealthy: " + err.Error())
return NewConnectionError(c.name, "health check", err)
}
c.healthCheckStatus = "healthy"
c.setHealth("healthy")
return nil
}
// Reconnect closes and re-establishes the connection
func (c *sqlConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
func (c *sqlConnection) setHealth(status string) {
c.mu.Lock()
c.lastHealthCheck = time.Now()
c.healthCheckStatus = status
c.mu.Unlock()
}
// Reconnect refreshes the connection as a single critical section.
//
// Providers that support it (PostgreSQL) retire their pooled connections and
// dial fresh ones without closing the *sql.DB, so handles handed out earlier
// keep working. Other providers fall back to Close+Connect, which invalidates
// earlier handles; that is meant for explicit operator use only, since
// *sql.DB already replaces broken connections by itself.
func (c *sqlConnection) Reconnect(ctx context.Context) (err error) {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
defer func() { RecordReconnectAttempt(c.name, c.dbType, err == nil) }()
if c.connected {
if r, ok := c.provider.(providers.Refresher); ok {
if err := r.Refresh(ctx); err != nil {
return NewConnectionError(c.name, "reconnect", err)
}
return nil
}
}
if err := c.closeLocked(); err != nil {
return err
}
return c.Connect(ctx)
return c.connectLocked(ctx)
}
// Native returns the native *sql.DB connection
@@ -250,6 +310,10 @@ func (c *sqlConnection) Bun() (*bun.DB, error) {
return c.bunDB, nil
}
if !c.connected {
return nil, ErrConnectionClosed
}
// Get native connection first
native, err := c.provider.GetNative()
if err != nil {
@@ -283,6 +347,10 @@ func (c *sqlConnection) GORM() (*gorm.DB, error) {
return c.gormDB, nil
}
if !c.connected {
return nil, ErrConnectionClosed
}
// Get native connection first
native, err := c.provider.GetNative()
if err != nil {
@@ -359,39 +427,18 @@ func (c *sqlConnection) Stats() *ConnectionStats {
return stats
}
func (c *sqlConnection) reconnectForAdapter() error {
timeout := c.config.ConnectTimeout
if timeout <= 0 {
timeout = 10 * time.Second
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
return c.Reconnect(ctx)
}
// The adapter factories only re-fetch the current handle. They must not close
// the shared pool: *sql.DB discards bad connections on its own, and closing it
// here would break every other holder of the pool.
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.Native()
}
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.Bun()
}
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.GORM()
}
@@ -512,15 +559,8 @@ func (c *sqlConnection) getNativeAdapter() (common.Database, error) {
// Create a native adapter based on database type
switch c.dbType {
case DatabaseTypePostgreSQL:
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
case DatabaseTypeSQLite:
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
case DatabaseTypeMSSQL:
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
// The adapter takes the driver name so it can adjust its dialect.
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
@@ -572,8 +612,9 @@ type mongoConnection struct {
client *mongo.Client
// State
connected bool
mu sync.RWMutex
connected bool
mu sync.RWMutex
lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu
// Health check
lastHealthCheck time.Time
@@ -601,9 +642,16 @@ func (c *mongoConnection) Type() DatabaseType {
// Connect establishes the MongoDB connection
func (c *mongoConnection) Connect(ctx context.Context) error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.connectLocked(ctx)
}
// connectLocked requires lifecycleMu and mu held for writing.
func (c *mongoConnection) connectLocked(ctx context.Context) error {
if c.connected {
return ErrAlreadyConnected
}
@@ -615,6 +663,7 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
// Get the mongo client
client, err := c.provider.GetMongo()
if err != nil {
_ = c.provider.Close()
return NewConnectionError(c.name, "get mongo client", err)
}
@@ -625,49 +674,75 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
// Close closes the MongoDB connection
func (c *mongoConnection) Close() error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.closeLocked()
}
// closeLocked requires lifecycleMu and mu held for writing. The connection is
// marked disconnected even when the provider fails to close cleanly.
func (c *mongoConnection) closeLocked() error {
if !c.connected {
return nil
}
if err := c.provider.Close(); err != nil {
return NewConnectionError(c.name, "close", err)
}
err := c.provider.Close()
c.connected = false
c.client = nil
if err != nil {
return NewConnectionError(c.name, "close", err)
}
return nil
}
// HealthCheck verifies the MongoDB connection is alive
// HealthCheck verifies the MongoDB connection is alive. The ping runs without
// holding mu so handle accessors are never blocked behind it.
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
c.mu.Lock()
defer c.mu.Unlock()
c.lifecycleMu.RLock()
defer c.lifecycleMu.RUnlock()
c.lastHealthCheck = time.Now()
c.mu.RLock()
connected := c.connected
c.mu.RUnlock()
if !c.connected {
c.healthCheckStatus = "disconnected"
if !connected {
c.setHealth("disconnected")
return ErrConnectionClosed
}
if err := c.provider.HealthCheck(ctx); err != nil {
c.healthCheckStatus = "unhealthy: " + err.Error()
c.setHealth("unhealthy: " + err.Error())
return NewConnectionError(c.name, "health check", err)
}
c.healthCheckStatus = "healthy"
c.setHealth("healthy")
return nil
}
// Reconnect closes and re-establishes the MongoDB connection
func (c *mongoConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
func (c *mongoConnection) setHealth(status string) {
c.mu.Lock()
c.lastHealthCheck = time.Now()
c.healthCheckStatus = status
c.mu.Unlock()
}
// Reconnect closes and re-establishes the MongoDB connection atomically.
func (c *mongoConnection) Reconnect(ctx context.Context) (err error) {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
defer func() { RecordReconnectAttempt(c.name, DatabaseTypeMongoDB, err == nil) }()
if err := c.closeLocked(); err != nil {
return err
}
return c.Connect(ctx)
return c.connectLocked(ctx)
}
// MongoDB returns the MongoDB client
-159
View File
@@ -4,13 +4,8 @@ import (
"context"
"database/sql"
"testing"
"time"
_ "github.com/mattn/go-sqlite3"
"gorm.io/gorm"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
)
func TestNewConnectionFromDB(t *testing.T) {
@@ -213,157 +208,3 @@ func TestNewConnectionFromDB_PostgreSQL(t *testing.T) {
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
}
}
func TestDatabaseNativeAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-native", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-native",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeNative),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.PgSQLAdapter)
if !ok {
t.Fatalf("Expected PgSQLAdapter, got %T", db)
}
underlyingBefore, ok := adapter.GetUnderlyingDB().(*sql.DB)
if !ok {
t.Fatalf("Expected underlying *sql.DB, got %T", adapter.GetUnderlyingDB())
}
if err := underlyingBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying database: %v", err)
}
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
t.Fatalf("Expected native adapter to reconnect, got error: %v", err)
}
underlyingAfter, ok := adapter.GetUnderlyingDB().(*sql.DB)
if !ok {
t.Fatalf("Expected reconnected *sql.DB, got %T", adapter.GetUnderlyingDB())
}
if underlyingAfter == underlyingBefore {
t.Fatal("Expected adapter to swap to a fresh *sql.DB after reconnect")
}
}
func TestDatabaseBunAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-bun", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-bun",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeBun),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.BunAdapter)
if !ok {
t.Fatalf("Expected BunAdapter, got %T", db)
}
underlyingBefore, ok := adapter.GetUnderlyingDB().(interface{ Close() error })
if !ok {
t.Fatalf("Expected underlying Bun DB with Close method, got %T", adapter.GetUnderlyingDB())
}
if err := underlyingBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying Bun database: %v", err)
}
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
t.Fatalf("Expected Bun adapter to reconnect, got error: %v", err)
}
underlyingAfter := adapter.GetUnderlyingDB()
if underlyingAfter == underlyingBefore {
t.Fatal("Expected adapter to swap to a fresh Bun DB after reconnect")
}
}
func TestDatabaseGormAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-gorm", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-gorm",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeGORM),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.GormAdapter)
if !ok {
t.Fatalf("Expected GormAdapter, got %T", db)
}
gormBefore, ok := adapter.GetUnderlyingDB().(*gorm.DB)
if !ok {
t.Fatalf("Expected underlying *gorm.DB, got %T", adapter.GetUnderlyingDB())
}
sqlBefore, err := gormBefore.DB()
if err != nil {
t.Fatalf("Failed to get underlying *sql.DB: %v", err)
}
if err := sqlBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying database: %v", err)
}
count, err := db.NewSelect().Table("sqlite_master").Count(ctx)
if err != nil {
t.Fatalf("Expected GORM query builder to reconnect, got error: %v", err)
}
if count < 0 {
t.Fatalf("Expected non-negative count, got %d", count)
}
gormAfter, ok := adapter.GetUnderlyingDB().(*gorm.DB)
if !ok {
t.Fatalf("Expected reconnected *gorm.DB, got %T", adapter.GetUnderlyingDB())
}
sqlAfter, err := gormAfter.DB()
if err != nil {
t.Fatalf("Failed to get reconnected *sql.DB: %v", err)
}
if sqlAfter == sqlBefore {
t.Fatal("Expected GORM adapter to use a fresh *sql.DB after reconnect")
}
}
+199
View File
@@ -0,0 +1,199 @@
package dbmanager
import (
"context"
"database/sql"
"sync"
"testing"
"time"
)
func sqliteManagerConfig() ManagerConfig {
return ManagerConfig{
DefaultConnection: "test",
Connections: map[string]ConnectionConfig{
"test": {Name: "test", Type: DatabaseTypeSQLite, FilePath: ":memory:"},
},
HealthCheckInterval: time.Hour,
}
}
func TestManagerConnectCloseCycleTwice(t *testing.T) {
mgr, err := NewManager(sqliteManagerConfig())
if err != nil {
t.Fatal(err)
}
ctx := context.Background()
for i := 0; i < 2; i++ {
if err := mgr.Connect(ctx); err != nil {
t.Fatalf("cycle %d connect: %v", i, err)
}
cm := mgr.(*connectionManager)
cm.healthMu.Lock()
running := cm.healthTicker != nil
cm.healthMu.Unlock()
if !running {
t.Fatalf("cycle %d: health checker not running", i)
}
if err := mgr.Close(); err != nil {
t.Fatalf("cycle %d close: %v", i, err)
}
}
// A further Close must not panic.
if err := mgr.Close(); err != nil {
t.Fatal(err)
}
}
func TestManagerConnectIsIdempotent(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
if got := mgr.Stats().TotalConnections; got != 1 {
t.Fatalf("expected 1 connection, got %d", got)
}
}
func TestConcurrentReconnectIsAtomic(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
var wg sync.WaitGroup
for i := 0; i < 20; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if err := conn.Reconnect(ctx); err != nil {
t.Errorf("reconnect: %v", err)
}
}()
}
wg.Wait()
db, err := conn.Native()
if err != nil {
t.Fatal(err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("pool unusable after concurrent reconnects: %v", err)
}
}
func TestAdapterFactoryDoesNotClosePool(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
held, _ := sc.Native()
if _, err := sc.reopenNativeForAdapter(); err != nil {
t.Fatal(err)
}
if err := held.PingContext(ctx); err != nil {
t.Fatalf("existing handle broken by adapter factory: %v", err)
}
}
func TestHealthCheckDoesNotBlockAccessors(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
// Simulate a health check in flight: it holds lifecycleMu (read) only.
sc.lifecycleMu.RLock()
defer sc.lifecycleMu.RUnlock()
done := make(chan struct{})
go func() {
_, _ = sc.Bun()
_, _ = sc.GORM()
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("accessors blocked while health check in flight")
}
}
func TestCloseAlwaysMarksDisconnected(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
_, _ = sc.Bun()
_ = mgr.Close()
if _, err := sc.Bun(); err == nil {
t.Fatal("Bun() should fail after Close")
}
if _, err := sc.GORM(); err == nil {
t.Fatal("GORM() should fail after Close")
}
}
func TestReconnectOnExistingDBKeepsCallersPool(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatal(err)
}
if err := conn.Reconnect(ctx); err != nil {
t.Fatalf("reconnect: %v", err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("caller's pool was closed by Reconnect: %v", err)
}
}
func TestCloseOnExistingDBLeavesCallersPoolOpen(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatal(err)
}
if _, err := conn.Bun(); err != nil { // bun.DB.Close would close the pool
t.Fatal(err)
}
if err := conn.Close(); err != nil {
t.Fatal(err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("caller's pool was closed: %v", err)
}
}
+66 -52
View File
@@ -2,9 +2,7 @@ package dbmanager
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"time"
@@ -49,6 +47,7 @@ type connectionManager struct {
// Background health check
healthTicker *time.Ticker
stopChan chan struct{}
healthMu sync.Mutex // guards healthTicker and stopChan
wg sync.WaitGroup
}
@@ -100,7 +99,9 @@ func ResetInstance() {
defer instanceMu.Unlock()
if instance != nil {
_ = instance.Close()
if err := instance.Close(); err != nil {
logger.Error("Failed to close manager during reset: %v", err)
}
}
instance = nil
}
@@ -116,7 +117,6 @@ func NewManager(cfg ManagerConfig) (Manager, error) {
mgr := &connectionManager{
connections: make(map[string]Connection),
config: cfg,
stopChan: make(chan struct{}),
}
return mgr, nil
@@ -195,11 +195,26 @@ func (m *connectionManager) SetDefaultDatabase(name string) error {
// Connect establishes all configured database connections
func (m *connectionManager) Connect(ctx context.Context) error {
m.mu.Lock()
defer m.mu.Unlock()
// Create connections from configuration
// Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck.
m.mu.RLock()
names := make([]string, 0, len(m.config.Connections))
for name := range m.config.Connections {
if _, exists := m.connections[name]; !exists {
names = append(names, name)
}
}
m.mu.RUnlock()
opened := make(map[string]Connection, len(names))
closeOpened := func() {
for name, conn := range opened {
if err := conn.Close(); err != nil {
logger.Error("Failed to close connection after failed Connect: name=%s, error=%v", name, err)
}
}
}
for _, name := range names {
// Get a copy of the connection config
connCfg := m.config.Connections[name]
// Apply global defaults to connection config
@@ -209,25 +224,39 @@ func (m *connectionManager) Connect(ctx context.Context) error {
// Create connection using factory
conn, err := createConnection(connCfg)
if err != nil {
closeOpened()
return fmt.Errorf("failed to create connection '%s': %w", name, err)
}
// Connect
if err := conn.Connect(ctx); err != nil {
closeOpened()
return fmt.Errorf("failed to connect '%s': %w", name, err)
}
m.connections[name] = conn
opened[name] = conn
logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type)
}
m.mu.Lock()
for name, conn := range opened {
if _, exists := m.connections[name]; exists {
// Lost a race with a concurrent Connect; drop our duplicate.
_ = conn.Close()
continue
}
m.connections[name] = conn
}
total := len(m.connections)
m.mu.Unlock()
// Always start background health checks
if m.config.HealthCheckInterval > 0 {
m.startHealthChecker()
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
}
logger.Info("Database manager initialized: connections=%d", len(m.connections))
logger.Info("Database manager initialized: connections=%d", total)
return nil
}
@@ -246,7 +275,7 @@ func (m *connectionManager) Close() error {
for name, conn := range m.connections {
if err := conn.Close(); err != nil {
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
logger.Error("Failed to close connection", "name", name, "error", err)
logger.Error("Failed to close connection: name=%s, error=%v", name, err)
} else {
logger.Info("Connection closed: name=%s", name)
}
@@ -311,11 +340,17 @@ func (m *connectionManager) Stats() *ManagerStats {
// startHealthChecker starts background health checking
func (m *connectionManager) startHealthChecker() {
m.healthMu.Lock()
defer m.healthMu.Unlock()
if m.healthTicker != nil {
return // Already running
}
m.healthTicker = time.NewTicker(m.config.HealthCheckInterval)
ticker := time.NewTicker(m.config.HealthCheckInterval)
stop := make(chan struct{})
m.healthTicker = ticker
m.stopChan = stop
m.wg.Add(1)
go func() {
@@ -324,9 +359,9 @@ func (m *connectionManager) startHealthChecker() {
for {
select {
case <-m.healthTicker.C:
case <-ticker.C:
m.performHealthCheck()
case <-m.stopChan:
case <-stop:
logger.Info("Health checker stopped")
return
}
@@ -334,14 +369,19 @@ func (m *connectionManager) startHealthChecker() {
}()
}
// stopHealthChecker stops background health checking
// stopHealthChecker stops background health checking. Safe to call repeatedly.
func (m *connectionManager) stopHealthChecker() {
if m.healthTicker != nil {
m.healthTicker.Stop()
close(m.stopChan)
m.wg.Wait()
m.healthTicker = nil
m.healthMu.Lock()
defer m.healthMu.Unlock()
if m.healthTicker == nil {
return
}
m.healthTicker.Stop()
close(m.stopChan)
m.wg.Wait()
m.healthTicker = nil
m.stopChan = nil
}
// performHealthCheck performs a health check on all connections
@@ -362,40 +402,14 @@ func (m *connectionManager) performHealthCheck() {
}
m.mu.RUnlock()
defer m.PublishMetrics()
for _, item := range connections {
if err := item.conn.HealthCheck(ctx); err != nil {
logger.Warn("Health check failed",
"connection", item.name,
"error", err)
// Only reconnect when the client handle itself is closed/disconnected.
// For transient database restarts or network blips, *sql.DB can recover
// on its own; forcing Close()+Connect() here invalidates any cached ORM
// wrappers and callers that still hold the old handle.
if m.config.EnableAutoReconnect && shouldReconnectAfterHealthCheck(err) {
logger.Info("Attempting reconnection: connection=%s", item.name)
if err := item.conn.Reconnect(ctx); err != nil {
logger.Error("Reconnection failed",
"connection", item.name,
"error", err)
} else {
logger.Info("Reconnection successful: connection=%s", item.name)
}
} else if m.config.EnableAutoReconnect {
logger.Info("Skipping reconnect for transient health check failure: connection=%s", item.name)
}
// Do not reconnect here: *sql.DB discards bad connections and dials
// new ones by itself, while Reconnect closes the pool and breaks
// every handle already handed out. Reconnect is operator-only.
logger.Warn("Health check failed: connection=%s, error=%v", item.name, err)
}
}
}
func shouldReconnectAfterHealthCheck(err error) bool {
if err == nil {
return false
}
if errors.Is(err, ErrConnectionClosed) {
return true
}
return strings.Contains(err.Error(), "sql: database is closed")
}
+36 -45
View File
@@ -21,19 +21,30 @@ type healthCheckStubConnection struct {
reconnectCalls int
}
func (c *healthCheckStubConnection) Name() string { return "stub" }
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Native() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Database() (common.Database, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil }
func (c *healthCheckStubConnection) Close() error { return nil }
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error { c.reconnectCalls++; return nil }
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
func (c *healthCheckStubConnection) Name() string { return "stub" }
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) Native() (*sql.DB, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Database() (common.Database, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil }
func (c *healthCheckStubConnection) Close() error { return nil }
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error {
c.reconnectCalls++
return nil
}
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
func TestBackgroundHealthChecker(t *testing.T) {
// Create a SQLite in-memory database
@@ -117,41 +128,21 @@ func TestDefaultHealthCheckInterval(t *testing.T) {
t.Errorf("Expected default health check interval to be %v, got %v",
expectedInterval, defaults.HealthCheckInterval)
}
if !defaults.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be true by default")
}
}
func TestApplyDefaultsEnablesAutoReconnect(t *testing.T) {
// Create a config without setting EnableAutoReconnect
cfg := ManagerConfig{
Connections: map[string]ConnectionConfig{
"test": {
Name: "test",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
},
},
}
// Verify it's false initially (Go's zero value for bool)
if cfg.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be false before ApplyDefaults")
}
// Apply defaults
func TestApplyDefaultsHealthCheckInterval(t *testing.T) {
cfg := ManagerConfig{}
cfg.ApplyDefaults()
// Verify it's now true
if !cfg.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be true after ApplyDefaults")
}
// Verify health check interval is also set
if cfg.HealthCheckInterval != 15*time.Second {
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
}
// A negative interval disables the background checker and is preserved.
cfg = ManagerConfig{HealthCheckInterval: -1}
cfg.ApplyDefaults()
if cfg.HealthCheckInterval >= 0 {
t.Errorf("Expected negative interval to be preserved, got %v", cfg.HealthCheckInterval)
}
}
func TestManagerHealthCheck(t *testing.T) {
@@ -270,7 +261,7 @@ func TestPerformHealthCheckSkipsReconnectForTransientFailures(t *testing.T) {
}
}
func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
func TestPerformHealthCheckNeverReconnects(t *testing.T) {
conn := &healthCheckStubConnection{
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
}
@@ -284,7 +275,7 @@ func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
mgr.performHealthCheck()
if conn.reconnectCalls != 1 {
t.Fatalf("expected reconnect attempt for closed database handle, got %d", conn.reconnectCalls)
if conn.reconnectCalls != 0 {
t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls)
}
}
+39 -15
View File
@@ -1,6 +1,8 @@
package dbmanager
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
)
@@ -34,8 +36,8 @@ var (
)
// connectionWaitCount tracks how many times connections had to wait for availability
connectionWaitCount = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionWaitCount = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_wait_count",
Help: "Number of times connections had to wait for availability",
},
@@ -43,8 +45,8 @@ var (
)
// connectionWaitDuration tracks total time connections spent waiting
connectionWaitDuration = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionWaitDuration = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_wait_duration_seconds",
Help: "Total time connections spent waiting for availability",
},
@@ -61,8 +63,8 @@ var (
)
// connectionLifetimeClosed tracks connections closed due to max lifetime
connectionLifetimeClosed = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionLifetimeClosed = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_lifetime_closed_total",
Help: "Total connections closed due to exceeding max lifetime",
},
@@ -70,8 +72,8 @@ var (
)
// connectionIdleClosed tracks connections closed due to max idle time
connectionIdleClosed = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionIdleClosed = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_idle_closed_total",
Help: "Total connections closed due to exceeding max idle time",
},
@@ -114,13 +116,13 @@ func (m *connectionManager) PublishMetrics() {
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
// Wait stats
connectionWaitCount.With(labels).Set(float64(connStats.WaitCount))
connectionWaitDuration.With(labels).Set(connStats.WaitDuration.Seconds())
// Lifetime/idle closed
connectionLifetimeClosed.With(labels).Set(float64(connStats.MaxLifetimeClosed))
connectionIdleClosed.With(labels).Set(float64(connStats.MaxIdleClosed))
// sql.DBStats values are cumulative, so add only the growth since
// the last publish to keep these true counters.
prev := lastPublished.swap(name, connStats)
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed))
}
}
}
@@ -134,3 +136,25 @@ func RecordReconnectAttempt(name string, dbType DatabaseType, success bool) {
reconnectAttempts.WithLabelValues(name, string(dbType), result).Inc()
}
// publishedStats remembers the cumulative pool stats last exported per
// connection so counters can be advanced by the delta.
type publishedStats struct {
mu sync.Mutex
last map[string]ConnectionStats
}
var lastPublished = &publishedStats{last: make(map[string]ConnectionStats)}
// swap stores cur and returns the previous value. A counter reset (a new pool
// after Close+Connect) is treated as starting from zero.
func (p *publishedStats) swap(name string, cur *ConnectionStats) ConnectionStats {
p.mu.Lock()
defer p.mu.Unlock()
prev := p.last[name]
if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed {
prev = ConnectionStats{}
}
p.last[name] = *cur
return prev
}
+93
View File
@@ -0,0 +1,93 @@
package dbmanager
import (
"context"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
"os"
"testing"
"time"
)
func TestLivePostgresRefreshKeepsHandles(t *testing.T) {
if os.Getenv("PG_LIVE") == "" {
t.Skip("PG_LIVE not set")
}
mgr, err := NewManager(ManagerConfig{
DefaultConnection: "pg",
Connections: map[string]ConnectionConfig{"pg": {
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", QueryTimeout: 30 * time.Second,
}},
HealthCheckInterval: -1,
})
if err != nil {
t.Fatal(err)
}
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
held, _ := conn.Bun()
gormDB, _ := conn.GORM()
var pid1, pid2 int
var st string
if err := held.DB.QueryRow("select pg_backend_pid(), current_setting('statement_timeout')").Scan(&pid1, &st); err != nil {
t.Fatal(err)
}
if st != "30s" {
t.Errorf("statement_timeout = %q", st)
}
if err := conn.Reconnect(ctx); err != nil {
t.Fatal(err)
}
if err := held.DB.QueryRow("select pg_backend_pid()").Scan(&pid2); err != nil {
t.Fatalf("held bun handle broken after reconnect: %v", err)
}
if pid1 == pid2 {
t.Error("expected a new backend after reconnect")
}
var n int
if err := gormDB.Raw("select 1").Scan(&n).Error; err != nil || n != 1 {
t.Fatalf("held gorm handle broken: %v", err)
}
}
func TestLiveListenerListenNotify(t *testing.T) {
if os.Getenv("PG_LIVE") == "" {
t.Skip("PG_LIVE not set")
}
cc := ConnectionConfig{Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", ConnectTimeout: 5 * time.Second}
p := providers.NewPostgresProvider()
ctx := context.Background()
if err := p.Connect(ctx, &cc); err != nil {
t.Fatal(err)
}
defer p.Close()
l, err := p.GetListener(ctx)
if err != nil {
t.Fatal(err)
}
got := make(chan string, 4)
for _, ch := range []string{"a", "b"} {
if err := l.Listen(ch, func(c, payload string) { got <- c + ":" + payload }); err != nil {
t.Fatal(err)
}
}
for i := 0; i < 3; i++ {
if err := l.Notify(ctx, "a", "x"); err != nil {
t.Fatalf("notify: %v", err)
}
}
select {
case v := <-got:
if v != "a:x" {
t.Fatalf("got %q", v)
}
case <-time.After(3 * time.Second):
t.Fatal("no notification")
}
}
+19 -6
View File
@@ -7,6 +7,8 @@ import (
"sync"
"go.mongodb.org/mongo-driver/mongo"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ExistingDBProvider wraps an existing *sql.DB connection
@@ -44,16 +46,27 @@ func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig)
return nil
}
// Close closes the underlying database connection
func (p *ExistingDBProvider) Close() error {
p.mu.Lock()
defer p.mu.Unlock()
// Refresh verifies the wrapped database is still reachable. The pool belongs to
// the caller and cannot be re-dialed here, so it is never closed to "reconnect".
func (p *ExistingDBProvider) Refresh(ctx context.Context) error {
p.mu.RLock()
defer p.mu.RUnlock()
if p.db == nil {
return nil
return fmt.Errorf("database connection is nil")
}
return p.db.PingContext(ctx)
}
return p.db.Close()
// OwnsDB reports whether Close releases the wrapped database. It never does:
// the *sql.DB was opened by the caller, who is responsible for closing it.
func (p *ExistingDBProvider) OwnsDB() bool { return false }
// Close is a no-op for the wrapped database. The pool belongs to the caller, so
// closing it here would break the caller's other users of it.
func (p *ExistingDBProvider) Close() error {
logger.Warn("Not closing externally provided database: name=%s; the caller owns this *sql.DB and must close it", p.name)
return nil
}
// HealthCheck verifies the connection is alive
+5 -5
View File
@@ -164,7 +164,7 @@ func TestExistingDBProvider_Stats(t *testing.T) {
}
}
func TestExistingDBProvider_Close(t *testing.T) {
func TestExistingDBProvider_Close_LeavesDBOpen(t *testing.T) {
db, err := sql.Open("sqlite3", ":memory:")
if err != nil {
t.Fatalf("Failed to open database: %v", err)
@@ -177,10 +177,10 @@ func TestExistingDBProvider_Close(t *testing.T) {
t.Errorf("Expected Close to succeed, got error: %v", err)
}
// Verify the database is closed
err = db.Ping()
if err == nil {
t.Error("Expected database to be closed")
// The caller owns the database, so Close must leave it open
defer db.Close()
if err := db.Ping(); err != nil {
t.Errorf("Expected caller's database to stay open, got: %v", err)
}
}
+8 -8
View File
@@ -41,9 +41,10 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
clientOpts.SetMaxPoolSize(maxPoolSize)
}
if cfg.GetMaxIdleConns() != nil {
minPoolSize := uint64(*cfg.GetMaxIdleConns())
clientOpts.SetMinPoolSize(minPoolSize)
// MaxIdleConns is a ceiling on idle connections, not a pre-warmed minimum
// (MinPoolSize), so only the idle-time limit maps onto the Mongo pool.
if cfg.GetConnMaxIdleTime() != nil {
clientOpts.SetMaxConnIdleTime(*cfg.GetConnMaxIdleTime())
}
// Set timeouts
@@ -65,12 +66,11 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
var client *mongo.Client
var lastErr error
retryAttempts := 3
retryDelay := 1 * time.Second
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if cfg.GetEnableLogging() {
logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
@@ -87,7 +87,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
if err != nil {
lastErr = err
if cfg.GetEnableLogging() {
logger.Warn("Failed to connect to MongoDB", "error", err)
logger.Warn("Failed to connect to MongoDB: %v", err)
}
continue
}
@@ -101,7 +101,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
lastErr = err
_ = client.Disconnect(ctx)
if cfg.GetEnableLogging() {
logger.Warn("Failed to ping MongoDB", "error", err)
logger.Warn("Failed to ping MongoDB: %v", err)
}
continue
}
+4 -5
View File
@@ -35,12 +35,11 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
var db *sql.DB
var lastErr error
retryAttempts := 3 // Default retry attempts
retryDelay := 1 * time.Second
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if cfg.GetEnableLogging() {
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
@@ -57,7 +56,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
if err != nil {
lastErr = err
if cfg.GetEnableLogging() {
logger.Warn("Failed to open MSSQL connection", "error", err)
logger.Warn("Failed to open MSSQL connection: %v", err)
}
continue
}
@@ -71,7 +70,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
lastErr = err
db.Close()
if cfg.GetEnableLogging() {
logger.Warn("Failed to ping MSSQL database", "error", err)
logger.Warn("Failed to ping MSSQL database: %v", err)
}
continue
}
+132
View File
@@ -0,0 +1,132 @@
package providers
import (
"context"
"database/sql/driver"
"fmt"
"net"
"sync/atomic"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/stdlib"
)
const (
// tcpKeepAlive is how often keepalive probes are sent on idle connections.
tcpKeepAlive = 30 * time.Second
// tcpUserTimeout bounds how long written data may stay unacknowledged before
// the kernel drops the socket. Without it a query on a silently dead peer
// waits for tcp_retries2 (about 15 minutes).
tcpUserTimeout = 30 * time.Second
// resetSessionTimeout bounds the liveness ping database/sql triggers when a
// pooled connection is reused, which otherwise runs on the request context.
resetSessionTimeout = 5 * time.Second
)
// pgConnector is a driver.Connector whose connections can be retired without
// closing the *sql.DB. Reconnecting bumps a generation; connections created
// under an older generation report themselves invalid and database/sql
// discards them and dials new ones. Every handle wrapping the *sql.DB keeps
// working across a reconnect.
type pgConnector struct {
inner atomic.Pointer[connectorState]
generation atomic.Uint64
}
type connectorState struct {
connector driver.Connector
gen uint64
}
func newPGConnector(cfg *pgx.ConnConfig) *pgConnector {
c := &pgConnector{}
c.swap(cfg)
return c
}
// swap installs a new connection config under a fresh generation.
func (c *pgConnector) swap(cfg *pgx.ConnConfig) {
gen := c.generation.Add(1)
c.inner.Store(&connectorState{connector: stdlib.GetConnector(*cfg), gen: gen})
}
func (c *pgConnector) Connect(ctx context.Context) (driver.Conn, error) {
st := c.inner.Load()
conn, err := st.connector.Connect(ctx)
if err != nil {
return nil, err
}
sc, ok := conn.(*stdlib.Conn)
if !ok {
conn.Close()
return nil, fmt.Errorf("unexpected pgx driver connection type %T", conn)
}
return &pgConn{Conn: sc, owner: c, gen: st.gen}, nil
}
func (c *pgConnector) Driver() driver.Driver {
return stdlib.GetDefaultDriver()
}
// pgConn embeds *stdlib.Conn, so every optional driver interface (context
// queries, Pinger, NamedValueChecker, ...) is promoted unchanged.
type pgConn struct {
*stdlib.Conn
owner *pgConnector
gen uint64
}
func (c *pgConn) stale() bool { return c.gen != c.owner.generation.Load() }
// IsValid implements driver.Validator: stale or closed connections are dropped
// when returned to the pool.
func (c *pgConn) IsValid() bool {
return !c.stale() && !c.Conn.Conn().IsClosed()
}
// ResetSession runs when a pooled connection is reused. It discards stale
// connections and bounds pgx's liveness ping so a dead socket fails in seconds
// rather than blocking on the caller's context.
func (c *pgConn) ResetSession(ctx context.Context) error {
if c.stale() {
return driver.ErrBadConn
}
ctx, cancel := context.WithTimeout(ctx, resetSessionTimeout)
defer cancel()
return c.Conn.ResetSession(ctx)
}
// newDialFunc returns a pgconn dial function with TCP keepalive and, where the
// platform supports it, TCP_USER_TIMEOUT.
func newDialFunc(connectTimeout time.Duration) func(ctx context.Context, network, addr string) (net.Conn, error) {
d := &net.Dialer{
Timeout: connectTimeout,
KeepAlive: tcpKeepAlive,
Control: setTCPUserTimeout(tcpUserTimeout),
}
return d.DialContext
}
// buildPGXConfig parses the DSN and applies client-side hardening: bounded
// dialing, TCP timeouts, and statement_timeout, which is set as a runtime
// parameter so it also applies to caller-supplied DSNs.
func buildPGXConfig(cfg ConnectionConfig) (*pgx.ConnConfig, error) {
dsn, err := cfg.BuildDSN()
if err != nil {
return nil, fmt.Errorf("failed to build DSN: %w", err)
}
cc, err := pgx.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("failed to parse connection config: %w", err)
}
cc.DialFunc = newDialFunc(cfg.GetConnectTimeout())
if cfg.GetQueryTimeout() > 0 {
if _, set := cc.RuntimeParams["statement_timeout"]; !set {
cc.RuntimeParams["statement_timeout"] = fmt.Sprintf("%d", cfg.GetQueryTimeout().Milliseconds())
}
}
return cc, nil
}
@@ -0,0 +1,34 @@
package providers
import (
"database/sql/driver"
"testing"
"time"
"github.com/jackc/pgx/v5"
)
func TestConnectorGenerationInvalidatesConns(t *testing.T) {
cfg, err := pgx.ParseConfig("postgres://u:p@127.0.0.1:1/db")
if err != nil {
t.Fatal(err)
}
c := newPGConnector(cfg)
conn := &pgConn{owner: c, gen: c.generation.Load()}
if conn.stale() {
t.Fatal("fresh connection reported stale")
}
c.swap(cfg)
if !conn.stale() {
t.Fatal("connection from an older generation must be stale")
}
if err := conn.ResetSession(t.Context()); err != driver.ErrBadConn {
t.Fatalf("ResetSession on stale conn = %v, want ErrBadConn", err)
}
}
func TestDialFuncHasTimeout(t *testing.T) {
if newDialFunc(2*time.Second) == nil {
t.Fatal("nil dial func")
}
}
+64 -48
View File
@@ -3,12 +3,12 @@ package providers
import (
"context"
"database/sql"
"errors"
"fmt"
"math"
"sync"
"time"
_ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver
"go.mongodb.org/mongo-driver/mongo"
"github.com/bitechdev/ResolveSpec/pkg/logger"
@@ -16,10 +16,11 @@ import (
// PostgresProvider implements Provider for PostgreSQL databases
type PostgresProvider struct {
db *sql.DB
config ConnectionConfig
listener *PostgresListener
mu sync.Mutex
db *sql.DB
connector *pgConnector
config ConnectionConfig
listener *PostgresListener
mu sync.Mutex
}
// NewPostgresProvider creates a new PostgreSQL provider
@@ -29,22 +30,24 @@ func NewPostgresProvider() *PostgresProvider {
// Connect establishes a PostgreSQL connection
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
// Build DSN
dsn, err := cfg.BuildDSN()
connCfg, err := buildPGXConfig(cfg)
if err != nil {
return fmt.Errorf("failed to build DSN: %w", err)
return err
}
// The connector and *sql.DB are created once; the pool is never closed to
// recover from errors (see Refresh).
connector := newPGConnector(connCfg)
db := sql.OpenDB(connector)
// Connect with retry logic
var db *sql.DB
var lastErr error
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
retryAttempts := 3 // Default retry attempts
retryDelay := 1 * time.Second
connected := false
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if cfg.GetEnableLogging() {
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
@@ -52,20 +55,11 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
select {
case <-time.After(delay):
case <-ctx.Done():
db.Close()
return ctx.Err()
}
}
// Open database connection
db, err = sql.Open("pgx", dsn)
if err != nil {
lastErr = err
if cfg.GetEnableLogging() {
logger.Warn("Failed to open PostgreSQL connection", "error", err)
}
continue
}
// Test the connection with context timeout
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
err = db.PingContext(connectCtx)
@@ -73,18 +67,18 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
if err != nil {
lastErr = err
db.Close()
if cfg.GetEnableLogging() {
logger.Warn("Failed to ping PostgreSQL database", "error", err)
logger.Warn("Failed to ping PostgreSQL database: %v", err)
}
continue
}
// Connection successful
connected = true
break
}
if err != nil {
if !connected {
db.Close()
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
}
@@ -103,6 +97,7 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
}
p.db = db
p.connector = connector
p.config = cfg
if cfg.GetEnableLogging() {
@@ -112,34 +107,55 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
return nil
}
// Close closes the PostgreSQL connection
func (p *PostgresProvider) Close() error {
// Close listener if it exists
p.mu.Lock()
if p.listener != nil {
if err := p.listener.Close(); err != nil {
p.mu.Unlock()
return fmt.Errorf("failed to close listener: %w", err)
}
p.listener = nil
// Refresh retires every pooled connection and dials fresh ones on demand,
// without closing the *sql.DB. Handles already handed out keep working:
// connections in use finish their current query and are then discarded.
func (p *PostgresProvider) Refresh(ctx context.Context) error {
if p.db == nil || p.connector == nil {
return fmt.Errorf("database connection is not initialized")
}
connCfg, err := buildPGXConfig(p.config)
if err != nil {
return err
}
p.connector.swap(connCfg)
pingCtx, cancel := context.WithTimeout(ctx, p.config.GetConnectTimeout())
defer cancel()
if err := p.db.PingContext(pingCtx); err != nil {
return fmt.Errorf("failed to ping after refresh: %w", err)
}
return nil
}
// Close closes the PostgreSQL connection. A listener failure does not stop the
// pool from being closed.
func (p *PostgresProvider) Close() error {
var errs []error
p.mu.Lock()
listener := p.listener
p.listener = nil
p.mu.Unlock()
if p.db == nil {
return nil
if listener != nil {
if err := listener.Close(); err != nil {
errs = append(errs, fmt.Errorf("failed to close listener: %w", err))
}
}
err := p.db.Close()
if err != nil {
return fmt.Errorf("failed to close PostgreSQL connection: %w", err)
if p.db != nil {
if err := p.db.Close(); err != nil {
errs = append(errs, fmt.Errorf("failed to close PostgreSQL connection: %w", err))
} else if p.config.GetEnableLogging() {
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
}
p.db = nil
p.connector = nil
}
if p.config.GetEnableLogging() {
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
}
p.db = nil
return nil
return errors.Join(errs...)
}
// HealthCheck verifies the PostgreSQL connection is alive
+217 -159
View File
@@ -23,6 +23,10 @@ type PostgresListener struct {
// Channel subscriptions
channels map[string]NotificationHandler
mu sync.RWMutex
// connMu serialises use of the single pgx.Conn: it is not safe for
// concurrent use, so the notification wait and LISTEN/UNLISTEN/NOTIFY take
// turns. Lock order: connMu before mu.
connMu sync.Mutex
// Lifecycle management
ctx context.Context
@@ -30,6 +34,7 @@ type PostgresListener struct {
closed bool
closeMu sync.Mutex
reconnectC chan struct{}
startOnce sync.Once // background goroutines start exactly once
}
// NewPostgresListener creates a new PostgreSQL listener
@@ -44,76 +49,20 @@ func NewPostgresListener(cfg ConnectionConfig) *PostgresListener {
}
}
// Connect establishes a dedicated connection for listening
// Connect establishes a dedicated connection for listening and starts the
// background loops (once per listener).
func (l *PostgresListener) Connect(ctx context.Context) error {
dsn, err := l.config.BuildDSN()
conn, err := l.dial(ctx)
if err != nil {
return fmt.Errorf("failed to build DSN: %w", err)
return err
}
// Parse connection config
connConfig, err := pgx.ParseConfig(dsn)
if err != nil {
return fmt.Errorf("failed to parse connection config: %w", err)
}
l.swapConn(conn)
// Connect with retry logic
var conn *pgx.Conn
var lastErr error
retryAttempts := 3
retryDelay := 1 * time.Second
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
if l.config.GetEnableLogging() {
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
select {
case <-time.After(delay):
case <-ctx.Done():
return ctx.Err()
}
}
conn, err = pgx.ConnectConfig(ctx, connConfig)
if err != nil {
lastErr = err
if l.config.GetEnableLogging() {
logger.Warn("Failed to connect PostgreSQL listener", "error", err)
}
continue
}
// Test the connection
if err = conn.Ping(ctx); err != nil {
lastErr = err
conn.Close(ctx)
if l.config.GetEnableLogging() {
logger.Warn("Failed to ping PostgreSQL listener", "error", err)
}
continue
}
// Connection successful
break
}
if err != nil {
return fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
}
l.mu.Lock()
l.conn = conn
l.mu.Unlock()
// Start notification handler
go l.handleNotifications()
// Start reconnection handler
go l.handleReconnection()
l.startOnce.Do(func() {
go l.handleNotifications()
go l.handleReconnection()
})
if l.config.GetEnableLogging() {
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
@@ -122,30 +71,105 @@ func (l *PostgresListener) Connect(ctx context.Context) error {
return nil
}
// dial opens and verifies a new dedicated connection, with retries.
func (l *PostgresListener) dial(ctx context.Context) (*pgx.Conn, error) {
connConfig, err := buildPGXConfig(l.config)
if err != nil {
return nil, err
}
var lastErr error
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(l.config)
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if l.config.GetEnableLogging() {
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
select {
case <-time.After(delay):
case <-ctx.Done():
return nil, ctx.Err()
}
}
conn, err := pgx.ConnectConfig(ctx, connConfig)
if err != nil {
lastErr = err
if l.config.GetEnableLogging() {
logger.Warn("Failed to connect PostgreSQL listener: %v", err)
}
continue
}
// Test the connection
if err = conn.Ping(ctx); err != nil {
lastErr = err
closeConnBounded(conn)
if l.config.GetEnableLogging() {
logger.Warn("Failed to ping PostgreSQL listener: %v", err)
}
continue
}
return conn, nil
}
return nil, fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
}
// closeConnBounded closes a pgx connection without ever waiting on a dead socket.
func closeConnBounded(conn *pgx.Conn) error {
ctx, cancel := context.WithTimeout(context.Background(), listenerCloseTimeout)
defer cancel()
return conn.Close(ctx)
}
const (
listenerCloseTimeout = 2 * time.Second
notificationPollInterval = 500 * time.Millisecond
)
// currentConn returns the live connection, or an error if the listener is
// closed or not yet connected.
func (l *PostgresListener) currentConn() (*pgx.Conn, error) {
l.closeMu.Lock()
closed := l.closed
l.closeMu.Unlock()
if closed {
return nil, fmt.Errorf("listener is closed")
}
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
return nil, fmt.Errorf("listener connection is not initialized")
}
return conn, nil
}
// Listen subscribes to a PostgreSQL notification channel
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
// Take the connection between notification waits (each wait is short).
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.Lock()
defer l.mu.Unlock()
if l.conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute LISTEN command
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize()))
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
}
// Store the handler
l.mu.Lock()
l.channels[channel] = handler
l.mu.Unlock()
if l.config.GetEnableLogging() {
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
@@ -156,28 +180,21 @@ func (l *PostgresListener) Listen(channel string, handler NotificationHandler) e
// Unlisten unsubscribes from a PostgreSQL notification channel
func (l *PostgresListener) Unlisten(channel string) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.Lock()
defer l.mu.Unlock()
if l.conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute UNLISTEN command
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
}
// Remove the handler
l.mu.Lock()
delete(l.channels, channel)
l.mu.Unlock()
if l.config.GetEnableLogging() {
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
@@ -188,31 +205,24 @@ func (l *PostgresListener) Unlisten(channel string) error {
// Notify sends a notification to a PostgreSQL channel
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute NOTIFY command
_, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload)
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload); err != nil {
return fmt.Errorf("failed to notify channel %s: %w", channel, err)
}
return nil
}
// Close closes the listener and all subscriptions
// Close closes the listener and all subscriptions. Closing the connection drops
// every subscription server-side, so no UNLISTEN round trips are needed, and
// the close itself is bounded so a dead socket cannot hang the caller.
func (l *PostgresListener) Close() error {
l.closeMu.Lock()
if l.closed {
@@ -225,27 +235,26 @@ func (l *PostgresListener) Close() error {
// Cancel context to stop background goroutines
l.cancel()
// The cancelled ctx makes the notification wait return promptly, releasing
// connMu; closing the conn while it is being read would race inside pgx.
l.connMu.Lock()
l.mu.Lock()
defer l.mu.Unlock()
conn := l.conn
l.conn = nil
l.channels = make(map[string]NotificationHandler)
l.mu.Unlock()
if l.conn == nil {
if conn == nil {
l.connMu.Unlock()
return nil
}
// Unlisten from all channels
for channel := range l.channels {
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
}
// Close connection
err := l.conn.Close(context.Background())
err := closeConnBounded(conn)
l.connMu.Unlock()
if err != nil {
return fmt.Errorf("failed to close listener connection: %w", err)
}
l.conn = nil
l.channels = make(map[string]NotificationHandler)
if l.config.GetEnableLogging() {
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
}
@@ -262,20 +271,26 @@ func (l *PostgresListener) handleNotifications() {
default:
}
l.connMu.Lock()
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
l.connMu.Unlock()
// Connection not available, wait for reconnection
time.Sleep(100 * time.Millisecond)
if !l.sleep(100 * time.Millisecond) {
return
}
continue
}
// Wait for notification with timeout
ctx, cancel := context.WithTimeout(l.ctx, 5*time.Second)
// Wait for a notification with a short timeout, so Listen/Unlisten/Notify
// waiting on connMu are served promptly.
ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval)
notification, err := conn.WaitForNotification(ctx)
cancel()
l.connMu.Unlock()
if err != nil {
// Check if context was cancelled
@@ -291,13 +306,15 @@ func (l *PostgresListener) handleNotifications() {
// Connection error, trigger reconnection
if l.config.GetEnableLogging() {
logger.Warn("Notification error, triggering reconnection", "error", err)
logger.Warn("Notification error, triggering reconnection: %v", err)
}
select {
case l.reconnectC <- struct{}{}:
default:
}
time.Sleep(1 * time.Second)
if !l.sleep(1 * time.Second) {
return
}
continue
}
@@ -322,7 +339,22 @@ func (l *PostgresListener) handleNotifications() {
}
}
// handleReconnection manages automatic reconnection
// sleep waits for d or until the listener is closed; it reports whether the
// listener is still running.
func (l *PostgresListener) sleep(d time.Duration) bool {
t := time.NewTimer(d)
defer t.Stop()
select {
case <-t.C:
return true
case <-l.ctx.Done():
return false
}
}
// handleReconnection manages automatic reconnection. It runs as a single
// goroutine and dials replacement connections directly rather than through the
// public Connect, so no extra loops are started.
func (l *PostgresListener) handleReconnection() {
for {
select {
@@ -333,31 +365,21 @@ func (l *PostgresListener) handleReconnection() {
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
}
// Close existing connection
l.mu.Lock()
if l.conn != nil {
l.conn.Close(context.Background())
l.conn = nil
}
// Save current subscriptions
channels := make(map[string]NotificationHandler)
for ch, handler := range l.channels {
channels[ch] = handler
}
l.mu.Unlock()
// Attempt reconnection
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
err := l.Connect(ctx)
ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second)
err := l.reconnect(ctx)
cancel()
if err != nil {
if l.ctx.Err() != nil {
return
}
if l.config.GetEnableLogging() {
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
}
// Retry after delay
time.Sleep(5 * time.Second)
if !l.sleep(5 * time.Second) {
return
}
select {
case l.reconnectC <- struct{}{}:
default:
@@ -365,15 +387,6 @@ func (l *PostgresListener) handleReconnection() {
continue
}
// Resubscribe to all channels
for channel, handler := range channels {
if err := l.Listen(channel, handler); err != nil {
if l.config.GetEnableLogging() {
logger.Error("Failed to resubscribe to channel: name=%s, channel=%s, error=%v", l.config.GetName(), channel, err)
}
}
}
if l.config.GetEnableLogging() {
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
}
@@ -381,6 +394,51 @@ func (l *PostgresListener) handleReconnection() {
}
}
// reconnect replaces the connection and resubscribes every channel on the new
// connection before publishing it, so the notification loop never touches a
// half-initialised conn.
func (l *PostgresListener) reconnect(ctx context.Context) error {
conn, err := l.dial(ctx)
if err != nil {
return err
}
l.mu.RLock()
channels := make([]string, 0, len(l.channels))
for ch := range l.channels {
channels = append(channels, ch)
}
l.mu.RUnlock()
for _, ch := range channels {
if _, err := conn.Exec(ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{ch}.Sanitize())); err != nil {
closeConnBounded(conn)
return fmt.Errorf("failed to resubscribe to channel %s: %w", ch, err)
}
}
if l.ctx.Err() != nil {
closeConnBounded(conn)
return l.ctx.Err()
}
l.swapConn(conn)
return nil
}
// swapConn installs conn and closes the previous one. The old connection is
// closed under connMu so it is never closed while another goroutine is using it.
func (l *PostgresListener) swapConn(conn *pgx.Conn) {
l.connMu.Lock()
l.mu.Lock()
old := l.conn
l.conn = conn
l.mu.Unlock()
if old != nil {
closeConnBounded(old)
}
l.connMu.Unlock()
}
// IsConnected returns true if the listener is connected
func (l *PostgresListener) IsConnected() bool {
l.mu.RLock()
+25 -6
View File
@@ -4,17 +4,11 @@ import (
"context"
"database/sql"
"errors"
"strings"
"time"
"go.mongodb.org/mongo-driver/mongo"
)
// isDBClosed reports whether err indicates the *sql.DB has been closed.
func isDBClosed(err error) bool {
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
}
// Common errors
var (
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
@@ -63,6 +57,31 @@ type ConnectionConfig interface {
GetConnMaxLifetime() *time.Duration
GetConnMaxIdleTime() *time.Duration
GetReadPreference() string
GetRetryAttempts() int
GetRetryDelay() time.Duration
GetRetryMaxDelay() time.Duration
}
// retryPolicy returns the configured retry settings, falling back to defaults.
func retryPolicy(cfg ConnectionConfig) (attempts int, delay, maxDelay time.Duration) {
attempts, delay, maxDelay = cfg.GetRetryAttempts(), cfg.GetRetryDelay(), cfg.GetRetryMaxDelay()
if attempts <= 0 {
attempts = 3
}
if delay <= 0 {
delay = time.Second
}
if maxDelay <= 0 {
maxDelay = 10 * time.Second
}
return
}
// Refresher is implemented by providers that can retire their pooled
// connections and dial fresh ones without closing the shared *sql.DB, so
// handles already handed out keep working.
type Refresher interface {
Refresh(ctx context.Context) error
}
// Provider creates and manages the underlying database connection
+44 -68
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"fmt"
"strings"
"sync"
"time"
@@ -15,10 +16,9 @@ import (
// SQLiteProvider implements Provider for SQLite databases
type SQLiteProvider struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
config ConnectionConfig
db *sql.DB
dbMu sync.RWMutex
config ConnectionConfig
}
// NewSQLiteProvider creates a new SQLite provider
@@ -26,6 +26,22 @@ func NewSQLiteProvider() *SQLiteProvider {
return &SQLiteProvider{}
}
// isMemoryDSN reports whether the SQLite DSN refers to a private in-memory
// database (each pooled connection would get its own empty database).
func isMemoryDSN(dsn string) bool {
path := dsn
if i := strings.IndexByte(path, '?'); i >= 0 {
path = path[:i]
}
if path == ":memory:" || path == "" {
return true
}
if strings.Contains(dsn, "mode=memory") && !strings.Contains(dsn, "cache=shared") {
return true
}
return path == "file::memory:" && !strings.Contains(dsn, "cache=shared")
}
// Connect establishes a SQLite connection
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
// Build DSN
@@ -50,48 +66,35 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro
return fmt.Errorf("failed to ping SQLite database: %w", err)
}
// Configure connection pool
// Note: SQLite works best with MaxOpenConns=1 for write operations
// but can handle multiple readers
if cfg.GetMaxOpenConns() != nil {
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
} else {
// Default to 1 for SQLite to avoid "database is locked" errors
if isMemoryDSN(dsn) {
// A private in-memory database exists per connection and disappears when
// that connection closes, so pin the pool to one connection that is
// never recycled.
db.SetMaxOpenConns(1)
}
if cfg.GetMaxIdleConns() != nil {
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
}
if cfg.GetConnMaxLifetime() != nil {
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
}
if cfg.GetConnMaxIdleTime() != nil {
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
}
// Enable WAL mode for better concurrent access
_, err = db.ExecContext(ctx, "PRAGMA journal_mode=WAL")
if err != nil {
if cfg.GetEnableLogging() {
logger.Warn("Failed to enable WAL mode for SQLite", "error", err)
db.SetMaxIdleConns(1)
db.SetConnMaxLifetime(0)
db.SetConnMaxIdleTime(0)
} else {
// SQLite works best with few writers; default to 1 unless configured.
if cfg.GetMaxOpenConns() != nil {
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
} else {
db.SetMaxOpenConns(1)
}
// Don't fail connection if WAL mode cannot be enabled
}
// Set busy timeout to handle locked database (minimum 2 minutes = 120000ms)
busyTimeout := cfg.GetQueryTimeout().Milliseconds()
if busyTimeout < 120000 {
busyTimeout = 120000 // Enforce minimum of 2 minutes
}
_, err = db.ExecContext(ctx, fmt.Sprintf("PRAGMA busy_timeout=%d", busyTimeout))
if err != nil {
if cfg.GetEnableLogging() {
logger.Warn("Failed to set busy timeout for SQLite", "error", err)
if cfg.GetMaxIdleConns() != nil {
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
}
if cfg.GetConnMaxLifetime() != nil {
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
}
if cfg.GetConnMaxIdleTime() != nil {
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
}
}
p.dbMu.Lock()
p.db = db
p.dbMu.Unlock()
p.config = cfg
if cfg.GetEnableLogging() {
@@ -132,14 +135,7 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
// Execute a simple query to verify the database is accessible
var result int
run := func() error { return p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result) }
err := run()
if isDBClosed(err) {
if reconnErr := p.reconnectDB(); reconnErr == nil {
err = run()
}
}
if err != nil {
if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil {
return fmt.Errorf("health check failed: %w", err)
}
@@ -150,32 +146,12 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
return nil
}
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
func (p *SQLiteProvider) WithDBFactory(factory func() (*sql.DB, error)) *SQLiteProvider {
p.dbFactory = factory
return p
}
func (p *SQLiteProvider) getDB() *sql.DB {
p.dbMu.RLock()
defer p.dbMu.RUnlock()
return p.db
}
func (p *SQLiteProvider) reconnectDB() error {
if p.dbFactory == nil {
return fmt.Errorf("no db factory configured for reconnect")
}
newDB, err := p.dbFactory()
if err != nil {
return err
}
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
return nil
}
// GetNative returns the native *sql.DB connection
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
if p.db == nil {
@@ -0,0 +1,24 @@
//go:build linux
package providers
import (
"syscall"
"time"
"golang.org/x/sys/unix"
)
// setTCPUserTimeout returns a net.Dialer Control func setting TCP_USER_TIMEOUT.
func setTCPUserTimeout(d time.Duration) func(network, address string, c syscall.RawConn) error {
return func(network, address string, c syscall.RawConn) error {
var sockErr error
err := c.Control(func(fd uintptr) {
sockErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_TCP, unix.TCP_USER_TIMEOUT, int(d.Milliseconds()))
})
if err != nil {
return err
}
return sockErr
}
}
@@ -0,0 +1,14 @@
//go:build !linux
package providers
import (
"syscall"
"time"
)
// setTCPUserTimeout is a no-op where TCP_USER_TIMEOUT is unavailable; TCP
// keepalive still applies.
func setTCPUserTimeout(time.Duration) func(network, address string, c syscall.RawConn) error {
return nil
}
+69
View File
@@ -0,0 +1,69 @@
package dbmanager
import (
"context"
"os"
"os/exec"
"testing"
"time"
)
func TestLiveServerRestart(t *testing.T) {
dir := os.Getenv("PG_RESTART_DIR")
if dir == "" {
t.Skip("PG_RESTART_DIR not set")
}
mgr, _ := NewManager(ManagerConfig{
DefaultConnection: "pg",
Connections: map[string]ConnectionConfig{"pg": {
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", ConnectTimeout: 2 * time.Second,
}},
HealthCheckInterval: 500 * time.Millisecond,
})
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
db, _ := conn.Bun()
gdb, _ := conn.GORM()
query := func() error { var n int; return db.DB.QueryRow("select 1").Scan(&n) }
if err := query(); err != nil {
t.Fatal(err)
}
run := func(args ...string) {
// Output must not be piped: the daemonised server would hold the pipe open.
if err := exec.Command("pg_ctl", append([]string{"-D", dir}, args...)...).Run(); err != nil {
t.Fatalf("pg_ctl %v: %v", args, err)
}
}
run("-m", "immediate", "-w", "stop") // crash-style shutdown
time.Sleep(1500 * time.Millisecond) // health checks fail meanwhile
if err := query(); err == nil {
t.Fatal("expected failure while server is down")
}
run("-l", dir+"/restart.log", "-o", "-p 54329 -k "+dir+" -c listen_addresses=127.0.0.1", "-w", "start")
var last error
for i := 0; i < 20; i++ {
if last = query(); last == nil {
break
}
t.Logf("attempt %d after restart: %v", i, last)
time.Sleep(200 * time.Millisecond)
}
if last != nil {
t.Fatalf("held bun handle never recovered: %v", last)
}
var n int
if err := gdb.Raw("select 1").Scan(&n).Error; err != nil {
t.Fatalf("held gorm handle: %v", err)
}
time.Sleep(time.Second)
if err := conn.HealthCheck(ctx); err != nil {
t.Fatalf("health check after restart: %v", err)
}
}