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
+22 -1
View File
@@ -576,11 +576,32 @@ Centralized management of multiple database connections with support for Postgre
- Multiple named database connections
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
- Automatic SQLite schema translation (`schema.table` → `schema_table`)
- Health checks with auto-reconnect
- Background health checks (report status; they never close the pool)
- Prometheus metrics for monitoring
- Configuration-driven via YAML
- Per-connection statistics and management
**How to use it correctly**:
```go
mgr, err := dbmanager.NewManager(cfg) // or dbmanager.SetupManager(cfg) + GetInstance()
if err != nil { /* handle */ }
if err := mgr.Connect(ctx); err != nil { /* handle */ } // SetupManager does NOT connect
defer mgr.Close() // once, at shutdown
conn, _ := mgr.GetDefault()
db, _ := conn.Bun() // or conn.GORM() / conn.Native() / conn.Database()
handler := restheadspec.NewHandlerWithBun(db)
```
- **Fetch a handle once and keep it.** `Bun()`, `GORM()`, `Native()` and `Database()` return handles over one long-lived `*sql.DB`. You do not need to re-fetch them per request, and they stay valid for the life of the connection.
- **Never close a handle yourself.** Closing a `*bun.DB`, `*gorm.DB` or the `*sql.DB` closes the shared pool for everyone. Only `mgr.Close()` (at shutdown) should close it. After `Close`, the handles are dead.
- **Don't reconnect to recover from errors.** `database/sql` already discards bad connections and dials new ones. The manager does not close the pool on errors or failed health checks. `conn.Reconnect(ctx)` is for explicit operator use only (for example after rotating credentials): on PostgreSQL it retires pooled connections without closing the pool, so held handles keep working. Other databases close and reopen the pool, which invalidates handles you already hold.
- **Bring your own `*sql.DB`.** `dbmanager.NewConnectionFromDB(name, type, db)` wraps a pool you opened. The manager never closes it (`Close` only logs a warning); you own it and must close it.
- **Set deadlines on request contexts.** `query_timeout` is applied to PostgreSQL as `statement_timeout` (server side) and TCP timeouts detect dead sockets, but pass a context with a deadline to your queries so callers fail fast.
- **Pool tuning.** Keep `conn_max_idle_time` below the shortest idle timeout of any NAT, load balancer or pgbouncer between you and the database (typically 60-240s). SQLite `:memory:` is pinned to a single connection.
- **Health checks** run every `health_check_interval` (default 15s; a negative value disables them) and publish Prometheus metrics. `enable_auto_reconnect` is deprecated and ignored.
For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md).
#### Cache
+630
View File
@@ -0,0 +1,630 @@
# Audit: `pkg/dbmanager`
| | |
|---|---|
| **Package** | `github.com/bitechdev/ResolveSpec/pkg/dbmanager` (+ `providers/`) |
| **Files** | `config.go` (489), `connection.go` (722), `manager.go` (401), `metrics.go` (136), `errors.go` (82), `factory.go` (67), `providers/postgres.go` (231), `providers/postgres_listener.go` (401), `providers/sqlite.go` (216), `providers/mongodb.go` (214), `providers/mssql.go` (184), `providers/existing_db.go` (111), `providers/provider.go` (89); tests `factory_test.go` (369), `manager_test.go` (290), `providers/existing_db_test.go` (194), `providers/postgres_listener_example_test.go` (229) |
| **Audit date** | 2026-09-30 |
| **Axes** | thread locking/waiting, slowness, security, panic handling & logging |
| **Threat model** | hostile internet client; request bodies, headers, query params, schema/table/column names all attacker-controlled |
| **Depth** | deep (hot package; every request's DB handle comes from here). Several findings were checked with a throw-away probe test against SQLite, and the probe was deleted afterwards |
## Summary
`pkg/dbmanager` owns every database pool in the process. It wraps a
`*sql.DB` (or a `mongo.Client`) in a `sqlConnection` and hands out lazily-built
`*bun.DB`, `*gorm.DB`, raw `*sql.DB` and `common.Database` adapters over it. A
background health checker pings each connection every 15 s, and it can
**reconnect**, which closes the pool and opens a new one.
This audit was started to answer one question: **"why does a database
connection that has been idle for a while become unusable?"** Several defects
in this package combine to give exactly that symptom. They are findings 1–5,
and the [Idle-connection failure chain](#idle-connection-failure-chain) section
below puts them together.
The root design problem is that **`Reconnect` destroys the shared `*sql.DB`**.
`*sql.DB` is already a self-healing pool: it throws away bad connections and
dials new ones. So "reconnecting" a pool is almost never needed, and here it
has a large blast radius. Every `*bun.DB`, `*gorm.DB` and `*sql.DB` handed out
before the reconnect now points at a closed pool, and it stays closed. Only the
`common.Database` adapters carry a factory that can re-fetch a handle, and even
they only use it on a subset of code paths (see `common.audit.md` finding 5).
Those adapter factories also *trigger* `Reconnect` themselves, so one stale
handle closes the pool for everyone else. `Reconnect` isn't atomic, so
concurrent callers turn this into a storm.
The other major theme is **missing client-side deadlines**. `QueryTimeout` is
only ever sent to the server as `statement_timeout`, which does nothing when
the TCP peer has vanished. No `context.WithTimeout` is applied to request
queries, and pgx's dialer sets no `TCP_USER_TIMEOUT`. So the first query on a
pooled connection whose peer silently disappeared (NAT/firewall idle drop,
failover, a pgbouncer restart) can block for minutes. One Close path does this
while holding the connection's write lock, which stalls every request.
## Findings
| # | Severity | Axis | Finding | Status |
|---|---|---|---|---|
| 1 | **Critical** | locking / availability | `Reconnect` closes the shared `*sql.DB`, so every `*bun.DB` / `*gorm.DB` / `*sql.DB` handed out earlier is permanently dead ("sql: database is closed") | Fixed |
| 2 | **High** | locking | Adapter reconnect factories call `Reconnect` on the *shared* connection, and `Reconnect` is not atomic, so one stale handle starts a reconnect storm that repeatedly closes the pool under in-flight requests | Fixed |
| 3 | **High** | slowness / locking | `sqlConnection.HealthCheck` holds the write lock across a network ping for up to 5 s; every `Bun()`/`GORM()`/`Native()`/`Database()`/`Stats()` call blocks for that time | Fixed |
| 4 | **High** | slowness | No client-side query deadline and no `TCP_USER_TIMEOUT`: a query on a silently-dead idle socket blocks for minutes (up to about 15 min); `QueryTimeout` is server-side only, and is forced to at least 2 min | Fixed |
| 5 | **High** | locking / slowness | `PostgresListener.Close` runs `UNLISTEN` with `context.Background()` while `sqlConnection.mu` (write), `PostgresProvider.mu` and `listener.mu` are all held; on a dead socket this freezes every request for minutes | Fixed |
| 6 | **High** | locking / leak | `PostgresListener.Connect` starts a new goroutine pair on every (re)connect; the old pair keeps running, so two loops call `WaitForNotification` on one `pgx.Conn` concurrently, which triggers more reconnects | Fixed |
| 7 | **High** | panic handling | `Connect → Close → Connect → Close` panics with "close of closed channel"; after the first cycle the health checker also exits immediately and silently | Fixed |
| 8 | **Medium** | availability | SQLite: `:memory:` with a 25-connection pool gives every connection its own empty database, and `ConnMaxIdleTime` then silently discards data; `busy_timeout` / WAL pragmas are applied to only one pooled connection | Fixed |
| 9 | **Medium** | availability | Partial failure in `sqlConnection.Close` leaves `connected=true` over a closed pool; partial failure in `Manager.Connect` leaks the connections already opened | Fixed |
| 10 | **Medium** | security | DSN builders concatenate unescaped credentials (postgres key=value, mssql/mongo URLs); `sslmode` defaults to `disable` | Fixed |
| 11 | **Medium** | config | Several config knobs are ignored or impossible to turn off: `EnableAutoReconnect`, `HealthCheckInterval`, `RetryAttempts`/`RetryDelay`/`RetryMaxDelay`, SQLite `_timeout`, and `statement_timeout` when a DSN is given | Fixed |
| 12 | **Medium** | locking | `Manager.Connect` holds `m.mu` across every network dial (up to 3 retries × `ConnectTimeout` per connection) | Fixed |
| 13 | **Low** | observability | `PublishMetrics` / `RecordReconnectAttempt` are never called, so all dbmanager metrics are permanently zero; `*_total` metrics are gauges | Fixed |
| 14 | **Low** | correctness | `Bun()`/`GORM()` do not check `connected`; `getNativeAdapter` uses `PgSQLAdapter` for SQLite and MSSQL; `ExistingDBProvider` applies no pool settings and closes the caller's DB | Fixed (partly, see notes) |
| 15 | **Low** | logging | `Close` / `performHealthCheck` pass key-value pairs to the printf-style logger, which produces `%!(EXTRA ...)` output; `ResetInstance` discards the close error | Fixed |
## Remediation status
Implemented 2026-09-30. `go build ./...` and `go test -race ./pkg/dbmanager/...`
pass. The Postgres behaviour was also verified against a live server (tests are
skipped unless `PG_LIVE=1` / `PG_RESTART_DIR` is set).
**Design decisions taken**
- No automatic reconnect. Adapter factories and the health checker never close
the pool; they only re-fetch the current handle. `*sql.DB` replaces bad
connections itself. `EnableAutoReconnect` is deprecated and ignored.
- `Reconnect` is atomic (one critical section) and operator-only. On PostgreSQL
it goes through a custom `driver.Connector` (`providers/pgconnector.go`): it
bumps a generation, stale pooled connections are discarded, and the `*sql.DB`
is never closed, so held Bun/GORM/`*sql.DB` handles keep working. Other
providers still close and reopen.
- Client-side deadlines are applied at the driver level rather than in the
adapters (a `context.WithTimeout` around a query is cancelled before the
caller has read the rows).
**Per finding**
1. Fixed. Postgres refresh keeps the pool; explicit `Reconnect` on other
providers still invalidates handles (documented in the README).
2. Fixed. Adapter factories no longer call `Reconnect`; `Reconnect` is a single
critical section under `lifecycleMu` + `mu`.
3. Fixed. The ping runs without `mu`; `lifecycleMu` (read) only keeps
`Close`/`Reconnect` from tearing the provider down mid-ping. Same for Mongo.
4. Fixed. TCP keepalive and `TCP_USER_TIMEOUT` (30 s, Linux) via `DialFunc`;
the reuse-time liveness ping is capped at 5 s; `statement_timeout` is set as
a runtime parameter so it also applies to a supplied DSN; the 2-minute floor
on `QueryTimeout` is removed. `SetConnMaxIdleTime` tuning remains a
configuration matter (documented in the README).
5. Fixed. Listener `Close` sends no `UNLISTEN`, closes with a 2 s bound, and
holds no lock across network I/O.
6. Fixed. Background goroutines start once (`sync.Once`); reconnect dials a
replacement, re-`LISTEN`s, then swaps it in; sleeps honour `ctx.Done()`.
Additionally, all use of the single `pgx.Conn` is serialised (`connMu`, 500 ms
notification poll), fixing "conn busy" from `Listen`/`Unlisten`/`Notify`, and
old connections are closed under `connMu` (a race found by the live test).
7. Fixed. Stop channel is created per start, guarded by `healthMu`; `Close` is
idempotent; `Connect` is idempotent.
8. Fixed. `:memory:` is pinned to one connection with no idle/lifetime limits;
`busy_timeout`/WAL are `_pragma` DSN parameters; `_timeout` and the dead
reconnect code are removed.
9. Fixed. `Close` always marks disconnected and returns joined errors;
`PostgresProvider.Close` closes the pool even if the listener fails;
`Manager.Connect` closes connections it opened when a later one fails.
10. Fixed. Postgres, MSSQL and Mongo DSNs are built as escaped URLs; default
`sslmode` is now `prefer` (was `disable`).
11. Fixed. Retry settings reach every provider; a negative
`HealthCheckInterval` disables the health checker; `EnableAutoReconnect`
deprecated; `statement_timeout` applies with a supplied DSN.
12. Fixed. `Manager.Connect` dials outside `m.mu` and publishes results under it.
13. Fixed. `PublishMetrics` runs on each health-check tick, `Reconnect` records
`RecordReconnectAttempt`, and the wait/closed metrics are true counters
(delta-tracked).
14. Partly fixed. `Bun()`/`GORM()` check `connected`; Mongo no longer maps
`MaxIdleConns` to `MinPoolSize`. `ExistingDBProvider`: `Close` is now a no-op
that logs a warning (the caller owns the `*sql.DB`; the connection's `Close`
also skips `bun.DB.Close`), and `Reconnect` only pings. Pool settings are
still not applied to a caller-owned pool. The `getNativeAdapter` claim was
stale: the adapter already receives the driver name; the three duplicate
cases were merged. Mongo `Stats()` is still empty.
15. Fixed. Printf-style logger calls corrected; `ResetInstance` logs the close
error. Unscrubbed driver errors in Sentry (X8) are not addressed here.
**Behaviour changes**
- Removed tests that closed the pool from outside and expected an adapter to
swap in a new one (three adapter tests, and the health-check reconnect test,
now asserting it never reconnects).
- `sslmode` default `prefer`; `NewConnectionFromDB` connections are no longer
closed by the manager.
**Regression tests added:** `lifecycle_test.go` (double Connect/Close cycle,
idempotent Connect, concurrent Reconnect, adapter factory leaves pool open,
accessors not blocked by health check, Close marks disconnected, existing-DB
Reconnect/Close leave the caller's pool open), `config_dsn_test.go`,
`providers/pgconnector_test.go`, `pg_live_test.go` (refresh keeps handles,
listener Listen/Notify) and `restart_live_test.go` (server crash and restart).
---
## Idle-connection failure chain
This is how findings 1–5 combine into "the connection sat idle and then could
not be used":
1. The app is idle. A NAT, firewall, load balancer or pgbouncer silently drops
the idle TCP flows. No FIN or RST reaches the process.
2. The next request takes a pooled connection. pgx's `ResetSession` pings it
because it has been idle for more than 1 s, and that ping uses the request
ctx, **which has no deadline** (finding 4). The write goes into the kernel
buffer and the read blocks until TCP retransmission gives up, which can
take minutes.
Meanwhile the health checker's 5 s ping times out and holds `c.mu`
**exclusively** for the whole time (finding 3), so every request trying to
get a handle queues behind it.
3. Eventually something returns "sql: database is closed" or
`ErrConnectionClosed`. That can be an adapter that hit a closed pool, or a
partial `Close` (finding 9). An adapter's `dbFactory` or the health checker
then calls `Reconnect` (finding 2).
4. `Reconnect` closes the `*sql.DB` (finding 1). If the Postgres listener has
subscriptions, `Close` first sends `UNLISTEN` on its own dead socket with no
deadline, still holding the write lock (finding 5), which freezes the
process again.
5. When the reconnect completes, every handle captured before it is
permanently broken. That includes the `*gorm.DB` given to
`resolvespec.NewHandlerWithGORM` in `cmd/testserver/main.go:142,56`, any
`*bun.DB` passed to `NewHandlerWithBun`, and every Bun `NewSelect`/`NewInsert`
path. **From this point on, every request that goes through those handles
fails until the process is restarted.** Concurrent failures run their own
`Reconnect`s, and each one closes the pool the previous one just opened
(finding 2).
### Fix order for this symptom
1. **Stop closing the pool to recover from connection errors.** Remove
`WithDBFactory(c.reopen*ForAdapter)` → `Reconnect`, and remove the
health-check → `Reconnect` path for SQL providers. `*sql.DB` already discards
bad connections (`driver.ErrBadConn`, `ResetSession`,
`SetConnMaxIdleTime`/`SetConnMaxLifetime`). Keep `Reconnect` for explicit
operator use only, and make it atomic (finding 2).
2. Give every request a deadline. Wrap the request ctx in
`context.WithTimeout(ctx, QueryTimeout)` in the adapters, or at the handler
boundary.
3. Set `SetConnMaxIdleTime` **below** the shortest idle timeout of any
middlebox (typically 60–240 s for cloud NATs and LBs) so idle connections
are recycled before they can be dropped silently. Also set TCP keepalive and
`TCP_USER_TIMEOUT` through a custom `pgconn.Config.DialFunc`.
4. Ping without the write lock (finding 3), and give the listener's `Close`
bounded ctxs (finding 5).
---
### 1. Critical — `Reconnect` kills every previously issued handle
`connection.go:129-160` (`Close`) and `connection.go:187-192` (`Reconnect`):
```go
func (c *sqlConnection) Close() error {
c.mu.Lock()
...
if c.bunDB != nil {
if err := c.bunDB.Close(); err != nil { // closes the shared *sql.DB
...
if err := c.provider.Close(); err != nil { // closes it again (idempotent)
...
c.nativeDB = nil
c.bunDB = nil
c.gormDB = nil
c.bunAdapter = nil
...
}
func (c *sqlConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
return err
}
return c.Connect(ctx)
}
```
`Bun()`, `GORM()` and `Native()` return the handle itself, and callers keep
it: every spec package has a `NewHandlerWithGORM(*gorm.DB)` /
`NewHandlerWithBun(*bun.DB)` constructor, and `cmd/testserver/main.go:142` does
exactly this. After `Reconnect`, the cached fields are nilled, a new pool is
built, and the handles the callers hold point at a `*sql.DB` whose `closed`
flag is set forever.
Verified with a probe: I obtained `conn.GORM()`, called `conn.Reconnect(ctx)`,
then ran a query through the old handle. It returned
`sql: database is closed`, and a fresh `conn.GORM()` worked.
The comment in `manager.go:371-374` shows the authors already knew about this
("forcing Close()+Connect() here invalidates any cached ORM wrappers and callers
that still hold the old handle"). Their mitigation was to narrow *when* the
health checker reconnects. But the adapters' own `dbFactory` still reconnects
unconditionally (finding 2).
**Failure scenario.** Any event that triggers a reconnect turns every
long-lived handler into a permanent 500 generator: a single adapter query hitting
"database is closed", or a health check returning `ErrConnectionClosed`. The
process does not recover without a restart. The same thing happens after a
normal `Manager.Close()` + `Connect()` in tests or hot-reload code.
**Recommendation.** Treat the `*sql.DB` as immortal for the life of the
`sqlConnection`. Don't close it to "reconnect": `database/sql` already replaces
broken connections. If a real re-dial is ever needed (for example after
changing credentials), build the new pool, atomically swap it in, and close the
old one only after a grace period. Give the handles returned by
`Bun()`/`GORM()`/`Native()` stable identity; one way is a `driver.Connector`
that indirects to the current pool.
---
### 2. High — Adapter-triggered, non-atomic `Reconnect` causes a reconnect storm
`connection.go:362-397` and `connection.go:431/474/517-525`:
```go
func (c *sqlConnection) reconnectForAdapter() error {
...
return c.Reconnect(ctx) // Close() then Connect(): two separate lock scopes
}
...
WithDBFactory(c.reopenBunForAdapter).
```
The adapters (`pkg/common/adapters/database/bun.go:131`, `gorm.go`,
`pgsql.go`) call `dbFactory` whenever an operation returns an error that
matches `"sql: database is closed"`. So:
- **One stale handle closes the pool for everyone.** If an adapter holds a
`*sql.DB` from before a previous reconnect, its first query fails with
"database is closed". Its factory then calls `c.Reconnect`, which closes the
*current, healthy* pool that every other adapter and request is using right
now.
- **`Reconnect` isn't atomic.** `Close` and `Connect` each take `c.mu`
separately. Under N concurrent failures, one goroutine closes and reconnects
while the others either close the brand-new pool again or fail with
`already connected`. The probe used 20 concurrent `Reconnect`s: 9 returned
"already connected", and every successful reconnect closed the pool the
previous winner had just handed to its adapter. Each of those adapters then
sees "database is closed" on its next query, and the cycle continues.
**Failure scenario.** A burst of traffic arrives just after a reconnect. Each
in-flight request whose adapter still holds the old pool triggers another
`Reconnect`, and each of those closes the pool that the previous request
reopened. The service flaps until traffic stops.
**Recommendation.** Remove the adapter → `Reconnect` path (see finding 1). If
it is kept, make `Reconnect` a single critical section, and add a generation
counter: a caller that saw generation N only reconnects if the current
generation is still N; otherwise it just re-fetches the handle.
---
### 3. High — Health check holds the write lock across a network ping
`connection.go:163-185`:
```go
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
c.mu.Lock() // exclusive
defer c.mu.Unlock()
...
if err := c.provider.HealthCheck(ctx); err != nil { // PingContext, 5 s timeout
```
Every handle accessor takes `c.mu.RLock()` first (`connection.go:199, 238, 271,
308, 335, 403, 441, 484`). While the health checker (every 15 s, `manager.go:348`)
is pinging, **every request that needs a DB handle waits**. On a healthy
network this is a few ms. On a dead idle socket it's the full 5 s ping timeout
(`providers/postgres.go:155`, inside a 10 s outer ctx).
Verified with a probe: while `c.mu` was held, `conn.Bun()` blocked for the whole
hold (200 ms in the test).
**Failure scenario.** A network blip or a silently dropped idle connection
makes the ping hang. Every 15 s the whole API pauses for up to 5 s. This fits
reports of "idle, then slow or unusable".
**Recommendation.** Snapshot `provider` under `RLock`, release the lock, ping,
then take the lock only to write `healthCheckStatus` / `lastHealthCheck`. Better
still, keep the status in an `atomic.Value`.
---
### 4. High — No client-side query deadline; `QueryTimeout` is server-side only and floored at 2 min
`config.go:223-228`:
```go
if cc.QueryTimeout == 0 {
cc.QueryTimeout = 2 * time.Minute
} else if cc.QueryTimeout < 2*time.Minute {
cc.QueryTimeout = 2 * time.Minute
}
```
`config.go:331-335` turns this into `statement_timeout=<ms>` in the Postgres DSN,
and it only does that when the DSN is *built*. A user-supplied `DSN` gets no
timeout at all. Nothing anywhere in the request path wraps ctx in a deadline.
`pkg/config`'s `query_timeout: 30s` default is silently raised to 2 min.
`statement_timeout` is enforced by the **server**, so it only helps if the
server is reachable. On a silently dropped connection:
- pgconn's default dialer is `&net.Dialer{}`: Go's default keepalive (15 s idle,
15 s interval, 9 probes) and **no `TCP_USER_TIMEOUT`**.
- Once a query has been written, there is unacknowledged data, so keepalive does
not apply. The socket then waits for TCP retransmission to give up
(`tcp_retries2`), which takes about 15 min on Linux defaults.
- `database/sql` calls pgx's `ResetSession`, which pings a connection that has
been idle for more than 1 s. That ping uses the **request ctx**, so with no
deadline it blocks just as long.
**Failure scenario.** An idle period longer than the NAT or LB idle timeout
causes the next request to hang for minutes rather than failing fast and being
retried on a fresh connection. With `MaxOpenConns` = 25, 25 such requests
exhaust the pool and every later request blocks on `db.conn()`.
**Recommendation.**
- Apply `context.WithTimeout(ctx, QueryTimeout)` in the adapters, or in a
handler middleware.
- Remove the 2-minute floor, and honour the configured value.
- Set `SetConnMaxIdleTime` below the middlebox idle timeout.
- Configure `pgconn.Config.DialFunc` with a `net.Dialer` that has `KeepAlive`
set and a `Control` func setting `TCP_USER_TIMEOUT` (for example 30 s).
- Apply `statement_timeout` through `RuntimeParams` so it also works with a
supplied DSN.
---
### 5. High — Listener `Close` does unbounded network I/O under three locks
`providers/postgres_listener.go:216-244`, reached from
`providers/postgres.go:116-126`, which is reached from `connection.go:147`:
```go
// sqlConnection.Close holds c.mu (write)
// PostgresProvider.Close holds p.mu
// PostgresListener.Close holds l.mu:
for channel := range l.channels {
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", ...))
}
err := l.conn.Close(context.Background())
```
If the listener's socket is dead, and it usually is in the situation that
triggers a reconnect, each `UNLISTEN` waits for a reply that never comes. This
is the same unbounded wait as in finding 4, and `c.mu` is held **for writing**
the whole time. Every request blocks. `bunDB` has already been closed at this
point, so there is no fallback either.
Also, if `listener.Close` returns an error, `PostgresProvider.Close` returns
early. `sqlConnection.Close` then returns with `connected=true` over a closed
pool (finding 9).
**Failure scenario.** An app with any `LISTEN` subscription hits a network
partition. The health checker or an adapter calls `Reconnect`, and the process
stops serving database requests for as long as the kernel takes to kill the
socket.
**Recommendation.** Skip `UNLISTEN` entirely, because closing the connection
drops all subscriptions server-side. Close with `context.WithTimeout(…, 2*time.Second)`.
Don't do network I/O while holding `l.mu`, and don't close the listener
inside `sqlConnection.Close`'s write lock.
---
### 6. High — Listener leaks a goroutine pair per reconnect, and they race on one `pgx.Conn`
`providers/postgres_listener.go:48-120` (Connect), `257-324` (handleNotifications),
`326-370` (handleReconnection).
`Connect()` ends by starting `go l.handleNotifications()` and
`go l.handleReconnection()`. `handleReconnection` responds to a reconnect
signal by calling `l.Connect(ctx)`, which starts **another** pair. The old pair
keeps running on the same `l.ctx`. After N reconnects there are N+1
notification loops. Each one snapshots `l.conn` and calls
`conn.WaitForNotification`. `pgx.Conn` is **not** safe for concurrent use, so
the second caller gets a "conn busy" error. That error isn't a timeout, so it
sends another reconnect signal, which adds another pair.
`handleReconnection` also waits with `time.Sleep(5 * time.Second)` instead of
selecting on `l.ctx.Done()`, so `Close` can't interrupt it. And `Listen` runs
`l.conn.Exec(LISTEN …)` while holding `l.mu`, which blocks `handleReconnection`
for as long as that Exec takes.
Once the parent `PostgresProvider` is closed (for example by any `Reconnect`,
finding 1), subscribers holding the old `*PostgresListener` get
"listener is closed" forever. Nothing re-subscribes them on the new provider.
**Failure scenario.** A flaky network causes a few listener reconnects. The
goroutine count grows without bound, notifications are delivered twice or
dropped, and CPU rises because of the busy/reconnect spiral.
**Recommendation.** Start the goroutines once, in the constructor or the first
`Connect`. Have `handleReconnection` dial a new conn without calling the public
`Connect`. Guard `WaitForNotification` so only one loop owns the conn. Replace
`time.Sleep` with `select { case <-time.After(d): case <-l.ctx.Done(): }`.
---
### 7. High — Second `Close` panics; health checker silently dead after first cycle
`manager.go:119, 313-345`:
```go
stopChan: make(chan struct{}), // created once, in the constructor
...
func (m *connectionManager) stopHealthChecker() {
if m.healthTicker != nil {
m.healthTicker.Stop()
close(m.stopChan) // never recreated
m.wg.Wait()
m.healthTicker = nil
}
}
```
After `Connect → Close`, `stopChan` is closed. A second `Connect` calls
`startHealthChecker`, which creates a new ticker and goroutine. That goroutine's
`select` sees the closed `stopChan` right away and **exits**, so health
checking is silently off. A second `Close` finds `healthTicker != nil` and
calls `close(m.stopChan)` again, which **panics**: `close of closed channel`.
`startHealthChecker` and `stopHealthChecker` also read and write `healthTicker`
without `m.mu` held (`Close` calls `stopHealthChecker` before locking), so a
concurrent `Connect`/`Close` pair is a data race.
Calling `Connect` twice without `Close` also leaks: `m.connections[name] = conn`
overwrites the previous connection without closing it.
**Failure scenario.** Anything that cycles the manager can crash the process
during shutdown: graceful restart, config hot-reload, or test suites using
`ResetInstance`.
**Recommendation.** Create `stopChan` in `startHealthChecker`. Guard both
functions with `m.mu`, or a dedicated mutex. Make `Connect` idempotent, or have
it close existing connections first.
---
### 8. Medium — SQLite: in-memory data loss and per-connection pragmas
`providers/sqlite.go:54-90`, `config.go:140-141, 202-204`:
- `ManagerConfig.ApplyDefaults` always gives `MaxOpenConns` a value (25), so the
"SQLite works best with MaxOpenConns=1" branch at `sqlite.go:60` never runs.
The probe reported `MaxOpenConnections=25`.
- With `:memory:` (the documented test setup), each pooled connection opens its
**own** private database. The probe created a table on one connection, and a
second connection reported `no such table: t`. `ConnMaxIdleTime` (default
5 min) then closes idle connections and their data with them.
- `PRAGMA journal_mode=WAL` and `PRAGMA busy_timeout` are `Exec`'d once on
whichever pooled connection runs them. `busy_timeout` is per-connection, so
the other 24 get `database is locked` immediately under write contention.
- `BuildDSN` adds `?_timeout=<ms>` (`config.go:347-351`), but
`glebarez/go-sqlite` only recognises `_pragma`, `_txlock` and `_time_format`,
so this parameter is silently ignored.
- `SQLiteProvider.reconnectDB` (`sqlite.go:165`) needs a `dbFactory` that
nothing ever sets, so it is dead code.
**Recommendation.** For SQLite, force `MaxOpenConns=1` for `:memory:` (or use
`file::memory:?cache=shared`), and never set an idle timeout there. Pass the
pragmas in the DSN (`_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)`) so
every connection gets them.
---
### 9. Medium — Partial-failure states in `Close` and `Connect`
- `connection.go:137-149`: if `bunDB.Close()` or `provider.Close()` fails, for
example because the listener's Close failed (finding 5), `Close` returns
early with `connected = true` and the pool already closed. Every accessor then
returns a handle to a closed pool until someone calls `Close` again.
- `manager.go:197-231`: if connection *k* of *n* fails to connect, `Connect`
returns an error. Connections 1…k-1 stay open but are never stored in
`m.connections`, so `Close` can't reach them and they leak.
**Recommendation.** In `Close`, mark the connection disconnected and nil the
fields regardless of errors, and return a joined error. In `Connect`, close any
connections opened so far when a later one fails.
---
### 10. Medium — DSN builders don't escape credentials; TLS off by default
`config.go` `buildPostgresDSN` / `buildMSSQLDSN` / `buildMongoDSN` use
`fmt.Sprintf` with raw `User`/`Password`/`Database` values:
- Postgres key=value format: a password containing a space or `'` breaks
parsing. A password like `x sslmode=disable` *overrides earlier parameters*.
- MSSQL and Mongo URLs: `@`, `:`, `/`, `?` or `&` in the password corrupt the
URL. They need `url.QueryEscape` / `url.UserPassword`.
- `sslmode` defaults to `disable` (`config.go:322-325`); see
`_CROSS-CUTTING.audit.md` X6.
These values come from config, not from clients, so this isn't directly
exploitable by the threat model. It is a correctness and hardening problem,
and it becomes a security problem wherever DSN parts come from a tenant or
operator UI.
**Recommendation.** Build the Postgres DSN as a URL with `url.URL{User: url.UserPassword(...)}`,
or quote key=value values properly. Default `sslmode` to `prefer` or `require`.
---
### 11. Medium — Config knobs that are ignored or cannot be disabled
- `config.go:161-168`: `HealthCheckInterval == 0` and
`EnableAutoReconnect == false` are both treated as "unset" and replaced with
the defaults (15 s, `true`). **Auto-reconnect, the trigger for findings 1–2,
cannot be switched off from config.**
- `RetryAttempts`, `RetryDelay` and `RetryMaxDelay` are defaulted and copied,
but no provider reads them. Every provider hardcodes `retryAttempts := 3`
and `retryDelay := 1 * time.Second`.
- `statement_timeout` is only added when the DSN is built (finding 4), and
SQLite `_timeout` is ignored by the driver (finding 8).
**Recommendation.** Use `*bool` / `*time.Duration`, or an explicit
`Disable…` flag, for the values that can legitimately be zero or false. Wire
the retry settings into the providers, or delete them.
---
### 12. Medium — `Manager.Connect` holds the manager lock across network dials
`manager.go:197-231` holds `m.mu` (write) while dialing every configured
connection, each with up to 3 attempts, backoff, and `ConnectTimeout`.
`GetConnection`, `HealthCheck`, `Stats` and the health checker all wait
behind it. That's harmless at startup, but it serialises the whole manager if
`Connect` is ever called at runtime (hot-reload, lazy init).
**Recommendation.** Dial outside the lock, then lock only to publish the
results into `m.connections`.
---
### 13. Low — dbmanager metrics are never published
`metrics.go` defines Prometheus collectors plus `PublishMetrics` and
`RecordReconnectAttempt`. A grep over the repository finds **no callers** of
either. The connection-pool gauges (open, in-use, idle, wait count) are exactly
what would have shown the idle-connection problem, and they are always zero.
The `*_total` names are registered as gauges, not counters.
**Recommendation.** Call `PublishMetrics` from the health-check tick, call
`RecordReconnectAttempt` from `Reconnect`, and make the totals counters.
---
### 14. Low — Assorted correctness issues
- `Native()` checks `c.connected` (`connection.go:214`); `Bun()` and `GORM()`
don't. After a partial `Close` they can build ORM wrappers over a nil or
closed DB.
- `getNativeAdapter` (`connection.go:500-525`) wraps SQLite and MSSQL in
`PgSQLAdapter`, which quotes and builds SQL in Postgres dialect.
- `ExistingDBProvider` (`NewConnectionFromDB`) applies no pool settings and no
idle or lifetime limits, and its `Close` closes the caller's `*sql.DB`.
- `MongoProvider` uses `MaxIdleConns` as `MinPoolSize`, and `Stats()` returns an
empty struct.
---
### 15. Low — Logging defects
- `manager.go:247, 367-369, 378-380` call `logger.Error("…", "name", name, "error", err)`.
`pkg/logger` is printf-style, so these print `%!(EXTRA string=name, …)`, and
the error text is buried in exactly the log lines needed during an outage.
- `ResetInstance` discards the error from `Close`.
- Connection errors wrap driver errors that can include the DSN host and user.
Together with `_CROSS-CUTTING.audit.md` X8, they reach Sentry unscrubbed.
---
## Test coverage
`manager_test.go` and `factory_test.go` cover construction and config defaults.
Nothing tests `Reconnect` while handles are held, concurrent `Reconnect`, a
`Connect`/`Close` cycle run twice, health-check lock hold time, or listener
reconnection. Each of findings 1, 2, 3, 6 and 7 can be reproduced with a short
SQLite-backed test (the probes used for this audit took about 20 lines each).
Add them as regression tests when the fixes land, and run them with `-race`
(`_CROSS-CUTTING.audit.md` X1).
+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)
}
}